mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-04 12:38:38 +00:00
Compare commits
65 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 |
@@ -36,7 +36,9 @@ jobs:
|
||||
permissions:
|
||||
contents: read
|
||||
outputs:
|
||||
checker_outcome: ${{ steps.lychee.outcome }}
|
||||
exit_code: ${{ steps.lychee.outputs.exit_code }}
|
||||
status: ${{ steps.validate.outputs.status }}
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v6
|
||||
@@ -50,6 +52,7 @@ jobs:
|
||||
|
||||
- name: Check links
|
||||
id: lychee
|
||||
continue-on-error: true
|
||||
uses: lycheeverse/lychee-action@e7477775783ea5526144ba13e8db5eec57747ce8 # v2.9.0
|
||||
with:
|
||||
# Restricted to http(s) on purpose. Much of docs/src is generated
|
||||
@@ -68,38 +71,50 @@ jobs:
|
||||
format: json
|
||||
output: ./lychee/out.json
|
||||
jobSummary: false
|
||||
# The report, not a red build, is the signal for broken links. The
|
||||
# validation step below still fails the run if the check itself
|
||||
# breaks.
|
||||
# The report issue, not a red workflow run, is the signal for link
|
||||
# findings and checker failures alike.
|
||||
fail: false
|
||||
|
||||
- name: Validate report
|
||||
id: validate
|
||||
# lychee does not reserve exit code 2 for broken links: its CLI
|
||||
# parser also exits 2 on an invalid option, before any link was
|
||||
# checked or any report written. Only a parseable report whose
|
||||
# counts agree with the exit code counts as a link verdict; anything
|
||||
# else fails here, and the report job below is skipped entirely, so
|
||||
# the tracking issue is never touched. Exit 2 covers timeouts as
|
||||
# well as errors, and a timed-out host is exactly the transient
|
||||
# unavailability this report exists to surface, so both count as
|
||||
# findings. Requiring total > 0 also catches a glob that silently
|
||||
# stopped matching any file.
|
||||
if: steps.lychee.outputs.exit_code == 0 || steps.lychee.outputs.exit_code == 2
|
||||
# counts agree with a completed exit code (0 or 2) counts as a link
|
||||
# verdict. Everything else becomes a checker-error report instead of
|
||||
# failing the workflow. Exit 2 covers timeouts as well as errors, and a
|
||||
# timed-out host is exactly the transient unavailability this report
|
||||
# exists to surface, so both count as findings. Requiring total > 0
|
||||
# also catches a glob that silently stopped matching any file.
|
||||
if: always()
|
||||
env:
|
||||
CHECKER_OUTCOME: ${{ steps.lychee.outcome }}
|
||||
EXIT_CODE: ${{ steps.lychee.outputs.exit_code }}
|
||||
run: |
|
||||
jq -e --argjson code "$EXIT_CODE" '
|
||||
(.total > 0) and
|
||||
(if $code == 0
|
||||
then .errors == 0 and .timeouts == 0
|
||||
and (.error_map | length == 0) and (.timeout_map | length == 0)
|
||||
else (.errors + .timeouts) > 0
|
||||
and ((.error_map | length) + (.timeout_map | length)) > 0
|
||||
end)
|
||||
' ./lychee/out.json
|
||||
status=checker-error
|
||||
if [[ "$CHECKER_OUTCOME" == success ]] &&
|
||||
[[ "$EXIT_CODE" == 0 || "$EXIT_CODE" == 2 ]] &&
|
||||
jq -e --argjson code "$EXIT_CODE" '
|
||||
(.total > 0) and
|
||||
(if $code == 0
|
||||
then .errors == 0 and .timeouts == 0
|
||||
and (.error_map | length == 0) and (.timeout_map | length == 0)
|
||||
else (.errors + .timeouts) > 0
|
||||
and ((.error_map | length) + (.timeout_map | length)) > 0
|
||||
end)
|
||||
' ./lychee/out.json
|
||||
then
|
||||
if [[ "$EXIT_CODE" == 0 ]]; then
|
||||
status=healthy
|
||||
else
|
||||
status=findings
|
||||
fi
|
||||
fi
|
||||
echo "status=$status" >> "$GITHUB_OUTPUT"
|
||||
echo "Validated link check as $status"
|
||||
|
||||
- name: Upload report
|
||||
if: steps.lychee.outputs.exit_code == 2
|
||||
if: steps.validate.outputs.status == 'findings'
|
||||
uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: link-report
|
||||
@@ -115,26 +130,11 @@ jobs:
|
||||
permissions:
|
||||
issues: write
|
||||
env:
|
||||
CHECKER_OUTCOME: ${{ needs.scan.outputs.checker_outcome }}
|
||||
EXIT_CODE: ${{ needs.scan.outputs.exit_code }}
|
||||
STATUS: ${{ needs.scan.outputs.status }}
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
steps:
|
||||
- name: Classify checker result
|
||||
# lychee exits 0 when every link resolves and 2 when links fail,
|
||||
# both already cross-checked against the report by the scan job's
|
||||
# validation step. Anything else (1 runtime, 3 bad config) means the
|
||||
# check never produced a link verdict, which must surface as a failed
|
||||
# run rather than be published as "broken documentation links".
|
||||
run: |
|
||||
case "$EXIT_CODE" in
|
||||
0|2)
|
||||
echo "lychee exit code $EXIT_CODE"
|
||||
;;
|
||||
*)
|
||||
echo "::error::lychee exited with '$EXIT_CODE': the link check did not complete. Leaving the report issue untouched."
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
- name: Find existing report issue
|
||||
id: report
|
||||
# Matched on title alone, and through search rather than a listing:
|
||||
@@ -144,7 +144,7 @@ jobs:
|
||||
# Closed issues are included because a healthy run closes the report:
|
||||
# an open-only lookup would forget that identity and the next failing
|
||||
# run would open a duplicate. The oldest match stays the canonical
|
||||
# report and is reopened below when links break again.
|
||||
# report and is reopened below when a problem recurs.
|
||||
run: |
|
||||
match=$(gh issue list --repo "$GITHUB_REPOSITORY" --state all \
|
||||
--search "in:title \"$REPORT_TITLE\" author:app/github-actions" \
|
||||
@@ -154,14 +154,14 @@ jobs:
|
||||
echo "state=$(jq -r '.state // empty' <<<"$match")" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Download report
|
||||
if: env.EXIT_CODE == 2
|
||||
if: env.STATUS == 'findings'
|
||||
uses: actions/download-artifact@v8
|
||||
with:
|
||||
name: link-report
|
||||
path: ./lychee
|
||||
|
||||
- name: Compose report
|
||||
if: env.EXIT_CODE == 2
|
||||
if: env.STATUS == 'findings'
|
||||
run: |
|
||||
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
|
||||
{
|
||||
@@ -185,22 +185,41 @@ jobs:
|
||||
' ./lychee/out.json
|
||||
} > ./lychee/issue.md
|
||||
|
||||
- name: Compose checker error report
|
||||
if: env.STATUS == 'checker-error'
|
||||
run: |
|
||||
mkdir -p ./lychee
|
||||
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
|
||||
{
|
||||
echo "The documentation link check did not complete in [the latest run]($run_url)."
|
||||
echo
|
||||
echo "This issue is rewritten by every scheduled run and closed automatically once a trustworthy run finds that all links resolve."
|
||||
echo
|
||||
echo "The checker did not produce a trustworthy link verdict. Treat the previous result, if any, as stale until a later run completes."
|
||||
echo
|
||||
echo "* Action outcome: \`$CHECKER_OUTCOME\`"
|
||||
echo "* Exit code: \`${EXIT_CODE:-not reported}\`"
|
||||
echo "* Verdict validation: \`failed\`"
|
||||
} > ./lychee/issue.md
|
||||
|
||||
- name: Reopen report issue
|
||||
# A healthy run closes the report, and the issue action below only
|
||||
# rewrites the body of whatever number it is given. Without an
|
||||
# explicit reopen, the 2 -> 0 -> 2 sequence would keep rewriting a
|
||||
# closed issue while links are broken. A CLOSED state implies the
|
||||
# lookup found a canonical issue, so no separate emptiness check.
|
||||
if: env.EXIT_CODE == 2 && steps.report.outputs.state == 'CLOSED'
|
||||
# explicit reopen, a later finding or checker error would rewrite a
|
||||
# closed issue. A CLOSED state implies the lookup found a canonical
|
||||
# issue, so no separate emptiness check.
|
||||
if: >-
|
||||
env.STATUS != 'healthy' &&
|
||||
steps.report.outputs.state == 'CLOSED'
|
||||
env:
|
||||
ISSUE_NUMBER: ${{ steps.report.outputs.number }}
|
||||
run: |
|
||||
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
|
||||
gh issue reopen "$ISSUE_NUMBER" --repo "$GITHUB_REPOSITORY" \
|
||||
--comment "Broken documentation links found again in [the latest run]($run_url)."
|
||||
--comment "The documentation link checker reported a problem again in [the latest run]($run_url)."
|
||||
|
||||
- name: Report broken links
|
||||
if: env.EXIT_CODE == 2
|
||||
- name: Report link-check problem
|
||||
if: env.STATUS != 'healthy'
|
||||
uses: peter-evans/create-issue-from-file@fca9117c27cdc29c6c4db3b86c48e4115a786710 # v6.0.0
|
||||
with:
|
||||
# Empty on the first failing run, which creates the issue; afterwards
|
||||
@@ -213,7 +232,9 @@ jobs:
|
||||
- name: Close report issue once links are healthy
|
||||
# An OPEN state implies the lookup found a canonical issue; a report
|
||||
# that is already closed needs nothing.
|
||||
if: env.EXIT_CODE == 0 && steps.report.outputs.state == 'OPEN'
|
||||
if: >-
|
||||
env.STATUS == 'healthy' &&
|
||||
steps.report.outputs.state == 'OPEN'
|
||||
env:
|
||||
ISSUE_NUMBER: ${{ steps.report.outputs.number }}
|
||||
run: |
|
||||
|
||||
Generated
+43
-43
@@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
|
||||
|
||||
[[package]]
|
||||
name = "fsst"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"rand 0.9.5",
|
||||
@@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a"
|
||||
|
||||
[[package]]
|
||||
name = "lance"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arc-swap",
|
||||
"arrow",
|
||||
@@ -4890,8 +4890,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-arrow"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4913,7 +4913,7 @@ dependencies = [
|
||||
[[package]]
|
||||
name = "lance-arrow-scalar"
|
||||
version = "58.0.0"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4927,7 +4927,7 @@ dependencies = [
|
||||
[[package]]
|
||||
name = "lance-arrow-stats"
|
||||
version = "58.0.0"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -4936,8 +4936,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-bitpacking"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrayref",
|
||||
"crunchy",
|
||||
@@ -4947,8 +4947,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-core"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -4988,8 +4988,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-datafusion"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5019,8 +5019,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-datagen"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5037,8 +5037,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-derive"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
@@ -5047,8 +5047,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-encoding"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-arith",
|
||||
"arrow-array",
|
||||
@@ -5082,8 +5082,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-file"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-arith",
|
||||
"arrow-array",
|
||||
@@ -5114,8 +5114,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-index"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arc-swap",
|
||||
"arrow",
|
||||
@@ -5182,8 +5182,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-index-core"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -5205,8 +5205,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-io"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5242,8 +5242,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-linalg"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -5259,8 +5259,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-namespace"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"async-trait",
|
||||
@@ -5272,8 +5272,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-namespace-impls"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-ipc",
|
||||
@@ -5326,8 +5326,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-select"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-buffer",
|
||||
@@ -5342,8 +5342,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-table"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow",
|
||||
"arrow-array",
|
||||
@@ -5383,8 +5383,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-testing"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
@@ -5397,8 +5397,8 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "lance-tokenizer"
|
||||
version = "11.0.0-beta.3"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.3#f7d475539cefbd140cc46a828f3d843e68cd10f1"
|
||||
version = "11.0.0-beta.4"
|
||||
source = "git+https://github.com/lance-format/lance.git?rev=11007a4c36fe3a8ecf236511b4ed034c6a980d34#11007a4c36fe3a8ecf236511b4ed034c6a980d34"
|
||||
dependencies = [
|
||||
"frostem",
|
||||
"icu_segmenter",
|
||||
@@ -5432,6 +5432,7 @@ dependencies = [
|
||||
"aws-sdk-kms",
|
||||
"aws-sdk-s3",
|
||||
"aws-smithy-runtime",
|
||||
"base64 0.22.1",
|
||||
"bytes",
|
||||
"candle-core",
|
||||
"candle-nn",
|
||||
@@ -5447,7 +5448,6 @@ dependencies = [
|
||||
"datafusion-physical-plan",
|
||||
"datafusion-sql",
|
||||
"futures",
|
||||
"goosefs-sdk",
|
||||
"half",
|
||||
"hf-hub",
|
||||
"http 1.5.0",
|
||||
|
||||
+14
-14
@@ -13,20 +13,20 @@ categories = ["database-implementations"]
|
||||
rust-version = "1.91.0"
|
||||
|
||||
[workspace.dependencies]
|
||||
lance = { "version" = "=11.0.0-beta.3", default-features = false, "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { "version" = "=11.0.0-beta.3", default-features = false, "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { "version" = "=11.0.0-beta.3", default-features = false, "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { "version" = "=11.0.0-beta.3", "tag" = "v11.0.0-beta.3", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
|
||||
ahash = "0.8"
|
||||
# Note that this one does not include pyarrow
|
||||
arrow = { version = "58.0.0", optional = false }
|
||||
|
||||
@@ -6,7 +6,9 @@ import * as arrow17 from "apache-arrow-17";
|
||||
import * as arrow18 from "apache-arrow-18";
|
||||
|
||||
import {
|
||||
Vector as CurrentVector,
|
||||
convertToTable,
|
||||
tableFromIPC as currentTableFromIPC,
|
||||
fromBufferToRecordBatch,
|
||||
fromDataToBuffer,
|
||||
fromRecordBatchToBuffer,
|
||||
@@ -19,6 +21,7 @@ import {
|
||||
FunctionOptions,
|
||||
} from "../lancedb/embedding/embedding_function";
|
||||
import { EmbeddingFunctionConfig } from "../lancedb/embedding/registry";
|
||||
import { sanitizeTable } from "../lancedb/sanitize";
|
||||
|
||||
// biome-ignore lint/suspicious/noExplicitAny: skip
|
||||
function sampleRecords(): Array<Record<string, any>> {
|
||||
@@ -64,7 +67,11 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
|
||||
tableFromIPC,
|
||||
DataType,
|
||||
Dictionary,
|
||||
RecordBatch: ArrowRecordBatch,
|
||||
Table: ArrowTable,
|
||||
Uint8: ArrowUint8,
|
||||
makeData: arrowMakeData,
|
||||
vectorFromArray,
|
||||
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
|
||||
} = <any>arrow;
|
||||
type Schema = ApacheArrow["Schema"];
|
||||
@@ -1054,6 +1061,114 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
|
||||
});
|
||||
|
||||
describe("when using two versions of arrow", function () {
|
||||
it("preserves a dictionary shared by multiple fields", async function () {
|
||||
const values = ["alpha", "beta", "alpha"];
|
||||
const dictionaryVector = vectorFromArray(values);
|
||||
const batch = new ArrowRecordBatch({
|
||||
first: dictionaryVector.data[0],
|
||||
second: dictionaryVector.data[0],
|
||||
});
|
||||
const table = new ArrowTable([batch]);
|
||||
|
||||
const sanitized = sanitizeTable(table);
|
||||
expect([...sanitized.getChild("first")!]).toEqual(values);
|
||||
expect([...sanitized.getChild("second")!]).toEqual(values);
|
||||
const firstType = sanitized.schema.fields[0].type as {
|
||||
dictionary: unknown;
|
||||
};
|
||||
const secondType = sanitized.schema.fields[1].type as {
|
||||
dictionary: unknown;
|
||||
};
|
||||
expect(secondType.dictionary).toBe(firstType.dictionary);
|
||||
expect(sanitized.batches[0].data.children[1].dictionary).toBe(
|
||||
sanitized.batches[0].data.children[0].dictionary,
|
||||
);
|
||||
|
||||
const buf = await fromDataToBuffer(table);
|
||||
const actual = currentTableFromIPC(buf);
|
||||
expect([...actual.getChild("first")!]).toEqual(values);
|
||||
expect([...actual.getChild("second")!]).toEqual(values);
|
||||
});
|
||||
|
||||
it("preserves shared dictionary data from another Arrow version", async function () {
|
||||
const values = ["alpha", "beta", "alpha"];
|
||||
const dictionaryVector = vectorFromArray(values);
|
||||
const firstBatch = new ArrowRecordBatch({
|
||||
label: dictionaryVector.slice(0, 2).data[0],
|
||||
});
|
||||
const secondBatch = new ArrowRecordBatch({
|
||||
label: dictionaryVector.slice(2).data[0],
|
||||
});
|
||||
const table = new ArrowTable([firstBatch, secondBatch]);
|
||||
|
||||
const sanitized = sanitizeTable(table);
|
||||
expect([...sanitized.getChild("label")!]).toEqual(values);
|
||||
|
||||
const dictionaries = sanitized.batches.map(
|
||||
(batch) => batch.data.children[0].dictionary,
|
||||
);
|
||||
expect(dictionaries[0]).toBeInstanceOf(CurrentVector);
|
||||
expect(dictionaries[1]).toBe(dictionaries[0]);
|
||||
|
||||
const buf = await fromDataToBuffer(table);
|
||||
const actual = currentTableFromIPC(buf);
|
||||
expect([...actual.getChild("label")!]).toEqual(values);
|
||||
});
|
||||
|
||||
it("preserves shared chunks in growing dictionaries", async function () {
|
||||
const type = new Dictionary(new Utf8(), new Int32(), 42, false);
|
||||
const firstDictionary = vectorFromArray(["alpha", "beta"], new Utf8());
|
||||
const secondDictionary = firstDictionary.concat(
|
||||
vectorFromArray(["gamma"], new Utf8()),
|
||||
);
|
||||
const firstData = arrowMakeData({
|
||||
type,
|
||||
data: Int32Array.from([0, 1]),
|
||||
dictionary: firstDictionary,
|
||||
});
|
||||
const secondData = arrowMakeData({
|
||||
type,
|
||||
data: Int32Array.from([2]),
|
||||
dictionary: secondDictionary,
|
||||
});
|
||||
const table = new ArrowTable([
|
||||
new ArrowRecordBatch({ label: firstData }),
|
||||
new ArrowRecordBatch({ label: secondData }),
|
||||
]);
|
||||
|
||||
const sanitized = sanitizeTable(table);
|
||||
const expected = ["alpha", "beta", "gamma"];
|
||||
expect([...sanitized.getChild("label")!]).toEqual(expected);
|
||||
const firstLocalDictionary =
|
||||
sanitized.batches[0].data.children[0].dictionary!;
|
||||
const secondLocalDictionary =
|
||||
sanitized.batches[1].data.children[0].dictionary!;
|
||||
expect(secondLocalDictionary.data[0]).toBe(
|
||||
firstLocalDictionary.data[0],
|
||||
);
|
||||
|
||||
const buf = await fromTableToBuffer(sanitized);
|
||||
const actual = currentTableFromIPC(buf);
|
||||
expect([...actual.getChild("label")!]).toEqual(expected);
|
||||
});
|
||||
|
||||
it("can serialize list data from another Arrow version", async function () {
|
||||
const values = [["anime", "action"], [], null];
|
||||
const vector = vectorFromArray(
|
||||
values,
|
||||
new List(new Field("item", new Utf8(), true)),
|
||||
);
|
||||
const table = new ArrowTable({ tags: vector });
|
||||
|
||||
const buf = await fromDataToBuffer(table);
|
||||
const actual = currentTableFromIPC(buf);
|
||||
const actualTags = actual.getChild("tags");
|
||||
|
||||
expect(actualTags?.get(0)?.toJSON()).toEqual(values[0]);
|
||||
expect(actualTags?.get(1)?.toJSON()).toEqual(values[1]);
|
||||
expect(actualTags?.get(2)).toBeNull();
|
||||
});
|
||||
|
||||
it("can still import data", async function () {
|
||||
const schema = new arrow15.Schema([
|
||||
new arrow15.Field("id", new arrow15.Int32()),
|
||||
|
||||
+174
-29
@@ -9,7 +9,7 @@
|
||||
// comes from the exact same library instance. This is not always the case
|
||||
// and so we must sanitize the input to ensure that it is compatible.
|
||||
|
||||
import { BufferType, Data } from "apache-arrow";
|
||||
import { BufferType, Data, Vector } from "apache-arrow";
|
||||
import type { IntBitWidth, TKeys, TimeBitWidth } from "apache-arrow/type";
|
||||
import {
|
||||
Binary,
|
||||
@@ -74,6 +74,20 @@ import {
|
||||
Utf8,
|
||||
} from "./arrow";
|
||||
|
||||
type SanitizationContext = {
|
||||
types: WeakMap<object, DataType>;
|
||||
vectors: WeakMap<object, Vector>;
|
||||
data: WeakMap<object, Data<DataType>>;
|
||||
};
|
||||
|
||||
function createSanitizationContext(): SanitizationContext {
|
||||
return {
|
||||
types: new WeakMap(),
|
||||
vectors: new WeakMap(),
|
||||
data: new WeakMap(),
|
||||
};
|
||||
}
|
||||
|
||||
export function sanitizeMetadata(
|
||||
metadataLike?: unknown,
|
||||
): Map<string, string> | undefined {
|
||||
@@ -186,6 +200,13 @@ export function sanitizeInterval(typeLike: object) {
|
||||
}
|
||||
|
||||
export function sanitizeList(typeLike: object) {
|
||||
return sanitizeListWithContext(typeLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeListWithContext(
|
||||
typeLike: object,
|
||||
context: SanitizationContext,
|
||||
) {
|
||||
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
|
||||
throw Error(
|
||||
"Expected a List type to have an array-like `children` property",
|
||||
@@ -194,19 +215,35 @@ export function sanitizeList(typeLike: object) {
|
||||
if (typeLike.children.length !== 1) {
|
||||
throw Error("Expected a List type to have exactly one child");
|
||||
}
|
||||
return new List(sanitizeField(typeLike.children[0]));
|
||||
return new List(sanitizeFieldWithContext(typeLike.children[0], context));
|
||||
}
|
||||
|
||||
export function sanitizeStruct(typeLike: object) {
|
||||
return sanitizeStructWithContext(typeLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeStructWithContext(
|
||||
typeLike: object,
|
||||
context: SanitizationContext,
|
||||
) {
|
||||
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
|
||||
throw Error(
|
||||
"Expected a Struct type to have an array-like `children` property",
|
||||
);
|
||||
}
|
||||
return new Struct(typeLike.children.map((child) => sanitizeField(child)));
|
||||
return new Struct(
|
||||
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
|
||||
);
|
||||
}
|
||||
|
||||
export function sanitizeUnion(typeLike: object) {
|
||||
return sanitizeUnionWithContext(typeLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeUnionWithContext(
|
||||
typeLike: object,
|
||||
context: SanitizationContext,
|
||||
) {
|
||||
if (
|
||||
!("typeIds" in typeLike) ||
|
||||
!("mode" in typeLike) ||
|
||||
@@ -226,7 +263,7 @@ export function sanitizeUnion(typeLike: object) {
|
||||
typeLike.mode,
|
||||
// biome-ignore lint/suspicious/noExplicitAny: skip
|
||||
typeLike.typeIds as any,
|
||||
typeLike.children.map((child) => sanitizeField(child)),
|
||||
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
|
||||
);
|
||||
}
|
||||
|
||||
@@ -234,6 +271,19 @@ export function sanitizeTypedUnion(
|
||||
typeLike: object,
|
||||
// eslint-disable-next-line @typescript-eslint/naming-convention
|
||||
UnionType: typeof DenseUnion | typeof SparseUnion,
|
||||
) {
|
||||
return sanitizeTypedUnionWithContext(
|
||||
typeLike,
|
||||
UnionType,
|
||||
createSanitizationContext(),
|
||||
);
|
||||
}
|
||||
|
||||
function sanitizeTypedUnionWithContext(
|
||||
typeLike: object,
|
||||
// eslint-disable-next-line @typescript-eslint/naming-convention
|
||||
UnionType: typeof DenseUnion | typeof SparseUnion,
|
||||
context: SanitizationContext,
|
||||
) {
|
||||
if (!("typeIds" in typeLike)) {
|
||||
throw Error(
|
||||
@@ -248,7 +298,7 @@ export function sanitizeTypedUnion(
|
||||
|
||||
return new UnionType(
|
||||
typeLike.typeIds as Int32Array | number[],
|
||||
typeLike.children.map((child) => sanitizeField(child)),
|
||||
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
|
||||
);
|
||||
}
|
||||
|
||||
@@ -262,6 +312,16 @@ export function sanitizeFixedSizeBinary(typeLike: object) {
|
||||
}
|
||||
|
||||
export function sanitizeFixedSizeList(typeLike: object) {
|
||||
return sanitizeFixedSizeListWithContext(
|
||||
typeLike,
|
||||
createSanitizationContext(),
|
||||
);
|
||||
}
|
||||
|
||||
function sanitizeFixedSizeListWithContext(
|
||||
typeLike: object,
|
||||
context: SanitizationContext,
|
||||
) {
|
||||
if (!("listSize" in typeLike) || typeof typeLike.listSize !== "number") {
|
||||
throw Error("Expected a FixedSizeList type to have a `listSize` property");
|
||||
}
|
||||
@@ -275,11 +335,18 @@ export function sanitizeFixedSizeList(typeLike: object) {
|
||||
}
|
||||
return new FixedSizeList(
|
||||
typeLike.listSize,
|
||||
sanitizeField(typeLike.children[0]),
|
||||
sanitizeFieldWithContext(typeLike.children[0], context),
|
||||
);
|
||||
}
|
||||
|
||||
export function sanitizeMap(typeLike: object) {
|
||||
return sanitizeMapWithContext(typeLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeMapWithContext(
|
||||
typeLike: object,
|
||||
context: SanitizationContext,
|
||||
) {
|
||||
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
|
||||
throw Error(
|
||||
"Expected a Map type to have an array-like `children` property",
|
||||
@@ -292,7 +359,10 @@ export function sanitizeMap(typeLike: object) {
|
||||
throw Error("Expected a Map type to have exactly one child");
|
||||
}
|
||||
|
||||
return new Map_(sanitizeField(typeLike.children[0]), typeLike.keysSorted);
|
||||
return new Map_(
|
||||
sanitizeFieldWithContext(typeLike.children[0], context),
|
||||
typeLike.keysSorted,
|
||||
);
|
||||
}
|
||||
|
||||
export function sanitizeDuration(typeLike: object) {
|
||||
@@ -303,6 +373,13 @@ export function sanitizeDuration(typeLike: object) {
|
||||
}
|
||||
|
||||
export function sanitizeDictionary(typeLike: object) {
|
||||
return sanitizeDictionaryWithContext(typeLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeDictionaryWithContext(
|
||||
typeLike: object,
|
||||
context: SanitizationContext,
|
||||
) {
|
||||
if (!("id" in typeLike) || typeof typeLike.id !== "number") {
|
||||
throw Error("Expected a Dictionary type to have an `id` property");
|
||||
}
|
||||
@@ -316,8 +393,8 @@ export function sanitizeDictionary(typeLike: object) {
|
||||
throw Error("Expected a Dictionary type to have an `isOrdered` property");
|
||||
}
|
||||
return new Dictionary(
|
||||
sanitizeType(typeLike.dictionary),
|
||||
sanitizeType(typeLike.indices) as TKeys,
|
||||
sanitizeTypeWithContext(typeLike.dictionary, context),
|
||||
sanitizeTypeWithContext(typeLike.indices, context) as TKeys,
|
||||
typeLike.id,
|
||||
typeLike.isOrdered,
|
||||
);
|
||||
@@ -325,12 +402,23 @@ export function sanitizeDictionary(typeLike: object) {
|
||||
|
||||
// biome-ignore lint/suspicious/noExplicitAny: skip
|
||||
export function sanitizeType(typeLike: unknown): DataType<any> {
|
||||
return sanitizeTypeWithContext(typeLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeTypeWithContext(
|
||||
typeLike: unknown,
|
||||
context: SanitizationContext,
|
||||
): DataType {
|
||||
if (typeof typeLike === "string") {
|
||||
return dataTypeFromName(typeLike);
|
||||
}
|
||||
if (typeof typeLike !== "object" || typeLike === null) {
|
||||
throw Error("Expected a Type but object was null/undefined");
|
||||
}
|
||||
const cached = context.types.get(typeLike);
|
||||
if (cached !== undefined) {
|
||||
return cached;
|
||||
}
|
||||
if (
|
||||
!("typeId" in typeLike) ||
|
||||
!(
|
||||
@@ -349,6 +437,16 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
|
||||
throw Error("Type's typeId property was not a function or number");
|
||||
}
|
||||
|
||||
const type = sanitizeTypeById(typeLike, typeId, context);
|
||||
context.types.set(typeLike, type);
|
||||
return type;
|
||||
}
|
||||
|
||||
function sanitizeTypeById(
|
||||
typeLike: object,
|
||||
typeId: Type,
|
||||
context: SanitizationContext,
|
||||
): DataType {
|
||||
switch (typeId) {
|
||||
case Type.NONE:
|
||||
throw Error("Received a Type with a typeId of NONE");
|
||||
@@ -375,21 +473,21 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
|
||||
case Type.Interval:
|
||||
return sanitizeInterval(typeLike);
|
||||
case Type.List:
|
||||
return sanitizeList(typeLike);
|
||||
return sanitizeListWithContext(typeLike, context);
|
||||
case Type.Struct:
|
||||
return sanitizeStruct(typeLike);
|
||||
return sanitizeStructWithContext(typeLike, context);
|
||||
case Type.Union:
|
||||
return sanitizeUnion(typeLike);
|
||||
return sanitizeUnionWithContext(typeLike, context);
|
||||
case Type.FixedSizeBinary:
|
||||
return sanitizeFixedSizeBinary(typeLike);
|
||||
case Type.FixedSizeList:
|
||||
return sanitizeFixedSizeList(typeLike);
|
||||
return sanitizeFixedSizeListWithContext(typeLike, context);
|
||||
case Type.Map:
|
||||
return sanitizeMap(typeLike);
|
||||
return sanitizeMapWithContext(typeLike, context);
|
||||
case Type.Duration:
|
||||
return sanitizeDuration(typeLike);
|
||||
case Type.Dictionary:
|
||||
return sanitizeDictionary(typeLike);
|
||||
return sanitizeDictionaryWithContext(typeLike, context);
|
||||
case Type.Int8:
|
||||
return new Int8();
|
||||
case Type.Int16:
|
||||
@@ -433,9 +531,9 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
|
||||
case Type.TimestampSecond:
|
||||
return sanitizeTypedTimestamp(typeLike, TimestampSecond);
|
||||
case Type.DenseUnion:
|
||||
return sanitizeTypedUnion(typeLike, DenseUnion);
|
||||
return sanitizeTypedUnionWithContext(typeLike, DenseUnion, context);
|
||||
case Type.SparseUnion:
|
||||
return sanitizeTypedUnion(typeLike, SparseUnion);
|
||||
return sanitizeTypedUnionWithContext(typeLike, SparseUnion, context);
|
||||
case Type.IntervalDayTime:
|
||||
return new IntervalDayTime();
|
||||
case Type.IntervalYearMonth:
|
||||
@@ -454,6 +552,13 @@ export function sanitizeType(typeLike: unknown): DataType<any> {
|
||||
}
|
||||
|
||||
export function sanitizeField(fieldLike: unknown): Field {
|
||||
return sanitizeFieldWithContext(fieldLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeFieldWithContext(
|
||||
fieldLike: unknown,
|
||||
context: SanitizationContext,
|
||||
): Field {
|
||||
if (fieldLike instanceof Field) {
|
||||
return fieldLike;
|
||||
}
|
||||
@@ -471,7 +576,7 @@ export function sanitizeField(fieldLike: unknown): Field {
|
||||
}
|
||||
let type: DataType;
|
||||
try {
|
||||
type = sanitizeType(fieldLike.type);
|
||||
type = sanitizeTypeWithContext(fieldLike.type, context);
|
||||
} catch (error: unknown) {
|
||||
throw Error(
|
||||
`Unable to sanitize type for field: ${fieldLike.name} due to error: ${error}`,
|
||||
@@ -501,6 +606,13 @@ export function sanitizeField(fieldLike: unknown): Field {
|
||||
* than lancedb is using.
|
||||
*/
|
||||
export function sanitizeSchema(schemaLike: SchemaLike): Schema {
|
||||
return sanitizeSchemaWithContext(schemaLike, createSanitizationContext());
|
||||
}
|
||||
|
||||
function sanitizeSchemaWithContext(
|
||||
schemaLike: SchemaLike,
|
||||
context: SanitizationContext,
|
||||
): Schema {
|
||||
if (schemaLike instanceof Schema) {
|
||||
return schemaLike;
|
||||
}
|
||||
@@ -522,7 +634,7 @@ export function sanitizeSchema(schemaLike: SchemaLike): Schema {
|
||||
);
|
||||
}
|
||||
const sanitizedFields = schemaLike.fields.map((field) =>
|
||||
sanitizeField(field),
|
||||
sanitizeFieldWithContext(field, context),
|
||||
);
|
||||
return new Schema(sanitizedFields, metadata);
|
||||
}
|
||||
@@ -544,13 +656,18 @@ export function sanitizeTable(tableLike: TableLike): Table {
|
||||
"The table passed in does not appear to be a table (no 'columns' property)",
|
||||
);
|
||||
}
|
||||
const schema = sanitizeSchema(tableLike.schema);
|
||||
|
||||
const batches = tableLike.batches.map(sanitizeRecordBatch);
|
||||
const context = createSanitizationContext();
|
||||
const schema = sanitizeSchemaWithContext(tableLike.schema, context);
|
||||
const batches = tableLike.batches.map((batch) =>
|
||||
sanitizeRecordBatch(batch, context),
|
||||
);
|
||||
return new Table(schema, batches);
|
||||
}
|
||||
|
||||
function sanitizeRecordBatch(batchLike: RecordBatchLike): RecordBatch {
|
||||
function sanitizeRecordBatch(
|
||||
batchLike: RecordBatchLike,
|
||||
context: SanitizationContext,
|
||||
): RecordBatch {
|
||||
if (batchLike instanceof RecordBatch) {
|
||||
return batchLike;
|
||||
}
|
||||
@@ -567,19 +684,43 @@ function sanitizeRecordBatch(batchLike: RecordBatchLike): RecordBatch {
|
||||
"The record batch passed in does not appear to be a record batch (no 'data' property)",
|
||||
);
|
||||
}
|
||||
const schema = sanitizeSchema(batchLike.schema);
|
||||
const data = sanitizeData(batchLike.data);
|
||||
const schema = sanitizeSchemaWithContext(batchLike.schema, context);
|
||||
const data = sanitizeData(batchLike.data, context) as Data<Struct>;
|
||||
return new RecordBatch(schema, data);
|
||||
}
|
||||
|
||||
type DictionaryVectorLike = {
|
||||
data: readonly DataLike[];
|
||||
};
|
||||
|
||||
type DictionaryDataLike = DataLike & {
|
||||
dictionary?: DictionaryVectorLike;
|
||||
};
|
||||
|
||||
function sanitizeData(
|
||||
dataLike: DataLike,
|
||||
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
|
||||
): import("apache-arrow").Data<Struct<any>> {
|
||||
context: SanitizationContext,
|
||||
): Data<DataType> {
|
||||
if (dataLike instanceof Data) {
|
||||
return dataLike;
|
||||
}
|
||||
return new Data(
|
||||
dataLike.type,
|
||||
const cachedData = context.data.get(dataLike);
|
||||
if (cachedData !== undefined) {
|
||||
return cachedData;
|
||||
}
|
||||
const dictionaryLike = (dataLike as DictionaryDataLike).dictionary;
|
||||
let dictionary: Vector | undefined;
|
||||
if (dictionaryLike !== undefined) {
|
||||
dictionary = context.vectors.get(dictionaryLike);
|
||||
if (dictionary === undefined) {
|
||||
dictionary = new Vector(
|
||||
dictionaryLike.data.map((data) => sanitizeData(data, context)),
|
||||
);
|
||||
context.vectors.set(dictionaryLike, dictionary);
|
||||
}
|
||||
}
|
||||
const data = new Data(
|
||||
sanitizeTypeWithContext(dataLike.type, context),
|
||||
dataLike.offset,
|
||||
dataLike.length,
|
||||
dataLike.nullCount,
|
||||
@@ -589,7 +730,11 @@ function sanitizeData(
|
||||
[BufferType.VALIDITY]: dataLike.nullBitmap,
|
||||
[BufferType.TYPE]: dataLike.typeIds,
|
||||
},
|
||||
dataLike.children.map((child) => sanitizeData(child, context)),
|
||||
dictionary,
|
||||
);
|
||||
context.data.set(dataLike, data);
|
||||
return data;
|
||||
}
|
||||
|
||||
const constructorsByTypeName = {
|
||||
|
||||
+11
-1
@@ -42,9 +42,19 @@ impl Job {
|
||||
}
|
||||
|
||||
/// Wait until the operation reaches a terminal state.
|
||||
///
|
||||
/// Jobs that complete without a resource result resolve successfully.
|
||||
/// Resource results are not exposed on this binding yet; unsupported
|
||||
/// success results reject with a generic error.
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn wait(&self) -> napi::Result<()> {
|
||||
self.inner.wait().await.default_error()
|
||||
match self.inner.wait().await.default_error()? {
|
||||
lancedb::JobResult::None => Ok(()),
|
||||
// JobResult is non_exhaustive; Function and future variants fail closed.
|
||||
_ => Err(napi::Error::from_reason(
|
||||
"unsupported job result".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Request cancellation. Cancelling a finished operation is a no-op.
|
||||
|
||||
@@ -12,6 +12,7 @@ __version__ = importlib.metadata.version("lancedb")
|
||||
|
||||
from ._lancedb import connect as lancedb_connect
|
||||
from ._lancedb import FtsToken
|
||||
from ._lancedb import Function
|
||||
from ._lancedb import tokenize as _tokenize
|
||||
from .common import URI, sanitize_uri
|
||||
from urllib.parse import urlparse
|
||||
@@ -23,6 +24,7 @@ from .schema import blob, vector, BlobType
|
||||
from .job import AsyncJob, Job
|
||||
from .table import AsyncTable, Table
|
||||
from .types import BaseTokenizerType
|
||||
from ._udf import FunctionCapability, udf
|
||||
from ._lancedb import Session
|
||||
from .namespace import (
|
||||
connect_namespace,
|
||||
@@ -507,6 +509,8 @@ __all__ = [
|
||||
"FtsToken",
|
||||
"col",
|
||||
"Expr",
|
||||
"Function",
|
||||
"FunctionCapability",
|
||||
"func",
|
||||
"lit",
|
||||
"URI",
|
||||
@@ -521,5 +525,6 @@ __all__ = [
|
||||
"RemoteDBConnection",
|
||||
"Session",
|
||||
"Table",
|
||||
"udf",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Private first-class Function namespace facades for database connections.
|
||||
|
||||
These helpers are internal submission and lookup surfaces. They are not durable
|
||||
resources and are not part of the public top-level export surface.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from . import _udf
|
||||
from ._lancedb import Function
|
||||
from .job import AsyncJob, Job
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .db import AsyncConnection, DBConnection
|
||||
|
||||
|
||||
class _SyncFunctions:
|
||||
"""Synchronous `db.functions` facade."""
|
||||
|
||||
__slots__ = ("_connection",)
|
||||
|
||||
def __init__(self, connection: DBConnection) -> None:
|
||||
self._connection = connection
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "_SyncFunctions()"
|
||||
|
||||
def register(self, name: str, decorated_udf: Callable[..., object]) -> Job:
|
||||
"""Register a decorated UDF and return a synchronous [Job][lancedb.job.Job]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = self._connection._submit_register_function(name, definition)
|
||||
return Job(AsyncJob(native_job))
|
||||
|
||||
def replace(
|
||||
self, name: str, current: Function, decorated_udf: Callable[..., object]
|
||||
) -> Job:
|
||||
"""Conditionally replace a Function; return sync [Job][lancedb.job.Job]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = self._connection._submit_replace_function(
|
||||
name, current, definition
|
||||
)
|
||||
return Job(AsyncJob(native_job))
|
||||
|
||||
def get(self, name: str) -> Function:
|
||||
"""Return the Function currently bound to a database-scoped name."""
|
||||
return self._connection._lookup_function_by_name(name)
|
||||
|
||||
def get_by_id(self, function_id: str) -> Function:
|
||||
"""Return the immutable Function for an exact Function ID."""
|
||||
return self._connection._lookup_function_by_id(function_id)
|
||||
|
||||
def remove(self, name: str, current: Function) -> None:
|
||||
"""Conditionally remove a Function catalog name binding."""
|
||||
return self._connection._remove_function_name(name, current)
|
||||
|
||||
def revoke(self, function: Function) -> None:
|
||||
"""Revoke an exact immutable Function by administrator set-bit."""
|
||||
return self._connection._revoke_function(function)
|
||||
|
||||
|
||||
class _AsyncFunctions:
|
||||
"""Asynchronous `async_db.functions` facade."""
|
||||
|
||||
__slots__ = ("_connection",)
|
||||
|
||||
def __init__(self, connection: AsyncConnection) -> None:
|
||||
self._connection = connection
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "_AsyncFunctions()"
|
||||
|
||||
async def register(
|
||||
self, name: str, decorated_udf: Callable[..., object]
|
||||
) -> AsyncJob:
|
||||
"""Register a decorated UDF and return an [AsyncJob][lancedb.job.AsyncJob]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = await self._connection._register_function(name, definition)
|
||||
return AsyncJob(native_job)
|
||||
|
||||
async def replace(
|
||||
self, name: str, current: Function, decorated_udf: Callable[..., object]
|
||||
) -> AsyncJob:
|
||||
"""Conditionally replace a Function; return [AsyncJob][lancedb.job.AsyncJob]."""
|
||||
definition = _udf._build_function_definition(decorated_udf)
|
||||
native_job = await self._connection._replace_function(name, current, definition)
|
||||
return AsyncJob(native_job)
|
||||
|
||||
async def get(self, name: str) -> Function:
|
||||
"""Return the Function currently bound to a database-scoped name."""
|
||||
return await self._connection._lookup_function_by_name(name)
|
||||
|
||||
async def get_by_id(self, function_id: str) -> Function:
|
||||
"""Return the immutable Function for an exact Function ID."""
|
||||
return await self._connection._lookup_function_by_id(function_id)
|
||||
|
||||
async def remove(self, name: str, current: Function) -> None:
|
||||
"""Conditionally remove a Function catalog name binding."""
|
||||
return await self._connection._remove_function_name(name, current)
|
||||
|
||||
async def revoke(self, function: Function) -> None:
|
||||
"""Revoke an exact immutable Function by administrator set-bit."""
|
||||
return await self._connection._revoke_function(function)
|
||||
@@ -153,6 +153,16 @@ class Connection(object):
|
||||
async def job_history(
|
||||
self, job_id: Optional[str] = None
|
||||
) -> List[pa.RecordBatch]: ...
|
||||
async def _register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> Job: ...
|
||||
async def _replace_function(
|
||||
self, name: str, current: Function, definition: "_FunctionDefinition"
|
||||
) -> Job: ...
|
||||
async def _lookup_function_by_name(self, name: str) -> Function: ...
|
||||
async def _lookup_function_by_id(self, function_id: str) -> Function: ...
|
||||
async def _remove_function_name(self, name: str, current: Function) -> None: ...
|
||||
async def _revoke_function(self, function: Function) -> None: ...
|
||||
async def create_table(
|
||||
self,
|
||||
name: str,
|
||||
@@ -216,11 +226,45 @@ class BlobFile:
|
||||
def read_range(self, offset: int, length: int) -> bytes: ...
|
||||
def read_up_to(self, length: int) -> bytes: ...
|
||||
|
||||
class Function:
|
||||
@property
|
||||
def id(self) -> str: ...
|
||||
@property
|
||||
def parameters(self) -> tuple[tuple[str, pa.DataType], ...]: ...
|
||||
@property
|
||||
def output_type(self) -> pa.DataType: ...
|
||||
@property
|
||||
def output_nullable(self) -> bool: ...
|
||||
def __call__(self, **kwargs: Any) -> "_FunctionCall": ...
|
||||
|
||||
class _FunctionCall:
|
||||
"""Private unresolved Function call authoring value (FF-028)."""
|
||||
|
||||
...
|
||||
|
||||
class _FunctionDefinition:
|
||||
"""Private owner of the Rust FunctionDefinition registration input."""
|
||||
|
||||
def _to_json(self) -> str: ...
|
||||
|
||||
def _new_function_definition(
|
||||
*,
|
||||
parameters: list[tuple[str, pa.DataType]],
|
||||
output_type: pa.DataType,
|
||||
output_nullable: bool,
|
||||
module: str,
|
||||
callable_name: str,
|
||||
source: str,
|
||||
python: str,
|
||||
packages: list[str],
|
||||
capabilities: list[tuple[str, str, Optional[str]]],
|
||||
) -> _FunctionDefinition: ...
|
||||
|
||||
class Job:
|
||||
@property
|
||||
def id(self) -> Optional[str]: ...
|
||||
async def status(self) -> str: ...
|
||||
async def wait(self) -> None: ...
|
||||
async def wait(self) -> Optional[Function]: ...
|
||||
async def cancel(self) -> None: ...
|
||||
|
||||
class JobInfo:
|
||||
@@ -242,6 +286,8 @@ class JobFailureInfo:
|
||||
def message(self) -> Optional[str]: ...
|
||||
@property
|
||||
def retryable(self) -> Optional[bool]: ...
|
||||
@property
|
||||
def error_code(self) -> Optional[str]: ...
|
||||
|
||||
class JobDescription:
|
||||
@property
|
||||
@@ -256,6 +302,8 @@ class JobDescription:
|
||||
def spec_json(self) -> Optional[str]: ...
|
||||
@property
|
||||
def failure(self) -> Optional[JobFailureInfo]: ...
|
||||
@property
|
||||
def result(self) -> Optional[Function]: ...
|
||||
|
||||
class Table:
|
||||
def name(self) -> str: ...
|
||||
@@ -318,6 +366,16 @@ class Table:
|
||||
name: Optional[str],
|
||||
train: Optional[bool],
|
||||
) -> Job: ...
|
||||
async def _add_generated_column(
|
||||
self, column_name: str, call: _FunctionCall
|
||||
) -> Job: ...
|
||||
async def _generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]: ...
|
||||
async def _refresh_generated_column(self, column_name: str) -> Job: ...
|
||||
async def _alter_generated_column(
|
||||
self, column_name: str, new_call: _FunctionCall
|
||||
) -> Job: ...
|
||||
async def list_versions(self) -> List[Dict[str, Any]]: ...
|
||||
async def version(self) -> int: ...
|
||||
async def checkout(self, version: Union[int, str]): ...
|
||||
|
||||
@@ -0,0 +1,538 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Local authoring declaration surface for first-class UDFs.
|
||||
|
||||
This module snapshots declaration metadata onto a Python function and privately
|
||||
validates packagable callables into a source snapshot. It does not mint durable
|
||||
identity or register anything with a database.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import inspect
|
||||
import stat
|
||||
import symtable
|
||||
import sys
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from types import CodeType, FunctionType
|
||||
from typing import NoReturn, ParamSpec, TypeVar
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from . import _lancedb
|
||||
|
||||
__all__ = ["FunctionCapability", "udf"]
|
||||
|
||||
_P = ParamSpec("_P")
|
||||
_R = TypeVar("_R")
|
||||
|
||||
_CONFIG_ATTR = "__lancedb_udf_config__"
|
||||
_SYNTHETIC_SOURCE_FILENAME = "<lancedb-udf>"
|
||||
_PACKAGING_ERROR = "udf is not packagable"
|
||||
_ALLOWED_PARAM_KINDS = (
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
inspect.Parameter.KEYWORD_ONLY,
|
||||
)
|
||||
|
||||
|
||||
class FunctionCapability:
|
||||
"""Local capability declaration for a first-class UDF.
|
||||
|
||||
Construct via :meth:`network` or :meth:`secret`. Direct construction is
|
||||
rejected so callers cannot create an uninitialized capability.
|
||||
"""
|
||||
|
||||
__slots__ = ("_kind", "_origin", "_reference", "_environment_variable")
|
||||
|
||||
def __new__(cls, *args: object, **kwargs: object) -> FunctionCapability:
|
||||
raise TypeError(
|
||||
"FunctionCapability cannot be constructed directly; "
|
||||
"use FunctionCapability.network() or FunctionCapability.secret()"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _create(
|
||||
cls,
|
||||
kind: str,
|
||||
origin: str | None,
|
||||
reference: str | None,
|
||||
environment_variable: str | None,
|
||||
) -> FunctionCapability:
|
||||
obj = object.__new__(cls)
|
||||
object.__setattr__(obj, "_kind", kind)
|
||||
object.__setattr__(obj, "_origin", origin)
|
||||
object.__setattr__(obj, "_reference", reference)
|
||||
object.__setattr__(obj, "_environment_variable", environment_variable)
|
||||
return obj
|
||||
|
||||
@classmethod
|
||||
def network(cls, origin: str) -> FunctionCapability:
|
||||
if not isinstance(origin, str):
|
||||
raise TypeError("origin must be a string")
|
||||
if origin == "":
|
||||
raise ValueError("origin must be non-empty")
|
||||
return cls._create("network", origin, None, None)
|
||||
|
||||
@classmethod
|
||||
def secret(cls, reference: str, *, environment_variable: str) -> FunctionCapability:
|
||||
if not isinstance(reference, str):
|
||||
raise TypeError("reference must be a string")
|
||||
if not isinstance(environment_variable, str):
|
||||
raise TypeError("environment_variable must be a string")
|
||||
if reference == "":
|
||||
raise ValueError("reference must be non-empty")
|
||||
if environment_variable == "":
|
||||
raise ValueError("environment_variable must be non-empty")
|
||||
return cls._create("secret", None, reference, environment_variable)
|
||||
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
return self._kind
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
return self._origin
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
return self._reference
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
return self._environment_variable
|
||||
|
||||
def __setattr__(self, name: str, value: object) -> None:
|
||||
raise AttributeError(
|
||||
f"{type(self).__name__!r} object attribute {name!r} is read-only"
|
||||
)
|
||||
|
||||
def __delattr__(self, name: str) -> None:
|
||||
raise AttributeError(
|
||||
f"{type(self).__name__!r} object attribute {name!r} is read-only"
|
||||
)
|
||||
|
||||
def __eq__(self, other: object) -> bool:
|
||||
if not isinstance(other, FunctionCapability):
|
||||
return NotImplemented
|
||||
return (
|
||||
self._kind == other._kind
|
||||
and self._origin == other._origin
|
||||
and self._reference == other._reference
|
||||
and self._environment_variable == other._environment_variable
|
||||
)
|
||||
|
||||
def __hash__(self) -> int:
|
||||
return hash(
|
||||
(
|
||||
self._kind,
|
||||
self._origin,
|
||||
self._reference,
|
||||
self._environment_variable,
|
||||
)
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
if self._kind == "network":
|
||||
return f"FunctionCapability.network({self._origin!r})"
|
||||
return (
|
||||
"FunctionCapability.secret("
|
||||
f"environment_variable={self._environment_variable!r})"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _UdfConfig:
|
||||
"""Private frozen snapshot of a ``@udf`` declaration."""
|
||||
|
||||
inputs: tuple[tuple[str, pa.DataType], ...]
|
||||
output: pa.DataType
|
||||
output_nullable: bool
|
||||
python: str
|
||||
packages: tuple[str, ...]
|
||||
capabilities: tuple[FunctionCapability, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PackagedUdf:
|
||||
"""Private frozen snapshot of a validated packagable UDF."""
|
||||
|
||||
source: str
|
||||
module: str
|
||||
callable_name: str
|
||||
config: _UdfConfig
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"_PackagedUdf(source=<redacted>, module={self.module!r}, "
|
||||
f"callable_name={self.callable_name!r}, config={self.config!r})"
|
||||
)
|
||||
|
||||
|
||||
def _validate_inputs(
|
||||
inputs: object,
|
||||
) -> tuple[tuple[str, pa.DataType], ...]:
|
||||
if not isinstance(inputs, Mapping):
|
||||
raise TypeError("udf inputs must be a Mapping of name to pyarrow DataType")
|
||||
snapshot: list[tuple[str, pa.DataType]] = []
|
||||
for key, value in inputs.items():
|
||||
if not isinstance(key, str):
|
||||
raise TypeError("udf input names must be strings")
|
||||
if key == "":
|
||||
raise ValueError("udf input names must be non-empty")
|
||||
if not isinstance(value, pa.DataType):
|
||||
raise TypeError("udf input types must be pyarrow DataType values")
|
||||
snapshot.append((key, value))
|
||||
return tuple(snapshot)
|
||||
|
||||
|
||||
def _validate_packages(packages: object) -> tuple[str, ...]:
|
||||
if isinstance(packages, (str, bytes, bytearray)):
|
||||
raise TypeError("udf packages must be a sequence of strings, not a string")
|
||||
if not isinstance(packages, Sequence):
|
||||
raise TypeError("udf packages must be a sequence of strings")
|
||||
snapshot: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for package in packages:
|
||||
if not isinstance(package, str):
|
||||
raise TypeError("udf packages must contain only strings")
|
||||
if package == "":
|
||||
raise ValueError("udf packages must be non-empty strings")
|
||||
if package in seen:
|
||||
raise ValueError(f"duplicate udf package: {package}")
|
||||
seen.add(package)
|
||||
snapshot.append(package)
|
||||
return tuple(snapshot)
|
||||
|
||||
|
||||
def _reject_non_exact_capability() -> NoReturn:
|
||||
# Exact-type only: subclasses are authoring inputs we never accept. Keep the
|
||||
# message fixed so hostile markers never enter exception text.
|
||||
raise TypeError(
|
||||
"udf capabilities must contain only FunctionCapability values"
|
||||
) from None
|
||||
|
||||
|
||||
def _require_exact_capability(capability: object) -> FunctionCapability:
|
||||
if type(capability) is not FunctionCapability:
|
||||
_reject_non_exact_capability()
|
||||
return capability
|
||||
|
||||
|
||||
def _validate_capabilities(capabilities: object) -> tuple[FunctionCapability, ...]:
|
||||
if isinstance(capabilities, (str, bytes, bytearray)):
|
||||
raise TypeError(
|
||||
"udf capabilities must be a sequence of FunctionCapability, not a string"
|
||||
)
|
||||
if not isinstance(capabilities, Sequence):
|
||||
raise TypeError("udf capabilities must be a sequence of FunctionCapability")
|
||||
return tuple(_require_exact_capability(capability) for capability in capabilities)
|
||||
|
||||
|
||||
def udf(
|
||||
*,
|
||||
inputs: Mapping[str, pa.DataType],
|
||||
output: pa.DataType,
|
||||
python: str,
|
||||
packages: Sequence[str] = (),
|
||||
output_nullable: bool = True,
|
||||
capabilities: Sequence[FunctionCapability] = (),
|
||||
) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]:
|
||||
"""Declare a local UDF without packaging or registration.
|
||||
|
||||
Applying the returned decorator attaches a private frozen config snapshot
|
||||
and returns the exact same function object.
|
||||
"""
|
||||
input_snapshot = _validate_inputs(inputs)
|
||||
if not isinstance(output, pa.DataType):
|
||||
raise TypeError("udf output must be a pyarrow DataType")
|
||||
if not isinstance(python, str):
|
||||
raise TypeError("udf python must be a string")
|
||||
if python == "":
|
||||
raise ValueError("udf python must be a non-empty string")
|
||||
package_snapshot = _validate_packages(packages)
|
||||
if not isinstance(output_nullable, bool):
|
||||
raise TypeError("udf output_nullable must be a bool")
|
||||
capability_snapshot = _validate_capabilities(capabilities)
|
||||
|
||||
config = _UdfConfig(
|
||||
inputs=input_snapshot,
|
||||
output=output,
|
||||
output_nullable=output_nullable,
|
||||
python=python,
|
||||
packages=package_snapshot,
|
||||
capabilities=capability_snapshot,
|
||||
)
|
||||
|
||||
def decorator(fn: Callable[_P, _R]) -> Callable[_P, _R]:
|
||||
if not inspect.isfunction(fn):
|
||||
raise TypeError("udf can only decorate a Python function")
|
||||
if hasattr(fn, _CONFIG_ATTR):
|
||||
raise ValueError("function is already decorated with @udf")
|
||||
setattr(fn, _CONFIG_ATTR, config)
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def _get_udf_config(fn: object) -> _UdfConfig:
|
||||
"""Return the private declaration snapshot for a ``@udf``-decorated function."""
|
||||
config = getattr(fn, _CONFIG_ATTR, None)
|
||||
if config is None:
|
||||
raise TypeError("function is not decorated with @udf")
|
||||
if not isinstance(config, _UdfConfig):
|
||||
raise TypeError("function is not decorated with @udf")
|
||||
return config
|
||||
|
||||
|
||||
def _packaging_reject() -> NoReturn:
|
||||
raise ValueError(_PACKAGING_ERROR) from None
|
||||
|
||||
|
||||
def _is_ordinary_function(fn: FunctionType) -> bool:
|
||||
if fn.__name__ == "<lambda>":
|
||||
return False
|
||||
if fn.__qualname__ != fn.__name__:
|
||||
return False
|
||||
if inspect.iscoroutinefunction(fn) or inspect.isasyncgenfunction(fn):
|
||||
return False
|
||||
if inspect.isgeneratorfunction(fn):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _resolve_source_path(fn: FunctionType, module: object) -> Path:
|
||||
try:
|
||||
fn_source: str | None = inspect.getsourcefile(fn)
|
||||
except TypeError:
|
||||
fn_source = None
|
||||
source_lookup_failed = True
|
||||
else:
|
||||
source_lookup_failed = False
|
||||
if source_lookup_failed:
|
||||
_packaging_reject()
|
||||
module_file = vars(module).get("__file__")
|
||||
if not fn_source or not isinstance(module_file, str) or module_file == "":
|
||||
_packaging_reject()
|
||||
try:
|
||||
resolved_paths: tuple[Path, Path] | None = (
|
||||
Path(fn_source).resolve(),
|
||||
Path(module_file).resolve(),
|
||||
)
|
||||
except (OSError, RuntimeError):
|
||||
resolved_paths = None
|
||||
if resolved_paths is None:
|
||||
_packaging_reject()
|
||||
fn_path, module_path = resolved_paths
|
||||
if fn_path != module_path:
|
||||
_packaging_reject()
|
||||
if fn_path.suffix != ".py":
|
||||
_packaging_reject()
|
||||
try:
|
||||
mode: int | None = fn_path.stat().st_mode
|
||||
except OSError:
|
||||
mode = None
|
||||
if mode is None:
|
||||
_packaging_reject()
|
||||
if not stat.S_ISREG(mode):
|
||||
_packaging_reject()
|
||||
return fn_path
|
||||
|
||||
|
||||
def _validate_source(
|
||||
source: str, callable_name: str
|
||||
) -> tuple[CodeType, symtable.SymbolTable]:
|
||||
try:
|
||||
module_code = compile(
|
||||
source,
|
||||
_SYNTHETIC_SOURCE_FILENAME,
|
||||
"exec",
|
||||
optimize=sys.flags.optimize,
|
||||
)
|
||||
ast.parse(source, filename=_SYNTHETIC_SOURCE_FILENAME, mode="exec")
|
||||
table = symtable.symtable(source, _SYNTHETIC_SOURCE_FILENAME, "exec")
|
||||
parsed: tuple[CodeType, symtable.SymbolTable] | None = (module_code, table)
|
||||
except Exception:
|
||||
parsed = None
|
||||
if parsed is None:
|
||||
_packaging_reject()
|
||||
module_code, table = parsed
|
||||
|
||||
for child in table.get_children():
|
||||
if child.get_name() == callable_name and child.get_type() == "function":
|
||||
return module_code, table
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _source_bound_names(table: symtable.SymbolTable) -> set[str]:
|
||||
names: set[str] = set()
|
||||
for symbol in table.get_symbols():
|
||||
if symbol.is_imported() or symbol.is_assigned() or symbol.is_namespace():
|
||||
names.add(symbol.get_name())
|
||||
return names
|
||||
|
||||
|
||||
def _code_fingerprint(code: CodeType) -> tuple[object, ...]:
|
||||
"""Structural fingerprint ignoring only location/debug fields."""
|
||||
constants = tuple(
|
||||
_code_fingerprint(constant) if isinstance(constant, CodeType) else constant
|
||||
for constant in code.co_consts
|
||||
)
|
||||
return (
|
||||
code.co_name,
|
||||
getattr(code, "co_qualname", code.co_name),
|
||||
code.co_argcount,
|
||||
code.co_posonlyargcount,
|
||||
code.co_kwonlyargcount,
|
||||
code.co_flags,
|
||||
code.co_code,
|
||||
code.co_names,
|
||||
code.co_varnames,
|
||||
code.co_freevars,
|
||||
code.co_cellvars,
|
||||
getattr(code, "co_exceptiontable", b""),
|
||||
constants,
|
||||
)
|
||||
|
||||
|
||||
def _toplevel_code_candidates(
|
||||
module_code: CodeType, callable_name: str
|
||||
) -> list[CodeType]:
|
||||
candidates: list[CodeType] = []
|
||||
for constant in module_code.co_consts:
|
||||
if not isinstance(constant, CodeType):
|
||||
continue
|
||||
if constant.co_name != callable_name:
|
||||
continue
|
||||
if getattr(constant, "co_qualname", callable_name) != callable_name:
|
||||
continue
|
||||
candidates.append(constant)
|
||||
return candidates
|
||||
|
||||
|
||||
def _validate_loaded_code_matches_source(
|
||||
fn: FunctionType, module_code: CodeType
|
||||
) -> None:
|
||||
candidates = _toplevel_code_candidates(module_code, fn.__name__)
|
||||
if not candidates:
|
||||
_packaging_reject()
|
||||
target = _code_fingerprint(fn.__code__)
|
||||
if not any(_code_fingerprint(candidate) == target for candidate in candidates):
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _validate_signature(fn: FunctionType, config: _UdfConfig) -> None:
|
||||
try:
|
||||
signature: inspect.Signature | None = inspect.signature(fn)
|
||||
except (TypeError, ValueError):
|
||||
signature = None
|
||||
if signature is None:
|
||||
_packaging_reject()
|
||||
parameters = list(signature.parameters.values())
|
||||
expected = [name for name, _ in config.inputs]
|
||||
actual = [parameter.name for parameter in parameters]
|
||||
if actual != expected:
|
||||
_packaging_reject()
|
||||
for parameter in parameters:
|
||||
if parameter.kind not in _ALLOWED_PARAM_KINDS:
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _validate_ambient_globals(fn: FunctionType, table: symtable.SymbolTable) -> None:
|
||||
try:
|
||||
closure_vars: inspect.ClosureVars | None = inspect.getclosurevars(fn)
|
||||
except (TypeError, ValueError):
|
||||
closure_vars = None
|
||||
if closure_vars is None:
|
||||
_packaging_reject()
|
||||
if closure_vars.nonlocals:
|
||||
_packaging_reject()
|
||||
bound_names = _source_bound_names(table)
|
||||
for name in closure_vars.globals:
|
||||
if name not in bound_names:
|
||||
_packaging_reject()
|
||||
|
||||
|
||||
def _package_udf(fn: object) -> _PackagedUdf:
|
||||
"""Validate and snapshot a packagable ``@udf``-decorated function."""
|
||||
config = _get_udf_config(fn)
|
||||
if not isinstance(fn, FunctionType) or not _is_ordinary_function(fn):
|
||||
_packaging_reject()
|
||||
|
||||
module_name = fn.__module__
|
||||
if (
|
||||
not isinstance(module_name, str)
|
||||
or module_name == ""
|
||||
or module_name == "__main__"
|
||||
):
|
||||
_packaging_reject()
|
||||
module = sys.modules.get(module_name)
|
||||
if module is None:
|
||||
_packaging_reject()
|
||||
callable_name = fn.__name__
|
||||
if vars(module).get(callable_name) is not fn:
|
||||
_packaging_reject()
|
||||
|
||||
source_path = _resolve_source_path(fn, module)
|
||||
try:
|
||||
source: str | None = source_path.read_text(encoding="utf-8")
|
||||
except (OSError, UnicodeError):
|
||||
source = None
|
||||
if source is None:
|
||||
_packaging_reject()
|
||||
|
||||
module_code, table = _validate_source(source, callable_name)
|
||||
_validate_signature(fn, config)
|
||||
_validate_ambient_globals(fn, table)
|
||||
_validate_loaded_code_matches_source(fn, module_code)
|
||||
|
||||
return _PackagedUdf(
|
||||
source=source,
|
||||
module=module_name,
|
||||
callable_name=callable_name,
|
||||
config=config,
|
||||
)
|
||||
|
||||
|
||||
def _normalize_capability_triple(
|
||||
capability: FunctionCapability,
|
||||
) -> tuple[str, str, str | None]:
|
||||
"""Normalize a local capability declaration to the native triple shape."""
|
||||
# Private config is untrusted; re-check exact type before any property access.
|
||||
capability = _require_exact_capability(capability)
|
||||
if capability.kind == "network":
|
||||
origin = capability.origin
|
||||
if origin is None:
|
||||
raise ValueError("invalid network capability") from None
|
||||
return ("network", origin, None)
|
||||
if capability.kind == "secret":
|
||||
reference = capability.reference
|
||||
environment_variable = capability.environment_variable
|
||||
if reference is None or environment_variable is None:
|
||||
raise ValueError("invalid secret capability") from None
|
||||
return ("secret", reference, environment_variable)
|
||||
# Fail closed without echoing the unknown kind.
|
||||
raise ValueError("unsupported capability kind") from None
|
||||
|
||||
|
||||
def _build_function_definition(fn: object) -> _lancedb._FunctionDefinition:
|
||||
"""Package a ``@udf`` and bridge it to the private native definition."""
|
||||
packaged = _package_udf(fn)
|
||||
config = packaged.config
|
||||
capabilities = [
|
||||
_normalize_capability_triple(capability) for capability in config.capabilities
|
||||
]
|
||||
return _lancedb._new_function_definition(
|
||||
parameters=list(config.inputs),
|
||||
output_type=config.output,
|
||||
output_nullable=config.output_nullable,
|
||||
module=packaged.module,
|
||||
callable_name=packaged.callable_name,
|
||||
source=packaged.source,
|
||||
python=config.python,
|
||||
packages=list(config.packages),
|
||||
capabilities=capabilities,
|
||||
)
|
||||
@@ -63,8 +63,12 @@ if TYPE_CHECKING:
|
||||
import pyarrow as pa
|
||||
from .pydantic import LanceModel
|
||||
|
||||
from ._functions import _AsyncFunctions, _SyncFunctions
|
||||
from ._lancedb import Connection as LanceDbConnection
|
||||
from ._lancedb import Function
|
||||
from ._lancedb import Job as NativeJob
|
||||
from ._lancedb import JobDescription, JobInfo
|
||||
from ._lancedb import _FunctionDefinition
|
||||
from .common import DATA, URI
|
||||
from .embeddings import EmbeddingFunctionConfig
|
||||
from ._lancedb import Session
|
||||
@@ -650,6 +654,71 @@ class DBConnection(EnforceOverrides):
|
||||
"job_history is not supported for this connection type"
|
||||
)
|
||||
|
||||
@property
|
||||
def functions(self) -> "_SyncFunctions":
|
||||
"""First-class Function operations for this connection."""
|
||||
from ._functions import _SyncFunctions
|
||||
|
||||
return _SyncFunctions(self)
|
||||
|
||||
def _submit_register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
"""Submit a Function registration job via the native connection.
|
||||
|
||||
Connection subclasses that support registration override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function registration is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _submit_replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
"""Submit a Function conditional replace job via the native connection.
|
||||
|
||||
Connection subclasses that support registration override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function replace is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
"""Look up a Function by database-scoped name via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function lookup is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
"""Look up a Function by exact Function ID via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function lookup is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
"""Conditionally remove a Function catalog name via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function name removal is not supported for this connection type"
|
||||
)
|
||||
|
||||
def _revoke_function(self, function: "Function") -> None:
|
||||
"""Revoke an exact immutable Function via the native connection.
|
||||
|
||||
Connection subclasses that share the native Connection override this hook.
|
||||
"""
|
||||
raise NotImplementedError(
|
||||
"function revocation is not supported for this connection type"
|
||||
)
|
||||
|
||||
|
||||
class LanceDBConnection(DBConnection):
|
||||
"""
|
||||
@@ -1267,6 +1336,34 @@ class LanceDBConnection(DBConnection):
|
||||
"""
|
||||
return LOOP.run(self._conn.job_history(job_id))
|
||||
|
||||
@override
|
||||
def _submit_register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return LOOP.run(self._conn._register_function(name, definition))
|
||||
|
||||
@override
|
||||
def _submit_replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return LOOP.run(self._conn._replace_function(name, current, definition))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_name(name))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_id(function_id))
|
||||
|
||||
@override
|
||||
def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
return LOOP.run(self._conn._remove_function_name(name, current))
|
||||
|
||||
@override
|
||||
def _revoke_function(self, function: "Function") -> None:
|
||||
return LOOP.run(self._conn._revoke_function(function))
|
||||
|
||||
@override
|
||||
def namespace_client(self) -> LanceNamespace:
|
||||
"""Get the equivalent namespace client for this connection.
|
||||
@@ -2013,6 +2110,35 @@ class AsyncConnection(object):
|
||||
"""
|
||||
return await self._inner.job_history(job_id)
|
||||
|
||||
@property
|
||||
def functions(self) -> "_AsyncFunctions":
|
||||
"""First-class Function operations for this connection."""
|
||||
from ._functions import _AsyncFunctions
|
||||
|
||||
return _AsyncFunctions(self)
|
||||
|
||||
async def _register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return await self._inner._register_function(name, definition)
|
||||
|
||||
async def _replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> NativeJob:
|
||||
return await self._inner._replace_function(name, current, definition)
|
||||
|
||||
async def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
return await self._inner._lookup_function_by_name(name)
|
||||
|
||||
async def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
return await self._inner._lookup_function_by_id(function_id)
|
||||
|
||||
async def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
return await self._inner._remove_function_name(name, current)
|
||||
|
||||
async def _revoke_function(self, function: "Function") -> None:
|
||||
return await self._inner._revoke_function(function)
|
||||
|
||||
async def namespace_client(self) -> LanceNamespace:
|
||||
"""Get the equivalent namespace client for this connection.
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
|
||||
"""Custom exception handling"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class MissingValueError(ValueError):
|
||||
"""Exception raised when a required value is missing."""
|
||||
@@ -26,12 +28,47 @@ class MissingColumnError(KeyError):
|
||||
|
||||
|
||||
class JobFailedError(RuntimeError):
|
||||
"""Exception raised when an asynchronous job reaches the failed state."""
|
||||
"""Exception raised when an asynchronous job reaches the failed state.
|
||||
|
||||
pass
|
||||
``error_code`` is the optional exact category string projected from the
|
||||
native job failure when the backend supplied one. The RuntimeError
|
||||
message remains the existing diagnostic text and must not be used to
|
||||
recover or override the code.
|
||||
"""
|
||||
|
||||
__slots__ = ("_error_code",)
|
||||
|
||||
def __init__(self, message: str, error_code: Optional[str] = None) -> None:
|
||||
super().__init__(message)
|
||||
self._error_code = error_code
|
||||
|
||||
@property
|
||||
def error_code(self) -> Optional[str]:
|
||||
"""Exact job failure error category string, when supplied."""
|
||||
return self._error_code
|
||||
|
||||
|
||||
class JobCancelledError(RuntimeError):
|
||||
"""Exception raised when an asynchronous job was cancelled."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class FunctionError(RuntimeError):
|
||||
"""Exception raised when a first-class Function operation fails.
|
||||
|
||||
``code`` is the stable semantic category from the native error. The
|
||||
message is a sanitized client diagnostic and must not be used to recover
|
||||
or override the code.
|
||||
"""
|
||||
|
||||
__slots__ = ("_code",)
|
||||
|
||||
def __init__(self, message: str, code: str) -> None:
|
||||
super().__init__(message)
|
||||
self._code = code
|
||||
|
||||
@property
|
||||
def code(self) -> str:
|
||||
"""Stable Function error category string."""
|
||||
return self._code
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import Optional
|
||||
from lancedb.background_loop import LOOP
|
||||
|
||||
from . import _lancedb
|
||||
from ._lancedb import Function
|
||||
|
||||
|
||||
class AsyncJob:
|
||||
@@ -44,18 +45,22 @@ class AsyncJob:
|
||||
return "finished"
|
||||
return await self._inner.status()
|
||||
|
||||
async def wait(self, timeout: Optional[timedelta] = None):
|
||||
async def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
"""Wait until the operation reaches a terminal state.
|
||||
|
||||
Returns the success result when present (currently a
|
||||
:class:`~lancedb.Function`), or `None` when the job finished without
|
||||
one.
|
||||
|
||||
Raises `JobFailedError` if the operation failed, `JobCancelledError`
|
||||
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return
|
||||
return None
|
||||
if timeout is None:
|
||||
await self._inner.wait()
|
||||
return await self._inner.wait()
|
||||
else:
|
||||
await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
return await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
|
||||
|
||||
async def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
@@ -88,15 +93,19 @@ class Job:
|
||||
return "finished"
|
||||
return LOOP.run(self._inner.status())
|
||||
|
||||
def wait(self, timeout: Optional[timedelta] = None):
|
||||
def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
|
||||
"""Block until the operation reaches a terminal state.
|
||||
|
||||
Returns the success result when present (currently a
|
||||
:class:`~lancedb.Function`), or `None` when the job finished without
|
||||
one.
|
||||
|
||||
Raises `JobFailedError` if the operation failed, `JobCancelledError`
|
||||
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
|
||||
"""
|
||||
if self._inner is None:
|
||||
return
|
||||
LOOP.run(self._inner.wait(timeout))
|
||||
return None
|
||||
return LOOP.run(self._inner.wait(timeout))
|
||||
|
||||
def cancel(self):
|
||||
"""Request cancellation. Cancelling a finished operation is a no-op."""
|
||||
|
||||
@@ -26,7 +26,9 @@ from ..db import DBConnection, LOOP
|
||||
from ..job import Job
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .._lancedb import JobDescription, JobInfo
|
||||
from .._lancedb import Function
|
||||
from .._lancedb import Job as NativeJob
|
||||
from .._lancedb import JobDescription, JobInfo, _FunctionDefinition
|
||||
from ..embeddings import EmbeddingFunctionConfig
|
||||
from lance_namespace import (
|
||||
LanceNamespace,
|
||||
@@ -734,6 +736,34 @@ class RemoteDBConnection(DBConnection):
|
||||
"""
|
||||
return LOOP.run(self._conn.job_history(job_id))
|
||||
|
||||
@override
|
||||
def _submit_register_function(
|
||||
self, name: str, definition: "_FunctionDefinition"
|
||||
) -> "NativeJob":
|
||||
return LOOP.run(self._conn._register_function(name, definition))
|
||||
|
||||
@override
|
||||
def _submit_replace_function(
|
||||
self, name: str, current: "Function", definition: "_FunctionDefinition"
|
||||
) -> "NativeJob":
|
||||
return LOOP.run(self._conn._replace_function(name, current, definition))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_name(self, name: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_name(name))
|
||||
|
||||
@override
|
||||
def _lookup_function_by_id(self, function_id: str) -> "Function":
|
||||
return LOOP.run(self._conn._lookup_function_by_id(function_id))
|
||||
|
||||
@override
|
||||
def _remove_function_name(self, name: str, current: "Function") -> None:
|
||||
return LOOP.run(self._conn._remove_function_name(name, current))
|
||||
|
||||
@override
|
||||
def _revoke_function(self, function: "Function") -> None:
|
||||
return LOOP.run(self._conn._revoke_function(function))
|
||||
|
||||
@override
|
||||
def namespace_client(self) -> LanceNamespace:
|
||||
"""Get the equivalent namespace client for this connection.
|
||||
|
||||
@@ -7,6 +7,7 @@ import logging
|
||||
from functools import cached_property
|
||||
import os
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
@@ -67,6 +68,9 @@ from ..query import (
|
||||
from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Table, Tags
|
||||
from ..types import BaseTokenizerType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lancedb._lancedb import _FunctionCall
|
||||
|
||||
|
||||
class RemoteTable(Table):
|
||||
def __init__(
|
||||
@@ -570,6 +574,45 @@ class RemoteTable(Table):
|
||||
)
|
||||
)
|
||||
|
||||
def add_generated_column(self, column_name: str, call: "_FunctionCall") -> Job:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
|
||||
of the Job does not publish the column; callers must wait and re-read
|
||||
the table to observe the new definition and values.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
|
||||
|
||||
def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
return LOOP.run(self._table.generated_column_status(column_name))
|
||||
|
||||
def refresh_generated_column(self, column_name: str) -> Job:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
|
||||
of the Job does not publish new values; callers must wait and re-read
|
||||
the table to observe refreshed results.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
|
||||
|
||||
def alter_generated_column(
|
||||
self, column_name: str, new_call: "_FunctionCall"
|
||||
) -> Job:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
|
||||
of the Job does not publish the new definition; callers must wait and
|
||||
re-read the table to observe the updated column.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
|
||||
|
||||
def _is_legacy_create_index_call(
|
||||
self,
|
||||
first_arg: str,
|
||||
|
||||
@@ -11,6 +11,11 @@ Provides StreamingDataset, a PyTorch IterableDataset that guarantees:
|
||||
- **Resumability**: state_dict / load_state_dict capture per-split consumption
|
||||
counts so training can resume from an exact mid-epoch position even when the
|
||||
distributed topology changes between runs.
|
||||
|
||||
Transform failures on bad rows (e.g. nulls or NaNs from incomplete data) can
|
||||
be tolerated with ``on_transform_error="skip"``; see the parameter
|
||||
documentation on StreamingDataset for how this interacts with the guarantees
|
||||
above.
|
||||
"""
|
||||
|
||||
import ctypes
|
||||
@@ -22,7 +27,7 @@ import time
|
||||
from collections import deque
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from multiprocessing import RawArray
|
||||
from typing import Any, Callable, Iterator, Optional
|
||||
from typing import Any, Callable, Iterator, Optional, Union
|
||||
|
||||
from torch.utils.data import IterableDataset, get_worker_info
|
||||
|
||||
@@ -127,6 +132,49 @@ class StreamingDataset(IterableDataset):
|
||||
Maximum number of transforms to run concurrently. Must be greater
|
||||
than zero. When ``None`` (the default), uses ``os.cpu_count()`` or 1
|
||||
when the CPU count is unavailable.
|
||||
on_transform_error:
|
||||
What to do when the transform raises an exception:
|
||||
|
||||
- ``"raise"`` (the default): the exception propagates and iteration
|
||||
aborts.
|
||||
- ``"skip"``: the failing rows are dropped and iteration continues.
|
||||
- ``"warn"``: like ``"skip"``, but a warning is logged for each
|
||||
failing batch.
|
||||
- a callable ``handler(exc) -> bool``: called with the exception;
|
||||
return ``True`` to skip the failing rows or ``False`` to re-raise.
|
||||
Useful to skip only expected error types (compatible with
|
||||
``webdataset.handlers`` style handlers).
|
||||
|
||||
When a batch fails, the transform is re-invoked on each single-row
|
||||
slice of the batch so that only the rows that actually fail are
|
||||
dropped. Transforms should therefore be deterministic and accept
|
||||
batches of any size (including one row). Skipped rows are counted in
|
||||
``rows_skipped``.
|
||||
|
||||
Skipping weakens the elastic-determinism guarantee at the end of the
|
||||
epoch: splits that lose more rows than others run dry earlier, and
|
||||
each rank's iterator ends at the last cycle where every split *it
|
||||
owns* still has a row. Because bad rows are not distributed evenly
|
||||
across splits, this means one rank's iterator can yield noticeably
|
||||
fewer or more steps than another rank's *in the same run* — there is
|
||||
no cross-rank coordination that stops every rank at the same global
|
||||
step. This is generally safe for asynchronous or single-rank use,
|
||||
but synchronous distributed training (e.g. ranks that call
|
||||
``all_reduce`` every step) can hang or deadlock if one rank's
|
||||
iterator is exhausted while others are still stepping; callers doing
|
||||
synchronous multi-rank training with ``on_transform_error != "raise"``
|
||||
are responsible for their own cross-rank stopping mechanism (e.g.
|
||||
broadcasting a stop signal on ``StopIteration``). The final few
|
||||
global steps can also differ across topologies (bounded by the skew
|
||||
in bad-row counts across splits). The sequence of samples yielded
|
||||
from each split remains deterministic. Mid-epoch
|
||||
checkpoints remain exact provided the transform fails
|
||||
deterministically; in multi-rank training each rank must save its
|
||||
own ``state_dict`` and the states must be combined with
|
||||
``merge_state_dicts`` before resuming on a different topology.
|
||||
Prefer the ``filter`` parameter when bad rows can be expressed as a
|
||||
SQL predicate (e.g. ``"col IS NOT NULL"``) — filtering happens before
|
||||
splits are built, so every guarantee is fully preserved.
|
||||
worker_info_override:
|
||||
If set, used in place of ``torch.utils.data.get_worker_info()`` to
|
||||
determine the DataLoader worker assignment. Intended for unit tests
|
||||
@@ -152,6 +200,7 @@ class StreamingDataset(IterableDataset):
|
||||
filter: Optional[str] = None,
|
||||
transform: Optional[Callable] = None,
|
||||
transform_parallelism: Optional[int] = None,
|
||||
on_transform_error: Union[str, Callable[[Exception], bool]] = "raise",
|
||||
connection_factory: Optional[Callable[[str], Any]] = None,
|
||||
worker_info_override=None,
|
||||
):
|
||||
@@ -167,6 +216,13 @@ class StreamingDataset(IterableDataset):
|
||||
)
|
||||
if transform_parallelism is not None and transform_parallelism <= 0:
|
||||
raise ValueError("transform_parallelism must be greater than 0")
|
||||
if on_transform_error not in ("raise", "skip", "warn") and not callable(
|
||||
on_transform_error
|
||||
):
|
||||
raise ValueError(
|
||||
"on_transform_error must be 'raise', 'skip', 'warn', or a "
|
||||
f"callable, got {on_transform_error!r}"
|
||||
)
|
||||
|
||||
self._table = table
|
||||
self._num_splits = num_splits
|
||||
@@ -182,6 +238,7 @@ class StreamingDataset(IterableDataset):
|
||||
self._filter = filter
|
||||
self._transform = transform
|
||||
self._transform_parallelism = transform_parallelism
|
||||
self._on_transform_error = on_transform_error
|
||||
self._connection_factory = connection_factory
|
||||
self._worker_info_override = worker_info_override
|
||||
|
||||
@@ -199,19 +256,28 @@ class StreamingDataset(IterableDataset):
|
||||
# in the main process. RawArray is picklable via the forkserver
|
||||
# reduction protocol so it survives the dataset pickle round-trip.
|
||||
# Layout: [unscanned_rows, raw_rows, cooked_rows, consumed_rows,
|
||||
# bytes_loaded, fetch_time_us, transform_time_us]
|
||||
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 7)
|
||||
# bytes_loaded, fetch_time_us, transform_time_us,
|
||||
# rows_skipped]
|
||||
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 8)
|
||||
|
||||
# Cumulative bytes of Arrow buffer data fetched across all iterations.
|
||||
self._bytes_loaded: int = 0
|
||||
# Cumulative seconds spent in LanceDB I/O and in transform functions.
|
||||
self._fetch_time: float = 0.0
|
||||
self._transform_time: float = 0.0
|
||||
# Cumulative rows dropped by on_transform_error across all iterations.
|
||||
self._rows_skipped: int = 0
|
||||
|
||||
# Number of samples each split has already been consumed. At global
|
||||
# step boundaries all splits have consumed this many samples, so a
|
||||
# single scalar captures the topology-independent checkpoint state.
|
||||
self._resume_offset: int = 0
|
||||
# Permutation position each split has consumed through, keyed by
|
||||
# global split index. Equal to _resume_offset for every split unless
|
||||
# on_transform_error skipped rows, in which case skipped positions
|
||||
# push the watermark of the affected splits further ahead. Splits
|
||||
# this instance has never iterated have no entry.
|
||||
self._resume_positions: dict[int, int] = {}
|
||||
|
||||
# Build the permutation table once, deterministically.
|
||||
builder = permutation_builder(table)
|
||||
@@ -275,6 +341,7 @@ class StreamingDataset(IterableDataset):
|
||||
# Set identity transform on each Permutation so __getitems__ returns
|
||||
# the raw RecordBatch. Stage 2 applies the real transform.
|
||||
permutations: list[Permutation] = []
|
||||
initial_positions: list[int] = []
|
||||
for split_idx in my_splits:
|
||||
perm = Permutation.from_tables(
|
||||
self._table, self._perm_table, split=split_idx
|
||||
@@ -282,14 +349,20 @@ class StreamingDataset(IterableDataset):
|
||||
if self._columns is not None:
|
||||
perm = perm.select_columns(self._columns)
|
||||
perm = perm.with_transform(lambda batch: batch)
|
||||
if self._resume_offset > 0:
|
||||
perm = perm.with_skip(self._resume_offset)
|
||||
start_pos = self._resume_positions.get(split_idx, self._resume_offset)
|
||||
if start_pos > 0:
|
||||
perm = perm.with_skip(start_pos)
|
||||
initial_positions.append(start_pos)
|
||||
permutations.append(perm)
|
||||
|
||||
n = len(permutations)
|
||||
split_sizes = [perm.num_rows for perm in permutations]
|
||||
initial_offset = self._resume_offset
|
||||
local_consumed = [0] * n
|
||||
# Permutation position each split has consumed through (absolute,
|
||||
# i.e. counted from the start of the unskipped split). Runs ahead of
|
||||
# initial + local_consumed when rows are skipped.
|
||||
pos_consumed = list(initial_positions)
|
||||
|
||||
batch_size = self._read_batch_size
|
||||
max_prefetch = self._prefetch_batches
|
||||
@@ -302,12 +375,14 @@ class StreamingDataset(IterableDataset):
|
||||
self._transform if self._transform is not None else Transforms.arrow2python
|
||||
)
|
||||
|
||||
# Per-split pipeline state.
|
||||
# Per-split pipeline state. Batches are paired with the absolute
|
||||
# permutation position of their first row so that skipped rows can be
|
||||
# accounted for in pos_consumed.
|
||||
fetch_head = [0] * n
|
||||
io_pending = [deque() for _ in range(n)] # Future[RecordBatch]
|
||||
raw_batches = [deque() for _ in range(n)] # RecordBatch — fetched, awaiting tx
|
||||
tx_pending = [deque() for _ in range(n)] # Future[list[Any]]
|
||||
cooked = [deque() for _ in range(n)] # rows ready to yield
|
||||
io_pending = [deque() for _ in range(n)] # (abs_start, Future[RecordBatch])
|
||||
raw_batches = [deque() for _ in range(n)] # (abs_start, RecordBatch)
|
||||
tx_pending = [deque() for _ in range(n)] # Future[list[(abs_pos, row)]]
|
||||
cooked = [deque() for _ in range(n)] # (abs_pos, row) ready to yield
|
||||
|
||||
# Limit simultaneous transforms to transform_workers across all splits.
|
||||
tx_semaphore = threading.Semaphore(transform_workers)
|
||||
@@ -330,7 +405,8 @@ class StreamingDataset(IterableDataset):
|
||||
fetch_head[i] += fetch
|
||||
perm_i = permutations[i]
|
||||
indices = list(range(start, start + fetch))
|
||||
io_pending[i].append(io_pool.submit(_io_call, perm_i, indices))
|
||||
abs_start = initial_positions[i] + start
|
||||
io_pending[i].append((abs_start, io_pool.submit(_io_call, perm_i, indices)))
|
||||
|
||||
def _fill_io(i: int) -> None:
|
||||
while len(io_pending[i]) < max_prefetch and fetch_head[i] < split_sizes[i]:
|
||||
@@ -338,15 +414,72 @@ class StreamingDataset(IterableDataset):
|
||||
|
||||
def _drain_io(i: int) -> None:
|
||||
"""Move completed I/O futures into raw_batches non-blockingly."""
|
||||
while io_pending[i] and io_pending[i][0].done():
|
||||
raw_batches[i].append(io_pending[i].popleft().result())
|
||||
while io_pending[i] and io_pending[i][0][1].done():
|
||||
abs_start, fut = io_pending[i].popleft()
|
||||
raw_batches[i].append((abs_start, fut.result()))
|
||||
|
||||
# ── Stage 2 helpers ───────────────────────────────────────────────────
|
||||
|
||||
def _tx_call_guarded(batch):
|
||||
on_error = self._on_transform_error
|
||||
|
||||
def _should_skip(exc: Exception) -> bool:
|
||||
if on_error == "raise":
|
||||
return False
|
||||
if callable(on_error):
|
||||
return bool(on_error(exc))
|
||||
return True # "skip" or "warn"
|
||||
|
||||
def _check_row_count(rows: list, num_rows: int) -> None:
|
||||
if len(rows) != num_rows:
|
||||
raise ValueError(
|
||||
f"transform returned {len(rows)} rows for a batch of "
|
||||
f"{num_rows}; transforms must return exactly one output "
|
||||
"row per input row. To drop bad rows, raise inside the "
|
||||
"transform and pass on_transform_error='skip'."
|
||||
)
|
||||
|
||||
def _transform_isolated(abs_start, batch, batch_exc):
|
||||
"""Re-run the transform on single-row slices, dropping failures."""
|
||||
out = []
|
||||
skipped = 0
|
||||
first_exc = None
|
||||
for j in range(batch.num_rows):
|
||||
try:
|
||||
rows = list(final_transform(batch.slice(j, 1)))
|
||||
except Exception as exc:
|
||||
if not _should_skip(exc):
|
||||
raise
|
||||
skipped += 1
|
||||
if first_exc is None:
|
||||
first_exc = exc
|
||||
continue
|
||||
_check_row_count(rows, 1)
|
||||
out.append((abs_start + j, rows[0]))
|
||||
self._rows_skipped += skipped
|
||||
if skipped and on_error == "warn":
|
||||
logger.warning(
|
||||
"Skipped %d of %d rows whose transform failed (first error: %r)",
|
||||
skipped,
|
||||
batch.num_rows,
|
||||
first_exc if first_exc is not None else batch_exc,
|
||||
)
|
||||
return out
|
||||
|
||||
def _transform_batch(abs_start, batch):
|
||||
"""Apply the transform, returning [(abs_pos, row), ...]."""
|
||||
try:
|
||||
rows = list(final_transform(batch))
|
||||
except Exception as exc:
|
||||
if not _should_skip(exc):
|
||||
raise
|
||||
return _transform_isolated(abs_start, batch, exc)
|
||||
_check_row_count(rows, batch.num_rows)
|
||||
return [(abs_start + j, row) for j, row in enumerate(rows)]
|
||||
|
||||
def _tx_call_guarded(abs_start, batch):
|
||||
try:
|
||||
t0 = time.perf_counter()
|
||||
result = final_transform(batch)
|
||||
result = _transform_batch(abs_start, batch)
|
||||
self._transform_time += time.perf_counter() - t0
|
||||
return result
|
||||
finally:
|
||||
@@ -355,8 +488,8 @@ class StreamingDataset(IterableDataset):
|
||||
def _try_submit_tx(i: int) -> None:
|
||||
"""Submit transforms for raw_batches[i] up to available capacity."""
|
||||
while raw_batches[i] and tx_semaphore.acquire(blocking=False):
|
||||
batch = raw_batches[i].popleft()
|
||||
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch))
|
||||
abs_start, batch = raw_batches[i].popleft()
|
||||
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, abs_start, batch))
|
||||
|
||||
def _drain_tx(i: int) -> None:
|
||||
"""Move completed transform futures into cooked non-blockingly."""
|
||||
@@ -384,11 +517,14 @@ class StreamingDataset(IterableDataset):
|
||||
# Acquire a transform slot (may block briefly if all
|
||||
# transform_workers are busy with other splits).
|
||||
tx_semaphore.acquire()
|
||||
batch = raw_batches[i].popleft()
|
||||
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch))
|
||||
abs_start, batch = raw_batches[i].popleft()
|
||||
tx_pending[i].append(
|
||||
tx_pool.submit(_tx_call_guarded, abs_start, batch)
|
||||
)
|
||||
elif io_pending[i]:
|
||||
# Block on the oldest in-flight I/O fetch.
|
||||
raw_batches[i].append(io_pending[i].popleft().result())
|
||||
abs_start, fut = io_pending[i].popleft()
|
||||
raw_batches[i].append((abs_start, fut.result()))
|
||||
_advance(i)
|
||||
else:
|
||||
break # split exhausted
|
||||
@@ -407,15 +543,28 @@ class StreamingDataset(IterableDataset):
|
||||
_fill_io(i)
|
||||
|
||||
while True:
|
||||
# Stop when any split is exhausted (all exhaust
|
||||
# simultaneously: equal split sizes + round-robin).
|
||||
if any(local_consumed[i] >= split_sizes[i] for i in range(n)):
|
||||
# A cycle only runs if every split can still produce a
|
||||
# row. Without skips all splits exhaust simultaneously
|
||||
# (equal split sizes + round-robin); when
|
||||
# on_transform_error drops rows a split can run dry
|
||||
# early, ending the epoch at the last complete cycle.
|
||||
# This check only sees splits owned by this rank/worker
|
||||
# (my_splits) — there is no cross-rank coordination, so
|
||||
# a different rank with fewer skipped rows keeps going;
|
||||
# see the on_transform_error docstring.
|
||||
exhausted = False
|
||||
for i in range(n):
|
||||
_ensure_cooked(i)
|
||||
if not cooked[i]:
|
||||
exhausted = True
|
||||
break
|
||||
if exhausted:
|
||||
break
|
||||
|
||||
for i in range(n):
|
||||
_ensure_cooked(i)
|
||||
row = cooked[i].popleft()
|
||||
pos, row = cooked[i].popleft()
|
||||
local_consumed[i] += 1
|
||||
pos_consumed[i] = pos + 1
|
||||
_advance(i)
|
||||
|
||||
# After the last split in each cycle: update the
|
||||
@@ -424,21 +573,39 @@ class StreamingDataset(IterableDataset):
|
||||
# even when __iter__ runs in a worker process.
|
||||
if i == n - 1:
|
||||
self._resume_offset = initial_offset + local_consumed[i]
|
||||
for j, split_idx in enumerate(my_splits):
|
||||
self._resume_positions[split_idx] = pos_consumed[j]
|
||||
ws = self._worker_stats
|
||||
ws[0] = sum(
|
||||
split_sizes[j] - fetch_head[j] for j in range(n)
|
||||
)
|
||||
ws[1] = sum(
|
||||
batch.num_rows for q in raw_batches for batch in q
|
||||
batch.num_rows
|
||||
for q in raw_batches
|
||||
for _, batch in q
|
||||
)
|
||||
ws[2] = sum(len(q) for q in cooked)
|
||||
ws[3] = sum(local_consumed)
|
||||
ws[4] = self._bytes_loaded
|
||||
ws[5] = int(self._fetch_time * 1_000_000)
|
||||
ws[6] = int(self._transform_time * 1_000_000)
|
||||
ws[7] = self._rows_skipped
|
||||
|
||||
yield row
|
||||
finally:
|
||||
# Final stats flush: the per-cycle write above never runs
|
||||
# when iteration ends mid-cycle (e.g. a split whose rows
|
||||
# were all skipped before completing a single cycle), so
|
||||
# counters like rows_skipped would otherwise be stale.
|
||||
ws = self._worker_stats
|
||||
ws[0] = sum(split_sizes[j] - fetch_head[j] for j in range(n))
|
||||
ws[1] = 0 # queue-depth properties document 0 when idle
|
||||
ws[2] = 0
|
||||
ws[3] = sum(local_consumed)
|
||||
ws[4] = self._bytes_loaded
|
||||
ws[5] = int(self._fetch_time * 1_000_000)
|
||||
ws[6] = int(self._transform_time * 1_000_000)
|
||||
ws[7] = self._rows_skipped
|
||||
self._raw_batches_ref = None
|
||||
self._cooked_ref = None
|
||||
self._fetch_head_ref = None
|
||||
@@ -492,7 +659,7 @@ class StreamingDataset(IterableDataset):
|
||||
batches. Returns 0 when not iterating.
|
||||
"""
|
||||
if self._raw_batches_ref is not None:
|
||||
return sum(batch.num_rows for q in self._raw_batches_ref for batch in q)
|
||||
return sum(batch.num_rows for q in self._raw_batches_ref for _, batch in q)
|
||||
return int(self._worker_stats[1])
|
||||
|
||||
@property
|
||||
@@ -522,6 +689,19 @@ class StreamingDataset(IterableDataset):
|
||||
)
|
||||
return int(self._worker_stats[0])
|
||||
|
||||
@property
|
||||
def rows_skipped(self) -> int:
|
||||
"""Number of rows dropped because their transform raised an exception.
|
||||
|
||||
Only ever non-zero when ``on_transform_error`` is set to ``"skip"``,
|
||||
``"warn"``, or a callable that returned ``True``. Accumulates across
|
||||
multiple iterations of the same dataset instance and is never reset
|
||||
automatically.
|
||||
"""
|
||||
if self._raw_batches_ref is not None:
|
||||
return self._rows_skipped
|
||||
return int(self._worker_stats[7])
|
||||
|
||||
@property
|
||||
def consumed_rows(self) -> int:
|
||||
"""Number of rows already yielded to the caller across all splits.
|
||||
@@ -587,12 +767,27 @@ class StreamingDataset(IterableDataset):
|
||||
every split has been consumed the same number of times (by the
|
||||
round-robin design), so the per-split count is a single uniform value
|
||||
that is identical across all ranks and DataLoader workers.
|
||||
|
||||
``positions_consumed_per_split`` records how far into each split's
|
||||
permutation iteration has advanced. It only differs from
|
||||
``samples_consumed_per_split`` when ``on_transform_error`` skipped
|
||||
rows, in which case entries are exact for the splits this instance
|
||||
iterated and a lower bound (the sample count) for splits owned by
|
||||
other ranks or workers. Combine the state dicts from all ranks with
|
||||
[merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts]
|
||||
to recover the exact value for every split before resuming on a
|
||||
different topology.
|
||||
"""
|
||||
positions = [
|
||||
self._resume_positions.get(split, self._resume_offset)
|
||||
for split in range(self._num_splits)
|
||||
]
|
||||
return {
|
||||
"shuffle_seed": self._shuffle_seed,
|
||||
"num_splits": self._num_splits,
|
||||
"epoch": self._epoch,
|
||||
"samples_consumed_per_split": [self._resume_offset] * self._num_splits,
|
||||
"positions_consumed_per_split": positions,
|
||||
}
|
||||
|
||||
def load_state_dict(self, state: dict) -> None:
|
||||
@@ -618,3 +813,96 @@ class StreamingDataset(IterableDataset):
|
||||
self._resume_offset = consumed[0] if consumed else 0
|
||||
else:
|
||||
self._resume_offset = int(consumed)
|
||||
# Older checkpoints predate positions_consumed_per_split; without
|
||||
# skipped rows positions equal sample counts, so falling back to
|
||||
# _resume_offset (the .get default in __iter__) is exact.
|
||||
positions = state.get("positions_consumed_per_split")
|
||||
if positions is None:
|
||||
self._resume_positions = {}
|
||||
else:
|
||||
self._resume_positions = {
|
||||
split: int(pos) for split, pos in enumerate(positions)
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def merge_state_dicts(states: list[dict]) -> dict:
|
||||
"""Merge state dicts saved by different ranks into one exact state.
|
||||
|
||||
Only needed when ``on_transform_error`` skips rows in multi-rank
|
||||
training: each rank then knows the exact permutation position only for
|
||||
its own splits, and records a lower bound for the rest. Because
|
||||
exactly one rank owns each split, the elementwise maximum across all
|
||||
ranks' ``positions_consumed_per_split`` recovers the exact position of
|
||||
every split. Without skipped rows every rank's state is already
|
||||
identical and merging is a no-op.
|
||||
|
||||
Raises ``ValueError`` if the states are empty or were not produced by
|
||||
the same run (mismatched seed, split count, epoch, or sample counts).
|
||||
|
||||
The merge is always all-to-all and topology-agnostic: collect the
|
||||
``state_dict()`` from every rank of the *previous* run into one list,
|
||||
merge that whole list, and hand the identical merged result to every
|
||||
rank of the *next* run — regardless of whether the rank count grew,
|
||||
shrank, or stayed the same. There is no pairwise or subset merging
|
||||
step, because each split's exact position is only known to whichever
|
||||
rank owned that split, and the elementwise maximum needs every rank's
|
||||
contribution to be correct.
|
||||
|
||||
For example, checkpointing 8 ranks and resuming on 4 (the same
|
||||
pattern applies when growing, e.g. 4 ranks resuming on 8)::
|
||||
|
||||
states = [ds.state_dict() for ds in previous_run_datasets] # 8
|
||||
merged = StreamingDataset.merge_state_dicts(states)
|
||||
for ds in resumed_datasets: # now only 4 ranks
|
||||
ds.load_state_dict(merged) # same dict on every rank
|
||||
|
||||
The rank count on either side never affects the merge itself, since
|
||||
``merge_state_dicts`` only cares about the list of states it is
|
||||
given. Each split's position is recovered by elementwise maximum;
|
||||
here rank 0 owned split 0 (and skipped two rows there) while rank 1
|
||||
owned split 1 (and skipped one row):
|
||||
|
||||
>>> rank0 = {
|
||||
... "shuffle_seed": 0, "num_splits": 2, "epoch": 0,
|
||||
... "samples_consumed_per_split": [3, 3],
|
||||
... "positions_consumed_per_split": [5, 3],
|
||||
... }
|
||||
>>> rank1 = {
|
||||
... "shuffle_seed": 0, "num_splits": 2, "epoch": 0,
|
||||
... "samples_consumed_per_split": [3, 3],
|
||||
... "positions_consumed_per_split": [3, 4],
|
||||
... }
|
||||
>>> merged = StreamingDataset.merge_state_dicts([rank0, rank1])
|
||||
>>> merged["positions_consumed_per_split"]
|
||||
[5, 4]
|
||||
"""
|
||||
if not states:
|
||||
raise ValueError("merge_state_dicts requires at least one state dict")
|
||||
first = states[0]
|
||||
for state in states[1:]:
|
||||
for key in ("shuffle_seed", "num_splits", "epoch"):
|
||||
if state[key] != first[key]:
|
||||
raise ValueError(
|
||||
f"{key} mismatch across state dicts: "
|
||||
f"{state[key]} != {first[key]}"
|
||||
)
|
||||
if (
|
||||
state["samples_consumed_per_split"]
|
||||
!= first["samples_consumed_per_split"]
|
||||
):
|
||||
raise ValueError(
|
||||
"samples_consumed_per_split mismatch across state dicts; "
|
||||
"state_dict() must be called at the same global step "
|
||||
"boundary on every rank"
|
||||
)
|
||||
merged = dict(first)
|
||||
all_positions = [
|
||||
state.get(
|
||||
"positions_consumed_per_split", state["samples_consumed_per_split"]
|
||||
)
|
||||
for state in states
|
||||
]
|
||||
merged["positions_consumed_per_split"] = [
|
||||
max(per_split) for per_split in zip(*all_positions)
|
||||
]
|
||||
return merged
|
||||
|
||||
@@ -185,6 +185,7 @@ if TYPE_CHECKING:
|
||||
LsmWriteSpec,
|
||||
MergeResult,
|
||||
UpdateResult,
|
||||
_FunctionCall,
|
||||
)
|
||||
from .index import IndexConfig
|
||||
import pandas
|
||||
@@ -1007,6 +1008,43 @@ class Table(ABC):
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
|
||||
of the Job does not publish the column; callers must wait and re-read
|
||||
the table to observe the new definition and values.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def refresh_generated_column(self, column_name: str) -> Job:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
|
||||
of the Job does not publish new values; callers must wait and re-read
|
||||
the table to observe refreshed results.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
|
||||
of the Job does not publish the new definition; callers must wait and
|
||||
re-read the table to observe the updated column.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def drop_index(self, name: str) -> None:
|
||||
"""
|
||||
Drop an index from the table.
|
||||
@@ -2849,6 +2887,43 @@ class LanceTable(Table):
|
||||
)
|
||||
)
|
||||
|
||||
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
|
||||
of the Job does not publish the column; callers must wait and re-read
|
||||
the table to observe the new definition and values.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
|
||||
|
||||
def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
return LOOP.run(self._table.generated_column_status(column_name))
|
||||
|
||||
def refresh_generated_column(self, column_name: str) -> Job:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
|
||||
of the Job does not publish new values; callers must wait and re-read
|
||||
the table to observe refreshed results.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
|
||||
|
||||
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
|
||||
of the Job does not publish the new definition; callers must wait and
|
||||
re-read the table to observe the updated column.
|
||||
"""
|
||||
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
|
||||
|
||||
def _is_legacy_create_index_call(
|
||||
self,
|
||||
first_arg: str,
|
||||
@@ -5066,6 +5141,50 @@ class AsyncTable:
|
||||
)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def add_generated_column(
|
||||
self, column_name: str, call: _FunctionCall
|
||||
) -> AsyncJob:
|
||||
"""Add a generated column from an authored Function call.
|
||||
|
||||
Returns an :class:`~lancedb.job.AsyncJob` for the create operation.
|
||||
Acceptance of the Job does not publish the column; callers must wait
|
||||
and re-read the table to observe the new definition and values.
|
||||
"""
|
||||
job = await self._inner._add_generated_column(column_name, call)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def generated_column_status(
|
||||
self, column_name: str
|
||||
) -> Literal["complete", "incomplete"]:
|
||||
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
|
||||
|
||||
Projection-only: reads the named column's stored definition status.
|
||||
Does not refresh values, submit a Job, or mutate table state.
|
||||
"""
|
||||
return await self._inner._generated_column_status(column_name)
|
||||
|
||||
async def refresh_generated_column(self, column_name: str) -> AsyncJob:
|
||||
"""Refresh values for an existing generated column.
|
||||
|
||||
Returns an :class:`~lancedb.job.AsyncJob` for the refresh operation.
|
||||
Acceptance of the Job does not publish new values; callers must wait
|
||||
and re-read the table to observe refreshed results.
|
||||
"""
|
||||
job = await self._inner._refresh_generated_column(column_name)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def alter_generated_column(
|
||||
self, column_name: str, new_call: _FunctionCall
|
||||
) -> AsyncJob:
|
||||
"""Alter the Function call for an existing generated column.
|
||||
|
||||
Returns an :class:`~lancedb.job.AsyncJob` for the change operation.
|
||||
Acceptance of the Job does not publish the new definition; callers must
|
||||
wait and re-read the table to observe the updated column.
|
||||
"""
|
||||
job = await self._inner._alter_generated_column(column_name, new_call)
|
||||
return AsyncJob(job)
|
||||
|
||||
async def drop_index(self, name: str) -> None:
|
||||
"""
|
||||
Drop an index from the table.
|
||||
|
||||
@@ -1456,6 +1456,408 @@ def test_shuffle_clump_size_yields_all_rows(lance_table):
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# on_transform_error tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BadRowError(ValueError):
|
||||
"""Raised by the failing transforms below when a batch contains a bad id."""
|
||||
|
||||
|
||||
def _failing_transform(bad_ids: set):
|
||||
"""A transform that raises BadRowError whenever the batch has a bad id.
|
||||
|
||||
Raises on the full batch and on any single-row slice containing a bad id,
|
||||
so per-row isolation drops exactly the bad rows.
|
||||
"""
|
||||
|
||||
def transform(batch: pa.RecordBatch) -> list:
|
||||
ids = batch.column("id").to_pylist()
|
||||
bad = sorted(set(ids) & bad_ids)
|
||||
if bad:
|
||||
raise BadRowError(f"bad ids in batch: {bad}")
|
||||
return [{"id": i} for i in ids]
|
||||
|
||||
return transform
|
||||
|
||||
|
||||
def _sequential_split_members(table) -> list[list[int]]:
|
||||
"""Return each split's ids in yield order for shuffle=False.
|
||||
|
||||
With a single rank and no workers the round-robin yields one row per split
|
||||
per cycle, so item k of a clean run belongs to split k % NUM_SPLITS.
|
||||
"""
|
||||
ds = StreamingDataset(table, num_splits=NUM_SPLITS, shuffle=False)
|
||||
members: list[list[int]] = [[] for _ in range(NUM_SPLITS)]
|
||||
for k, row in enumerate(ds):
|
||||
members[k % NUM_SPLITS].append(row["id"])
|
||||
return members
|
||||
|
||||
|
||||
def test_on_transform_error_default_raises(lance_table):
|
||||
"""By default a transform exception propagates and aborts iteration."""
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle_seed=SHUFFLE_SEED,
|
||||
transform=_failing_transform({7}),
|
||||
)
|
||||
with pytest.raises(BadRowError):
|
||||
list(ds)
|
||||
|
||||
|
||||
def test_on_transform_error_invalid_value(lance_table):
|
||||
with pytest.raises(ValueError, match="on_transform_error"):
|
||||
StreamingDataset(lance_table, num_splits=NUM_SPLITS, on_transform_error="bogus")
|
||||
|
||||
|
||||
def test_on_transform_error_skip_drops_bad_rows(lance_table):
|
||||
"""With one bad row per split, 'skip' yields every good row exactly once
|
||||
and counts the dropped rows in rows_skipped."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
bad_ids = {members[i][4] for i in range(NUM_SPLITS)}
|
||||
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
assert ds.rows_skipped == 0
|
||||
|
||||
ids = [row["id"] for row in ds]
|
||||
|
||||
assert sorted(ids) == sorted(set(range(NUM_ROWS)) - bad_ids)
|
||||
assert ds.rows_skipped == NUM_SPLITS
|
||||
|
||||
|
||||
def test_on_transform_error_skip_uneven_ends_at_last_complete_cycle(lance_table):
|
||||
"""When one split loses more rows than the others, the epoch ends at the
|
||||
last cycle where every split still has a row — no crash, no bad rows, and
|
||||
every step remains one sample per split."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
bad_ids = set(members[0][:3]) # all 3 bad rows in split 0
|
||||
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
items = [row["id"] for row in ds]
|
||||
|
||||
rows_per_split = NUM_ROWS // NUM_SPLITS
|
||||
expected_cycles = rows_per_split - len(bad_ids)
|
||||
assert len(items) == expected_cycles * NUM_SPLITS
|
||||
assert len(set(items)) == len(items), "duplicate samples yielded"
|
||||
assert not set(items) & bad_ids, "a bad row was yielded"
|
||||
# Split 0 contributed exactly its surviving rows, in order, one per cycle.
|
||||
survivors = [i for i in members[0] if i not in bad_ids]
|
||||
assert items[0::NUM_SPLITS] == survivors[:expected_cycles]
|
||||
|
||||
|
||||
def test_on_transform_error_warn_logs(lance_table, caplog):
|
||||
"""'warn' skips like 'skip' but logs a warning for the failing batch."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
bad_ids = {members[i][3] for i in range(NUM_SPLITS)}
|
||||
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="warn",
|
||||
)
|
||||
with caplog.at_level(logging.WARNING, logger="lancedb.streaming"):
|
||||
items = list(ds)
|
||||
|
||||
assert len(items) == NUM_ROWS - NUM_SPLITS
|
||||
assert ds.rows_skipped == NUM_SPLITS
|
||||
assert "Skipped" in caplog.text
|
||||
assert "BadRowError" in caplog.text
|
||||
|
||||
|
||||
def test_on_transform_error_callable_selective(lance_table):
|
||||
"""A callable handler can skip expected errors and re-raise the rest."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
bad_ids = {members[i][0] for i in range(NUM_SPLITS)}
|
||||
|
||||
handled: list[Exception] = []
|
||||
|
||||
def handler(exc: Exception) -> bool:
|
||||
handled.append(exc)
|
||||
return isinstance(exc, BadRowError)
|
||||
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error=handler,
|
||||
)
|
||||
items = list(ds)
|
||||
assert len(items) == NUM_ROWS - NUM_SPLITS
|
||||
assert handled and all(isinstance(exc, BadRowError) for exc in handled)
|
||||
|
||||
def broken_transform(batch: pa.RecordBatch) -> list:
|
||||
raise TypeError("boom")
|
||||
|
||||
ds2 = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=broken_transform,
|
||||
on_transform_error=handler,
|
||||
)
|
||||
with pytest.raises(TypeError, match="boom"):
|
||||
list(ds2)
|
||||
|
||||
|
||||
def test_transform_wrong_row_count_raises(lance_table):
|
||||
"""A transform that returns the wrong number of rows is an error even with
|
||||
on_transform_error='skip' — silent shrinkage would corrupt accounting."""
|
||||
|
||||
def drops_rows(batch: pa.RecordBatch) -> list:
|
||||
return batch.column("id").to_pylist()[:-1]
|
||||
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle_seed=SHUFFLE_SEED,
|
||||
transform=drops_rows,
|
||||
on_transform_error="skip",
|
||||
)
|
||||
with pytest.raises(ValueError, match="one output row per input row"):
|
||||
list(ds)
|
||||
|
||||
|
||||
def test_skip_deterministic_across_runs(lance_table):
|
||||
"""With a fixed seed, skipping produces the identical sample sequence on
|
||||
every run — skips are data-dependent, not run-dependent."""
|
||||
bad_ids = {5, 17, 46}
|
||||
|
||||
def run() -> tuple[list[int], int]:
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle_seed=SHUFFLE_SEED,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
return [row["id"] for row in ds], ds.rows_skipped
|
||||
|
||||
ids_a, skipped_a = run()
|
||||
ids_b, skipped_b = run()
|
||||
assert ids_a == ids_b
|
||||
assert skipped_a == skipped_b
|
||||
assert not set(ids_a) & bad_ids
|
||||
|
||||
|
||||
def test_skip_elastic_det_across_world_sizes(lance_table):
|
||||
"""With equal bad-row counts per split, skipping preserves the full
|
||||
elastic-determinism guarantee: identical global batches at every step for
|
||||
every compatible world_size."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
bad_ids = {members[i][6] for i in range(NUM_SPLITS)}
|
||||
|
||||
def collect(world_size: int) -> list[frozenset[int]]:
|
||||
micro = GLOBAL_BATCH_SIZE // world_size
|
||||
iters = [
|
||||
iter(
|
||||
StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
)
|
||||
for rank in range(world_size)
|
||||
]
|
||||
_STOP = object()
|
||||
batches: list[frozenset[int]] = []
|
||||
while True:
|
||||
step_samples: set[int] = set()
|
||||
exhausted = 0
|
||||
for it in iters:
|
||||
for _ in range(micro):
|
||||
val = next(it, _STOP)
|
||||
if val is _STOP:
|
||||
exhausted += 1
|
||||
break
|
||||
step_samples.add(val["id"])
|
||||
if exhausted == len(iters):
|
||||
break
|
||||
assert exhausted == 0, (
|
||||
"Rank iterators exhausted at different steps despite equal "
|
||||
"bad-row counts per split"
|
||||
)
|
||||
batches.append(frozenset(step_samples))
|
||||
return batches
|
||||
|
||||
reference = collect(1)
|
||||
assert len(reference) == NUM_ROWS // NUM_SPLITS - 1
|
||||
for ws in (2, 3, 4):
|
||||
assert collect(ws) == reference, f"world_size={ws} diverged"
|
||||
|
||||
|
||||
def test_resumability_with_skips_same_topology(lance_table):
|
||||
"""Checkpointing mid-epoch with skipped rows resumes exactly: no sample
|
||||
repeated, no sample lost, skipped rows stay skipped."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
# Uneven skips: positions diverge across splits (2 bad in split 0, 1 in
|
||||
# split 5), which only a position-based checkpoint can resume exactly.
|
||||
bad_ids = {members[0][2], members[0][3], members[5][7]}
|
||||
kwargs = dict(
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
|
||||
reference = [row["id"] for row in StreamingDataset(lance_table, **kwargs)]
|
||||
rows_per_split = NUM_ROWS // NUM_SPLITS
|
||||
assert len(reference) == (rows_per_split - 2) * NUM_SPLITS
|
||||
|
||||
steps = 3
|
||||
ds = StreamingDataset(lance_table, **kwargs)
|
||||
it = iter(ds)
|
||||
consumed = [next(it)["id"] for _ in range(steps * NUM_SPLITS)]
|
||||
checkpoint = ds.state_dict()
|
||||
it.close()
|
||||
|
||||
# Split 0 skipped positions 2 and 3 within its first 3 yields; split 5's
|
||||
# bad row is beyond the checkpoint. Everything else is at 3 = the sample
|
||||
# count.
|
||||
positions = checkpoint["positions_consumed_per_split"]
|
||||
assert positions[0] == 5
|
||||
assert positions[1:] == [3] * (NUM_SPLITS - 1)
|
||||
assert checkpoint["samples_consumed_per_split"] == [3] * NUM_SPLITS
|
||||
|
||||
ds2 = StreamingDataset(lance_table, **kwargs)
|
||||
ds2.load_state_dict(checkpoint)
|
||||
resumed = [row["id"] for row in ds2]
|
||||
|
||||
assert consumed == reference[: steps * NUM_SPLITS]
|
||||
assert resumed == reference[steps * NUM_SPLITS :]
|
||||
|
||||
|
||||
def test_resumability_with_skips_elastic_merge(lance_table):
|
||||
"""Elastic resume with skips: each rank's checkpoint knows exact positions
|
||||
only for its own splits; merge_state_dicts recovers the global state, and
|
||||
a run on a different world_size continues exactly."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
# Bad rows early in split 0 (rank 0) and split 6 (rank 1 of a ws=2 run) so
|
||||
# both ranks' position vectors diverge before the checkpoint.
|
||||
bad_ids = {members[0][0], members[0][2], members[6][1]}
|
||||
kwargs = dict(
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
|
||||
reference = [row["id"] for row in StreamingDataset(lance_table, **kwargs)]
|
||||
|
||||
steps = 3
|
||||
world_size = 2
|
||||
micro = GLOBAL_BATCH_SIZE // world_size
|
||||
datasets = [
|
||||
StreamingDataset(lance_table, rank=rank, world_size=world_size, **kwargs)
|
||||
for rank in range(world_size)
|
||||
]
|
||||
iters = [iter(ds) for ds in datasets]
|
||||
seen: list[frozenset[int]] = []
|
||||
for _ in range(steps):
|
||||
step_samples = set()
|
||||
for it in iters:
|
||||
for _ in range(micro):
|
||||
step_samples.add(next(it)["id"])
|
||||
seen.append(frozenset(step_samples))
|
||||
states = [ds.state_dict() for ds in datasets]
|
||||
for it in iters:
|
||||
it.close()
|
||||
|
||||
merged = StreamingDataset.merge_state_dicts(states)
|
||||
expected_positions = [3] * NUM_SPLITS
|
||||
expected_positions[0] = 5 # skipped positions 0 and 2
|
||||
expected_positions[6] = 4 # skipped position 1
|
||||
assert merged["positions_consumed_per_split"] == expected_positions
|
||||
|
||||
# The first 3 global batches match the world_size=1 reference.
|
||||
ref_batches = [
|
||||
frozenset(reference[s * NUM_SPLITS : (s + 1) * NUM_SPLITS])
|
||||
for s in range(len(reference) // NUM_SPLITS)
|
||||
]
|
||||
assert seen == ref_batches[:steps]
|
||||
|
||||
# Resume on world_size=1 from the merged state.
|
||||
ds_resume = StreamingDataset(lance_table, **kwargs)
|
||||
ds_resume.load_state_dict(merged)
|
||||
resumed = [row["id"] for row in ds_resume]
|
||||
assert resumed == reference[steps * NUM_SPLITS :]
|
||||
|
||||
|
||||
def test_rows_skipped_flushed_when_split_entirely_bad(lance_table):
|
||||
"""A split whose rows all fail never completes a cycle, so the epoch ends
|
||||
immediately — but rows_skipped must still report the drops after the
|
||||
iterator exits (the shared-memory counter is flushed on exhaustion)."""
|
||||
members = _sequential_split_members(lance_table)
|
||||
bad_ids = set(members[0]) # every row of split 0 is bad
|
||||
|
||||
ds = StreamingDataset(
|
||||
lance_table,
|
||||
num_splits=NUM_SPLITS,
|
||||
shuffle=False,
|
||||
transform=_failing_transform(bad_ids),
|
||||
on_transform_error="skip",
|
||||
)
|
||||
assert list(ds) == []
|
||||
assert ds.rows_skipped == len(bad_ids)
|
||||
|
||||
|
||||
def test_merge_state_dicts_validates_consistency(lance_table):
|
||||
ds = StreamingDataset(lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED)
|
||||
state = ds.state_dict()
|
||||
other = dict(state, shuffle_seed=SHUFFLE_SEED + 1)
|
||||
with pytest.raises(ValueError, match="shuffle_seed mismatch"):
|
||||
StreamingDataset.merge_state_dicts([state, other])
|
||||
with pytest.raises(ValueError, match="at least one"):
|
||||
StreamingDataset.merge_state_dicts([])
|
||||
|
||||
|
||||
def test_load_state_dict_without_positions_key(lance_table):
|
||||
"""Checkpoints from before positions_consumed_per_split existed still
|
||||
resume exactly (positions equal sample counts when nothing is skipped)."""
|
||||
reference = [
|
||||
row["id"]
|
||||
for row in StreamingDataset(
|
||||
lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED
|
||||
)
|
||||
]
|
||||
|
||||
steps = 4
|
||||
ds = StreamingDataset(lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED)
|
||||
it = iter(ds)
|
||||
for _ in range(steps * NUM_SPLITS):
|
||||
next(it)
|
||||
checkpoint = ds.state_dict()
|
||||
it.close()
|
||||
del checkpoint["positions_consumed_per_split"]
|
||||
|
||||
ds2 = StreamingDataset(
|
||||
lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED
|
||||
)
|
||||
ds2.load_state_dict(checkpoint)
|
||||
resumed = [row["id"] for row in ds2]
|
||||
assert resumed == reference[steps * NUM_SPLITS :]
|
||||
|
||||
|
||||
def test_num_splits_defaults_to_world_size(lance_table):
|
||||
"""Omitting num_splits gives world_size splits (one per rank)."""
|
||||
ds = StreamingDataset(
|
||||
|
||||
@@ -0,0 +1,372 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python exact Function handle call authoring (FF-028)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.expr import Expr, col, func, lit
|
||||
|
||||
_CALL_PATH = "/v1/functions/lookup"
|
||||
_CALL_CATALOG_NAME = "text.normalize.call-name"
|
||||
_CALL_FUNCTION_ID = "fn.exact.call-handle"
|
||||
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_call_xyz_42"
|
||||
_INT_PAYLOAD_SENTINEL = 2_147_000_123
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64).
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
_LIST_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////+4AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAABAAAANz///8c"
|
||||
"AAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECH"
|
||||
"AAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////8AAAAAFAAAAAAAAAAMABQAEgAMAAgABAAMAAAAnAAAAKAAAAAQAAAAAAAEAAgACAAAAAQACAAAAAQAAAA"
|
||||
"BAAAABAAAANz///8cAAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAA"
|
||||
"GAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAAAwAAAAEFSUk9XMQ=="
|
||||
)
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"name",
|
||||
"connection",
|
||||
"table",
|
||||
"snapshot",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"job",
|
||||
"job_id",
|
||||
"artifact",
|
||||
"digest",
|
||||
"retry_key",
|
||||
"idempotency_key",
|
||||
"user_version",
|
||||
"execute",
|
||||
"status",
|
||||
"wait",
|
||||
"cancel",
|
||||
"to_json",
|
||||
"_to_json",
|
||||
"serialize",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire(
|
||||
*,
|
||||
function_id: str = _CALL_FUNCTION_ID,
|
||||
parameters: list[dict[str, str]] | None = None,
|
||||
output_type_ipc: str = _UTF8_TYPE_IPC_B64,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": function_id,
|
||||
"signature": {
|
||||
"parameters": parameters
|
||||
or [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": output_type_ipc,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body(function: dict[str, Any] | None = None) -> bytes:
|
||||
return json.dumps({"function": function or _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _lookup_function(function: dict[str, Any] | None = None):
|
||||
body = _lookup_success_body(function)
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _CALL_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(body)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
return db.functions.get(_CALL_CATALOG_NAME)
|
||||
|
||||
|
||||
def _authored_call_type():
|
||||
cls = getattr(_native, "_FunctionCall", None)
|
||||
if cls is None:
|
||||
pytest.fail("lancedb._lancedb._FunctionCall is missing")
|
||||
return cls
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def test_function_keyword_call_returns_private_frozen_authored_value():
|
||||
function = _lookup_function()
|
||||
assert callable(function)
|
||||
|
||||
authored = function(text=col("text"), limit=8)
|
||||
authored_type = _authored_call_type()
|
||||
assert type(authored) is authored_type
|
||||
assert authored_type.__module__ == "lancedb._lancedb"
|
||||
assert authored_type.__name__ == "_FunctionCall"
|
||||
|
||||
# Keyword order must not matter; bindings store/render in signature order.
|
||||
authored_reversed = function(limit=8, text=col("text"))
|
||||
assert type(authored_reversed) is authored_type
|
||||
rendered = repr(authored_reversed)
|
||||
assert rendered.index("text=") < rendered.index("limit=")
|
||||
assert 'text=field("text")' in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
|
||||
def test_function_call_rejects_positional_missing_and_unknown_args():
|
||||
function = _lookup_function()
|
||||
|
||||
with pytest.raises(TypeError, match="keyword"):
|
||||
function(col("text"), 8)
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="limit"):
|
||||
function(text=col("text"))
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="text"):
|
||||
function(limit=8)
|
||||
|
||||
with pytest.raises((TypeError, ValueError), match="unknown|extra"):
|
||||
function(text=col("text"), limit=8, extra=1)
|
||||
|
||||
|
||||
def test_function_call_accepts_direct_case_sensitive_column_and_rejects_complex_exprs():
|
||||
function = _lookup_function()
|
||||
|
||||
authored = function(text=col("firstName"), limit=1)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert 'text=field("firstName")' in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
complex_exprs = (
|
||||
col("text") + lit("x"),
|
||||
col("text").cast(pa.string()),
|
||||
func("lower", col("text")),
|
||||
col("text") == lit("x"),
|
||||
col("text").lower(),
|
||||
)
|
||||
for expr in complex_exprs:
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=expr, limit=1)
|
||||
|
||||
# Raw native PyExpr is not the public col() wrapper.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=col("text")._inner, limit=1)
|
||||
|
||||
# Non-expression / non-literal objects are rejected for field-shaped misuse
|
||||
# when a column binding is required; plain strings are literals for utf8.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=object(), limit=1)
|
||||
|
||||
|
||||
def test_function_call_plain_literal_declared_type_null_and_nested():
|
||||
function = _lookup_function()
|
||||
|
||||
authored = function(text="hello", limit=7)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
|
||||
# Plain Python int normalizes to declared Int32 and non-null.
|
||||
authored_int32 = function(text="hello", limit=2_147_483_647)
|
||||
assert type(authored_int32) is _authored_call_type()
|
||||
rendered_int32 = repr(authored_int32)
|
||||
assert "limit=literal(Int32, null=false)" in rendered_int32
|
||||
assert "Int64" not in rendered_int32
|
||||
assert "2147483647" not in rendered_int32
|
||||
|
||||
# Plain None keeps each declared parameter type with null=true.
|
||||
authored_null = function(text=None, limit=None)
|
||||
assert type(authored_null) is _authored_call_type()
|
||||
rendered_null = repr(authored_null)
|
||||
assert "text=literal(Utf8, null=true)" in rendered_null
|
||||
assert "limit=literal(Int32, null=true)" in rendered_null
|
||||
|
||||
list_function = _lookup_function(
|
||||
_sample_function_wire(
|
||||
parameters=[
|
||||
{"name": "values", "data_type_ipc": _LIST_INT32_TYPE_IPC_B64},
|
||||
]
|
||||
)
|
||||
)
|
||||
authored_list = list_function(values=[1, 2, 3])
|
||||
assert type(authored_list) is _authored_call_type()
|
||||
rendered_list = repr(authored_list)
|
||||
assert "values=literal(List(Int32), null=false)" in rendered_list
|
||||
assert "[1, 2, 3]" not in rendered_list
|
||||
|
||||
authored_list_null = list_function(values=None)
|
||||
assert type(authored_list_null) is _authored_call_type()
|
||||
rendered_list_null = repr(authored_list_null)
|
||||
assert "values=literal(List(Int32), null=true)" in rendered_list_null
|
||||
|
||||
|
||||
def test_function_call_direct_literal_expr_exact_type_only():
|
||||
function = _lookup_function()
|
||||
|
||||
# lit(int) is Int64 in the expression builder; int32 parameter must reject it.
|
||||
with pytest.raises((TypeError, ValueError), match="limit|int32|type") as raised:
|
||||
function(text="hello", limit=lit(8))
|
||||
reject_text = _exception_text(raised.value)
|
||||
assert "Int64" in reject_text or "int64" in reject_text.lower()
|
||||
assert "Int32" in reject_text or "int32" in reject_text.lower()
|
||||
|
||||
# Exact utf8 literal expression is accepted and stored as Utf8/non-null.
|
||||
authored = function(text=lit("hello"), limit=8)
|
||||
assert type(authored) is _authored_call_type()
|
||||
rendered = repr(authored)
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
assert "hello" not in rendered
|
||||
|
||||
# Cast / arithmetic around a literal is not a direct Literal node.
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
function(text=lit("hello").cast(pa.string()), limit=8)
|
||||
|
||||
|
||||
def test_function_call_conversion_error_and_repr_are_payload_free():
|
||||
function = _lookup_function()
|
||||
|
||||
with pytest.raises((TypeError, ValueError)) as raised:
|
||||
function(text="ok", limit=_LITERAL_PAYLOAD_SENTINEL)
|
||||
text = _exception_text(raised.value)
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
assert "limit" in text
|
||||
assert "int32" in text.lower() or "Int32" in text
|
||||
|
||||
authored = function(text=_LITERAL_PAYLOAD_SENTINEL, limit=_INT_PAYLOAD_SENTINEL)
|
||||
rendered = f"{authored!r}\n{authored!s}"
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in rendered
|
||||
assert str(_INT_PAYLOAD_SENTINEL) not in rendered
|
||||
assert "text=literal(Utf8, null=false)" in rendered
|
||||
assert "limit=literal(Int32, null=false)" in rendered
|
||||
assert type(authored).__name__ == "_FunctionCall"
|
||||
assert "_FunctionCall" in rendered
|
||||
|
||||
|
||||
def test_function_call_private_type_nonconstructible_immutable_and_not_exported():
|
||||
function = _lookup_function()
|
||||
authored = function(text=col("text"), limit=1)
|
||||
authored_type = _authored_call_type()
|
||||
|
||||
assert "_FunctionCall" not in getattr(lancedb, "__all__", [])
|
||||
assert not hasattr(lancedb, "_FunctionCall")
|
||||
assert getattr(_native, "_FunctionCall", None) is authored_type
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
authored_type()
|
||||
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(authored, attr)
|
||||
|
||||
for attr in ("function", "bindings", "arguments", "parameters", "text", "limit"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(authored, attr, None)
|
||||
|
||||
# Existing Function handle stays frozen / connection-free / name-free.
|
||||
assert not hasattr(function, "name")
|
||||
assert not hasattr(function, "connection")
|
||||
with pytest.raises(AttributeError):
|
||||
function.id = "mutated"
|
||||
|
||||
|
||||
def test_function_call_does_not_change_col_query_expression_behavior():
|
||||
# Regression guard: authoring must not alter public col()/Expr query behavior.
|
||||
expr = col("firstName") > lit(1)
|
||||
assert isinstance(expr, Expr)
|
||||
assert expr.to_sql() == "(`firstName` > 1)"
|
||||
@@ -0,0 +1,181 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb import udf
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def double_nullable(value):
|
||||
if value is None:
|
||||
return None
|
||||
return value * 2
|
||||
|
||||
|
||||
def test_first_class_function_enterprise_lifecycle():
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from datetime import timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError
|
||||
from lancedb.expr import col
|
||||
|
||||
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
|
||||
if not host:
|
||||
pytest.skip("LANCEDB_FCF_E2E_HOST is required for the live enterprise test")
|
||||
|
||||
database_uri = os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local")
|
||||
api_key = os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake")
|
||||
run_suffix = uuid.uuid4().hex[:12]
|
||||
table_name = f"fcf_e2e_{run_suffix}"
|
||||
function_name = f"fcf_e2e.double_{run_suffix}"
|
||||
job_timeout = timedelta(minutes=5)
|
||||
query_timeout = timedelta(seconds=30)
|
||||
|
||||
def connect():
|
||||
return lancedb.connect(
|
||||
database_uri,
|
||||
api_key=api_key,
|
||||
host_override=host,
|
||||
)
|
||||
|
||||
setup_db = connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.Table.from_pylist(
|
||||
[
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 5},
|
||||
{"row_id": 3, "value": None},
|
||||
],
|
||||
schema=pa.schema(
|
||||
[
|
||||
pa.field("row_id", pa.int64(), nullable=False),
|
||||
pa.field("value", pa.int64(), nullable=True),
|
||||
]
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, double_nullable)
|
||||
registration_job_id = registration_job.id
|
||||
assert isinstance(registration_job_id, str) and registration_job_id
|
||||
registered_function = registration_job.wait(timeout=job_timeout)
|
||||
assert type(registered_function) is lancedb.Function
|
||||
assert isinstance(registered_function.id, str) and registered_function.id
|
||||
with pytest.raises(AttributeError):
|
||||
registered_function.id = "mutated"
|
||||
|
||||
catalog_reader = connect()
|
||||
function_by_name = catalog_reader.functions.get(function_name)
|
||||
function_by_id = catalog_reader.functions.get_by_id(registered_function.id)
|
||||
expected_signature = ((("value", pa.int64()),), pa.int64(), True)
|
||||
expected_identity = (
|
||||
registered_function.id,
|
||||
*expected_signature,
|
||||
)
|
||||
for function in (registered_function, function_by_name, function_by_id):
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
generated_column_table = catalog_reader.open_table(table_name)
|
||||
generated_column_job = generated_column_table.add_generated_column(
|
||||
"derived",
|
||||
registered_function(value=col("value")),
|
||||
)
|
||||
generated_column_job_id = generated_column_job.id
|
||||
assert isinstance(generated_column_job_id, str) and generated_column_job_id
|
||||
assert generated_column_job.wait(timeout=job_timeout) is None
|
||||
|
||||
complete_reader = connect().open_table(table_name)
|
||||
complete_status = complete_reader.generated_column_status("derived")
|
||||
assert complete_status == "complete"
|
||||
initial_rows = sorted(
|
||||
complete_reader.search()
|
||||
.select(["row_id", "value", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
assert initial_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 5, "derived": 10},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
update_result = complete_reader.update(
|
||||
where="row_id = 2",
|
||||
values={"value": 7},
|
||||
)
|
||||
assert update_result.rows_updated == 1
|
||||
|
||||
incomplete_reader = connect().open_table(table_name)
|
||||
incomplete_status = incomplete_reader.generated_column_status("derived")
|
||||
assert incomplete_status == "incomplete"
|
||||
with pytest.raises(FunctionError) as raised:
|
||||
(
|
||||
incomplete_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout)
|
||||
)
|
||||
assert raised.value.code == "generated_column_incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
assert refresh_job.wait(timeout=job_timeout) is None
|
||||
|
||||
refreshed_reader = connect().open_table(table_name)
|
||||
refreshed_status = refreshed_reader.generated_column_status("derived")
|
||||
assert refreshed_status == "complete"
|
||||
final_rows = sorted(
|
||||
refreshed_reader.search()
|
||||
.select(["row_id", "value", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=query_timeout),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
evidence = {
|
||||
"run_suffix": run_suffix,
|
||||
"database": database_uri.removeprefix("db://"),
|
||||
"table": table_name,
|
||||
"function": function_name,
|
||||
"function_id": registered_function.id,
|
||||
"job_ids": {
|
||||
"register": registration_job_id,
|
||||
"add_generated_column": generated_column_job_id,
|
||||
"refresh_generated_column": refresh_job_id,
|
||||
},
|
||||
"status_transitions": [
|
||||
complete_status,
|
||||
incomplete_status,
|
||||
refreshed_status,
|
||||
],
|
||||
"final_rows": final_rows,
|
||||
}
|
||||
print(json.dumps(evidence, sort_keys=True, separators=(",", ":")))
|
||||
@@ -0,0 +1,595 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb import udf
|
||||
|
||||
|
||||
_RUNNING_DEADLINE_SECONDS = 30
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def reliable_double(value):
|
||||
if value is None:
|
||||
return None
|
||||
return value * 2
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
def terminate_worker_on_input(value):
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
if len(value) == 0:
|
||||
return value
|
||||
except TypeError:
|
||||
pass
|
||||
|
||||
import os
|
||||
|
||||
os._exit(73)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"value": pa.int64()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pyarrow==24.0.0"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def slow_triple(value):
|
||||
import time
|
||||
|
||||
time.sleep(0.02)
|
||||
return value * 3
|
||||
|
||||
|
||||
def _require_live() -> str:
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
|
||||
if not host:
|
||||
pytest.skip(
|
||||
"LANCEDB_FCF_E2E_HOST is required for live enterprise reliability tests"
|
||||
)
|
||||
return host
|
||||
|
||||
|
||||
def _job_timeout():
|
||||
from datetime import timedelta
|
||||
|
||||
return timedelta(minutes=5)
|
||||
|
||||
|
||||
def _query_timeout():
|
||||
from datetime import timedelta
|
||||
|
||||
return timedelta(seconds=30)
|
||||
|
||||
|
||||
def _connect():
|
||||
import os
|
||||
|
||||
import lancedb
|
||||
|
||||
return lancedb.connect(
|
||||
os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local"),
|
||||
api_key=os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake"),
|
||||
host_override=_require_live(),
|
||||
)
|
||||
|
||||
|
||||
def _run_names(case: str) -> tuple[str, str]:
|
||||
import uuid
|
||||
|
||||
suffix = uuid.uuid4().hex[:12]
|
||||
return f"fcf_rel_{case}_{suffix}", f"fcf_rel.{case}_{suffix}"
|
||||
|
||||
|
||||
def _read_rows(table, columns: list[str], row_count: int) -> list[dict]:
|
||||
return sorted(
|
||||
table.search()
|
||||
.select(columns)
|
||||
.limit(row_count)
|
||||
.to_list(timeout=_query_timeout()),
|
||||
key=lambda row: row["row_id"],
|
||||
)
|
||||
|
||||
|
||||
def _emit_evidence(case: str, evidence: dict) -> None:
|
||||
import json
|
||||
|
||||
print(
|
||||
json.dumps(
|
||||
{"case": case, **evidence},
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_core_lifecycle():
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, function_name = _run_names("lifecycle")
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.Table.from_pylist(
|
||||
[
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 5},
|
||||
{"row_id": 3, "value": None},
|
||||
],
|
||||
schema=pa.schema(
|
||||
[
|
||||
pa.field("row_id", pa.int64(), nullable=False),
|
||||
pa.field("value", pa.int64(), nullable=True),
|
||||
]
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, reliable_double)
|
||||
registration_job_id = registration_job.id
|
||||
assert isinstance(registration_job_id, str) and registration_job_id
|
||||
registered = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(registered) is lancedb.Function
|
||||
assert isinstance(registered.id, str) and registered.id
|
||||
with pytest.raises(AttributeError):
|
||||
registered.id = "mutated"
|
||||
|
||||
catalog_reader = _connect()
|
||||
by_name = catalog_reader.functions.get(function_name)
|
||||
by_id = catalog_reader.functions.get_by_id(registered.id)
|
||||
expected_identity = (
|
||||
registered.id,
|
||||
(("value", pa.int64()),),
|
||||
pa.int64(),
|
||||
True,
|
||||
)
|
||||
for function in (registered, by_name, by_id):
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
table = catalog_reader.open_table(table_name)
|
||||
create_job = table.add_generated_column(
|
||||
"derived",
|
||||
registered(value=col("value")),
|
||||
)
|
||||
create_job_id = create_job.id
|
||||
assert isinstance(create_job_id, str) and create_job_id
|
||||
assert create_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
complete_reader = _connect().open_table(table_name)
|
||||
complete_status = complete_reader.generated_column_status("derived")
|
||||
assert complete_status == "complete"
|
||||
initial_rows = _read_rows(
|
||||
complete_reader,
|
||||
["row_id", "value", "derived"],
|
||||
3,
|
||||
)
|
||||
assert initial_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 5, "derived": 10},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
complete_reader.update(where="row_id = 2", values={"value": 7})
|
||||
|
||||
incomplete_reader = _connect().open_table(table_name)
|
||||
changed_rows = _read_rows(incomplete_reader, ["row_id", "value"], 3)
|
||||
assert changed_rows == [
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 7},
|
||||
{"row_id": 3, "value": None},
|
||||
]
|
||||
incomplete_status = incomplete_reader.generated_column_status("derived")
|
||||
assert incomplete_status == "incomplete"
|
||||
with pytest.raises(FunctionError) as raised:
|
||||
(
|
||||
incomplete_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(3)
|
||||
.to_list(timeout=_query_timeout())
|
||||
)
|
||||
assert raised.value.code == "generated_column_incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
assert refresh_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
refreshed_reader = _connect().open_table(table_name)
|
||||
refreshed_status = refreshed_reader.generated_column_status("derived")
|
||||
assert refreshed_status == "complete"
|
||||
final_rows = _read_rows(
|
||||
refreshed_reader,
|
||||
["row_id", "value", "derived"],
|
||||
3,
|
||||
)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"core_lifecycle",
|
||||
{
|
||||
"final_rows": final_rows,
|
||||
"function_id": registered.id,
|
||||
"job_ids": {
|
||||
"create": create_job_id,
|
||||
"refresh": refresh_job_id,
|
||||
"register": registration_job_id,
|
||||
},
|
||||
"status": [
|
||||
complete_status,
|
||||
incomplete_status,
|
||||
refreshed_status,
|
||||
],
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_restart_retention():
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
|
||||
_require_live()
|
||||
raw_evidence = os.environ.get("LANCEDB_FCF_E2E_RESTART_EVIDENCE")
|
||||
if not raw_evidence:
|
||||
pytest.skip(
|
||||
"LANCEDB_FCF_E2E_RESTART_EVIDENCE is required for restart retention"
|
||||
)
|
||||
|
||||
try:
|
||||
evidence = json.loads(raw_evidence)
|
||||
except json.JSONDecodeError as error:
|
||||
pytest.fail(f"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be valid JSON: {error.msg}")
|
||||
|
||||
assert isinstance(evidence, dict), (
|
||||
"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be a JSON object"
|
||||
)
|
||||
table_name = evidence.get("table")
|
||||
function_id = evidence.get("function_id")
|
||||
raw_job_ids = evidence.get("job_ids")
|
||||
assert isinstance(table_name, str) and table_name, (
|
||||
"restart evidence must contain a non-empty table"
|
||||
)
|
||||
assert isinstance(function_id, str) and function_id, (
|
||||
"restart evidence must contain a non-empty function_id"
|
||||
)
|
||||
assert isinstance(raw_job_ids, dict), (
|
||||
"restart evidence must contain a job_ids object"
|
||||
)
|
||||
job_ids = {}
|
||||
for job_kind in ("register", "create", "refresh"):
|
||||
job_id = raw_job_ids.get(job_kind)
|
||||
assert isinstance(job_id, str) and job_id, (
|
||||
f"restart evidence must contain a non-empty job_ids.{job_kind}"
|
||||
)
|
||||
job_ids[job_kind] = job_id
|
||||
|
||||
db = _connect()
|
||||
function = db.functions.get_by_id(function_id)
|
||||
expected_identity = (
|
||||
function_id,
|
||||
(("value", pa.int64()),),
|
||||
pa.int64(),
|
||||
True,
|
||||
)
|
||||
assert type(function) is lancedb.Function
|
||||
assert (
|
||||
function.id,
|
||||
function.parameters,
|
||||
function.output_type,
|
||||
function.output_nullable,
|
||||
) == expected_identity
|
||||
|
||||
jobs = {}
|
||||
for job_kind in ("register", "create", "refresh"):
|
||||
job = db.get_job(job_ids[job_kind])
|
||||
assert job is not None
|
||||
assert job.job_id == job_ids[job_kind]
|
||||
assert job.state == "finished"
|
||||
assert job.failure is None
|
||||
jobs[job_kind] = job
|
||||
|
||||
registered_result = jobs["register"].result
|
||||
assert type(registered_result) is lancedb.Function
|
||||
assert (
|
||||
registered_result.id,
|
||||
registered_result.parameters,
|
||||
registered_result.output_type,
|
||||
registered_result.output_nullable,
|
||||
) == expected_identity
|
||||
assert jobs["create"].result is None
|
||||
assert jobs["refresh"].result is None
|
||||
|
||||
table = db.open_table(table_name)
|
||||
status = table.generated_column_status("derived")
|
||||
assert status == "complete"
|
||||
assert table.count_rows() == 3
|
||||
final_rows = _read_rows(table, ["row_id", "value", "derived"], 3)
|
||||
assert final_rows == [
|
||||
{"row_id": 1, "value": 2, "derived": 4},
|
||||
{"row_id": 2, "value": 7, "derived": 14},
|
||||
{"row_id": 3, "value": None, "derived": None},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"restart_retention",
|
||||
{
|
||||
"final_rows": final_rows,
|
||||
"function_id": function_id,
|
||||
"generated_column_status": status,
|
||||
"job_ids": job_ids,
|
||||
"job_states": {
|
||||
job_kind: jobs[job_kind].state
|
||||
for job_kind in ("register", "create", "refresh")
|
||||
},
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_failure_atomicity_and_worker_recovery():
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import JobFailedError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, failing_function_name = _run_names("worker_failure")
|
||||
_, healthy_function_name = _run_names("worker_recovery")
|
||||
row_count = 4
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.table(
|
||||
{
|
||||
"row_id": list(range(row_count)),
|
||||
"value": [1, 2, 3, 4],
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(
|
||||
failing_function_name,
|
||||
terminate_worker_on_input,
|
||||
)
|
||||
failing_function = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(failing_function) is lancedb.Function
|
||||
|
||||
table = setup_db.open_table(table_name)
|
||||
failed_create_job = table.add_generated_column(
|
||||
"must_not_publish",
|
||||
failing_function(value=col("value")),
|
||||
)
|
||||
failed_job_id = failed_create_job.id
|
||||
assert isinstance(failed_job_id, str) and failed_job_id
|
||||
with pytest.raises(JobFailedError) as raised:
|
||||
failed_create_job.wait(timeout=_job_timeout())
|
||||
assert raised.value.error_code == "udf_execution_failure"
|
||||
|
||||
first_description = _connect().get_job(failed_job_id)
|
||||
second_description = _connect().get_job(failed_job_id)
|
||||
for description in (first_description, second_description):
|
||||
assert description is not None
|
||||
assert description.job_id == failed_job_id
|
||||
assert description.state == "failed"
|
||||
assert description.failure is not None
|
||||
assert description.failure.error_code == "udf_execution_failure"
|
||||
|
||||
atomic_reader = _connect().open_table(table_name)
|
||||
assert "must_not_publish" not in atomic_reader.schema.names
|
||||
assert _read_rows(atomic_reader, ["row_id", "value"], row_count) == [
|
||||
{"row_id": 0, "value": 1},
|
||||
{"row_id": 1, "value": 2},
|
||||
{"row_id": 2, "value": 3},
|
||||
{"row_id": 3, "value": 4},
|
||||
]
|
||||
|
||||
healthy_registration_job = setup_db.functions.register(
|
||||
healthy_function_name,
|
||||
reliable_double,
|
||||
)
|
||||
healthy_function = healthy_registration_job.wait(timeout=_job_timeout())
|
||||
assert type(healthy_function) is lancedb.Function
|
||||
recovery_job = atomic_reader.add_generated_column(
|
||||
"recovered",
|
||||
healthy_function(value=col("value")),
|
||||
)
|
||||
recovery_job_id = recovery_job.id
|
||||
assert isinstance(recovery_job_id, str) and recovery_job_id
|
||||
assert recovery_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
recovered_reader = _connect().open_table(table_name)
|
||||
assert "must_not_publish" not in recovered_reader.schema.names
|
||||
assert recovered_reader.generated_column_status("recovered") == "complete"
|
||||
recovered_rows = _read_rows(
|
||||
recovered_reader,
|
||||
["row_id", "value", "recovered"],
|
||||
row_count,
|
||||
)
|
||||
assert recovered_rows == [
|
||||
{"row_id": 0, "value": 1, "recovered": 2},
|
||||
{"row_id": 1, "value": 2, "recovered": 4},
|
||||
{"row_id": 2, "value": 3, "recovered": 6},
|
||||
{"row_id": 3, "value": 4, "recovered": 8},
|
||||
]
|
||||
|
||||
_emit_evidence(
|
||||
"failure_atomicity_and_worker_recovery",
|
||||
{
|
||||
"failure_code": first_description.failure.error_code,
|
||||
"failed_job_id": failed_job_id,
|
||||
"recovered_rows": recovered_rows,
|
||||
"recovery_job_id": recovery_job_id,
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_enterprise_reliability_concurrent_refresh_fencing():
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import FunctionError, JobFailedError
|
||||
from lancedb.expr import col
|
||||
|
||||
_require_live()
|
||||
table_name, function_name = _run_names("refresh_fencing")
|
||||
row_count = 1024
|
||||
setup_db = _connect()
|
||||
setup_db.create_table(
|
||||
table_name,
|
||||
data=pa.table(
|
||||
{
|
||||
"row_id": list(range(row_count)),
|
||||
"value": list(range(row_count)),
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
registration_job = setup_db.functions.register(function_name, slow_triple)
|
||||
function = registration_job.wait(timeout=_job_timeout())
|
||||
assert type(function) is lancedb.Function
|
||||
|
||||
table = setup_db.open_table(table_name)
|
||||
create_job = table.add_generated_column(
|
||||
"derived",
|
||||
function(value=col("value")),
|
||||
)
|
||||
assert create_job.wait(timeout=_job_timeout()) is None
|
||||
initial_reader = _connect().open_table(table_name)
|
||||
assert initial_reader.generated_column_status("derived") == "complete"
|
||||
|
||||
initial_reader.update(where="row_id = 0", values={"value": 10_000})
|
||||
incomplete_reader = _connect().open_table(table_name)
|
||||
assert incomplete_reader.generated_column_status("derived") == "incomplete"
|
||||
|
||||
refresh_job = incomplete_reader.refresh_generated_column("derived")
|
||||
refresh_job_id = refresh_job.id
|
||||
assert isinstance(refresh_job_id, str) and refresh_job_id
|
||||
deadline = time.monotonic() + _RUNNING_DEADLINE_SECONDS
|
||||
observed_states = []
|
||||
running_observations = 0
|
||||
while running_observations < 2:
|
||||
state = refresh_job.status()
|
||||
if not observed_states or observed_states[-1] != state:
|
||||
observed_states.append(state)
|
||||
if state == "running":
|
||||
running_observations += 1
|
||||
else:
|
||||
running_observations = 0
|
||||
assert state not in {"finished", "failed", "cancelled"}
|
||||
assert time.monotonic() < deadline
|
||||
if running_observations < 2:
|
||||
time.sleep(0.05)
|
||||
|
||||
concurrent_writer = _connect().open_table(table_name)
|
||||
concurrent_writer.update(where="row_id = 1", values={"value": 20_000})
|
||||
with pytest.raises(JobFailedError) as raised:
|
||||
refresh_job.wait(timeout=_job_timeout())
|
||||
assert raised.value.error_code == "stale_or_conflicting_input"
|
||||
|
||||
stale_job = _connect().get_job(refresh_job_id)
|
||||
assert stale_job is not None
|
||||
assert stale_job.job_id == refresh_job_id
|
||||
assert stale_job.state == "failed"
|
||||
assert stale_job.failure is not None
|
||||
assert stale_job.failure.error_code == raised.value.error_code
|
||||
if observed_states[-1] != stale_job.state:
|
||||
observed_states.append(stale_job.state)
|
||||
|
||||
stale_reader = _connect().open_table(table_name)
|
||||
stale_rows = _read_rows(stale_reader, ["row_id", "value"], row_count)
|
||||
assert len(stale_rows) == row_count
|
||||
for row_id, row in enumerate(stale_rows):
|
||||
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
|
||||
assert (row["row_id"], row["value"]) == (row_id, expected_value)
|
||||
assert stale_reader.generated_column_status("derived") == "incomplete"
|
||||
with pytest.raises(FunctionError) as incomplete:
|
||||
(
|
||||
stale_reader.search()
|
||||
.select(["row_id", "derived"])
|
||||
.limit(row_count)
|
||||
.to_list(timeout=_query_timeout())
|
||||
)
|
||||
assert incomplete.value.code == "generated_column_incomplete"
|
||||
|
||||
resubmitted_job = stale_reader.refresh_generated_column("derived")
|
||||
resubmitted_job_id = resubmitted_job.id
|
||||
assert isinstance(resubmitted_job_id, str) and resubmitted_job_id
|
||||
assert resubmitted_job.wait(timeout=_job_timeout()) is None
|
||||
|
||||
final_reader = _connect().open_table(table_name)
|
||||
final_status = final_reader.generated_column_status("derived")
|
||||
assert final_status == "complete"
|
||||
final_rows = _read_rows(
|
||||
final_reader,
|
||||
["row_id", "value", "derived"],
|
||||
row_count,
|
||||
)
|
||||
assert len(final_rows) == row_count
|
||||
final_checksum = 0
|
||||
for row_id, row in enumerate(final_rows):
|
||||
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
|
||||
assert (row["row_id"], row["value"], row["derived"]) == (
|
||||
row_id,
|
||||
expected_value,
|
||||
expected_value * 3,
|
||||
)
|
||||
final_checksum += row["derived"]
|
||||
|
||||
_emit_evidence(
|
||||
"concurrent_refresh_fencing",
|
||||
{
|
||||
"failure_code": stale_job.failure.error_code,
|
||||
"final_checksum": final_checksum,
|
||||
"final_status": final_status,
|
||||
"observed_states": observed_states,
|
||||
"resubmitted_job_id": resubmitted_job_id,
|
||||
"row_count": row_count,
|
||||
"stale_job_id": refresh_job_id,
|
||||
"table": table_name,
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,268 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract: Python projection of JobFailure.error_code / JobFailedError.error_code.
|
||||
|
||||
Public Function failures expose eight stable string categories. Asynchronous
|
||||
errors remain the unified JobFailedError and JobFailureInfo. Python must
|
||||
project the optional exact error_code string already supplied structurally by
|
||||
Rust: preserve a known code, preserve an unknown nonempty future code
|
||||
byte-for-byte, and return None for legacy failure payloads without error_code.
|
||||
Never infer or override a code from message, phase, retryable, HTTP status,
|
||||
job type, or state.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.exceptions import JobFailedError
|
||||
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_KNOWN_CODE = "name_or_function_not_found"
|
||||
_CONFLICTING_STABLE_IN_MESSAGE = "definition_validation_failure"
|
||||
_UNKNOWN_CODE = "enterprise_future_category_xyz"
|
||||
_WAIT_KNOWN_CODE = "unsupported_runtime_or_capability"
|
||||
_WAIT_CONFLICTING_IN_MESSAGE = "revoked_function"
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _failed_describe_body(
|
||||
*,
|
||||
job_id: str,
|
||||
error_code: Optional[str] = None,
|
||||
include_error_code: bool = True,
|
||||
phase: str = "execute",
|
||||
message: str = "worker died",
|
||||
retryable: bool = False,
|
||||
job_type: str = "create_index",
|
||||
) -> dict[str, Any]:
|
||||
failure: dict[str, Any] = {
|
||||
"phase": phase,
|
||||
"message": message,
|
||||
"retryable": retryable,
|
||||
}
|
||||
if include_error_code:
|
||||
failure["error_code"] = error_code
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"job_type": job_type,
|
||||
"job_state": "FAILED",
|
||||
"creation_ms": 1000,
|
||||
"spec": {},
|
||||
"failure": failure,
|
||||
}
|
||||
|
||||
|
||||
def _describe_handler(bodies_by_job_id: dict[str, dict[str, Any]]):
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
payload = json.loads(_read_body(request).decode("utf-8") or "{}")
|
||||
job_id = payload["job_id"]
|
||||
body = bodies_by_job_id.get(job_id)
|
||||
if body is None:
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_known_not_inferred_from_message():
|
||||
"""Structural error_code wins; conflicting message text must not override."""
|
||||
body = _failed_describe_body(
|
||||
job_id="job-known",
|
||||
error_code=_KNOWN_CODE,
|
||||
phase="validate",
|
||||
message=f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}",
|
||||
retryable=False,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-known": body})) as db:
|
||||
description = db.get_job("job-known")
|
||||
assert description is not None
|
||||
failure = description.failure
|
||||
assert failure is not None
|
||||
assert failure.error_code == _KNOWN_CODE
|
||||
assert failure.error_code != _CONFLICTING_STABLE_IN_MESSAGE
|
||||
assert failure.phase == "validate"
|
||||
assert failure.message == f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}"
|
||||
assert failure.retryable is False
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_unknown_preserved_byte_for_byte():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-unknown",
|
||||
error_code=_UNKNOWN_CODE,
|
||||
phase="execute",
|
||||
message=f"new category mentioning {_KNOWN_CODE}",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-unknown": body})) as db:
|
||||
failure = db.get_job("job-unknown").failure
|
||||
assert failure.error_code == _UNKNOWN_CODE
|
||||
assert failure.error_code != _KNOWN_CODE
|
||||
|
||||
|
||||
def test_get_job_failure_error_code_absent_is_none():
|
||||
"""Legacy describe payloads without error_code must not invent a category."""
|
||||
body = _failed_describe_body(
|
||||
job_id="job-legacy",
|
||||
include_error_code=False,
|
||||
phase="execute",
|
||||
message=f"{_KNOWN_CODE} in logs",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-legacy": body})) as db:
|
||||
failure = db.get_job("job-legacy").failure
|
||||
assert failure.error_code is None
|
||||
assert failure.phase == "execute"
|
||||
assert failure.retryable is True
|
||||
|
||||
|
||||
def test_sync_job_wait_job_failed_error_code_known_not_inferred():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-known",
|
||||
error_code=_WAIT_KNOWN_CODE,
|
||||
phase="dispatch",
|
||||
message=f"{_WAIT_CONFLICTING_IN_MESSAGE} in transport logs",
|
||||
retryable=False,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-wait-known": body})) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
db.job("job-wait-known").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, JobFailedError)
|
||||
assert err.error_code == _WAIT_KNOWN_CODE
|
||||
assert err.error_code != _WAIT_CONFLICTING_IN_MESSAGE
|
||||
|
||||
|
||||
def test_sync_job_wait_job_failed_error_code_absent_is_none():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-legacy",
|
||||
include_error_code=False,
|
||||
phase="execute",
|
||||
message=f"{_WAIT_KNOWN_CODE} mentioned only in message",
|
||||
retryable=True,
|
||||
)
|
||||
with _mock_remote_db(_describe_handler({"job-wait-legacy": body})) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
db.job("job-wait-legacy").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert exc_info.value.error_code is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_job_wait_job_failed_error_code_unknown_preserved():
|
||||
body = _failed_describe_body(
|
||||
job_id="job-wait-unknown",
|
||||
error_code=_UNKNOWN_CODE,
|
||||
phase="execute",
|
||||
message=f"future code with {_WAIT_KNOWN_CODE} in text",
|
||||
retryable=False,
|
||||
)
|
||||
async with _mock_remote_db_async(
|
||||
_describe_handler({"job-wait-unknown": body})
|
||||
) as db:
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
await db.job("job-wait-unknown").wait(timeout=timedelta(seconds=5))
|
||||
|
||||
err = exc_info.value
|
||||
assert err.error_code == _UNKNOWN_CODE
|
||||
assert err.error_code != _WAIT_KNOWN_CODE
|
||||
|
||||
|
||||
def test_job_failed_error_legacy_message_construction_error_code_is_none():
|
||||
err = JobFailedError("legacy construction with only a message")
|
||||
assert err.error_code is None
|
||||
|
||||
|
||||
def test_job_failed_error_error_code_is_read_only():
|
||||
err = JobFailedError("message")
|
||||
with pytest.raises(AttributeError):
|
||||
err.error_code = _KNOWN_CODE
|
||||
@@ -0,0 +1,634 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python first-class Function catalog lookup."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_LOOKUP_CATALOG_NAME = "text.normalize.lookup-name"
|
||||
_LOOKUP_FUNCTION_ID = "fn.exact.lookup-handle"
|
||||
_LOOKUP_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_LOOKUP_DIAGNOSTIC_MARKER name=text.normalize.lookup-name "
|
||||
"id=fn.exact.lookup-handle"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_LOOKUP_BODY_MARKER"
|
||||
_UNKNOWN_CODE = "enterprise_future_lookup_category_xyz"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as job-result
|
||||
# tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust FileWriter.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_LOOKUP_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
"list",
|
||||
"alias",
|
||||
"lineage",
|
||||
"FunctionVersion",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _LOOKUP_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body(
|
||||
*,
|
||||
function: dict[str, Any] | None = None,
|
||||
extra_outer: dict[str, Any] | None = None,
|
||||
) -> bytes:
|
||||
body: dict[str, Any] = {"function": function or _sample_function_wire()}
|
||||
if extra_outer:
|
||||
body.update(extra_outer)
|
||||
return json.dumps(body).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _LOOKUP_SERVER_MESSAGE_MARKER not in text
|
||||
assert _LOOKUP_CATALOG_NAME not in text
|
||||
assert _LOOKUP_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_lookup_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _LOOKUP_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _LOOKUP_CATALOG_NAME not in repr(function)
|
||||
assert _LOOKUP_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_name_request(raw: bytes, body: dict[str, Any]) -> None:
|
||||
assert raw
|
||||
assert body == {"name": _LOOKUP_CATALOG_NAME}
|
||||
assert "function_id" not in body
|
||||
|
||||
|
||||
def _assert_id_request(raw: bytes, body: dict[str, Any]) -> None:
|
||||
assert raw
|
||||
assert body == {"function_id": _LOOKUP_FUNCTION_ID}
|
||||
assert "name" not in body
|
||||
|
||||
|
||||
def _assert_native_lookup_methods_present() -> None:
|
||||
assert hasattr(_native.Connection, "_lookup_function_by_name")
|
||||
assert hasattr(_native.Connection, "_lookup_function_by_id")
|
||||
assert callable(getattr(_native.Connection, "_lookup_function_by_name"))
|
||||
assert callable(getattr(_native.Connection, "_lookup_function_by_id"))
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_lookup_methods():
|
||||
_assert_native_lookup_methods_present()
|
||||
|
||||
|
||||
def test_sync_remote_get_by_name_exact_request_and_function_shape():
|
||||
_assert_native_lookup_methods_present()
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
assert not hasattr(db, "get_function")
|
||||
function = db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
_assert_name_request(seen["raw"], seen["body"])
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
def test_sync_remote_get_by_id_exact_request_and_function_shape():
|
||||
_assert_native_lookup_methods_present()
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
function = db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
_assert_id_request(seen["raw"], seen["body"])
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_get_by_name_and_id():
|
||||
_assert_native_lookup_methods_present()
|
||||
name_seen: dict[str, Any] = {}
|
||||
id_seen: dict[str, Any] = {}
|
||||
stage = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
stage["n"] += 1
|
||||
if stage["n"] == 1:
|
||||
name_seen["raw"] = raw
|
||||
name_seen["body"] = body
|
||||
else:
|
||||
id_seen["raw"] = raw
|
||||
id_seen["body"] = body
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
by_name = await db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
by_id = await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
_assert_name_request(name_seen["raw"], name_seen["body"])
|
||||
_assert_id_request(id_seen["raw"], id_seen["body"])
|
||||
_assert_exact_lookup_function(by_name)
|
||||
_assert_exact_lookup_function(by_id)
|
||||
|
||||
|
||||
def test_sync_remote_get_accepts_additive_outer_success_fields():
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
_lookup_success_body(
|
||||
extra_outer={
|
||||
"server_extra": {"ok": True},
|
||||
"request_echo_name": _LOOKUP_CATALOG_NAME,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
function = db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
_assert_exact_lookup_function(function)
|
||||
|
||||
|
||||
def test_empty_name_and_id_reject_before_transport():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.get("")
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.get_by_id("")
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_local_sync_lookup_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_lookup_methods_present()
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_lookup_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_lookup_methods_present()
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
def test_explicit_known_code_is_function_error_with_exact_code():
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
"looks_like": "definition_validation_failure",
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_or_function_not_found"
|
||||
assert err.code != "definition_validation_failure"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_explicit_unknown_code_preserved_despite_status_and_message():
|
||||
body = {
|
||||
"error_code": _UNKNOWN_CODE,
|
||||
"message": (
|
||||
f"{_LOOKUP_SERVER_MESSAGE_MARKER} revoked_function "
|
||||
"name_or_function_not_found"
|
||||
),
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
raw = _read_body(request)
|
||||
assert json.loads(raw.decode("utf-8")) == {"function_id": _LOOKUP_FUNCTION_ID}
|
||||
request.send_response(409)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
|
||||
|
||||
err = exc_info.value
|
||||
assert err.code == _UNKNOWN_CODE
|
||||
assert err.code != "revoked_function"
|
||||
assert err.code != "name_or_function_not_found"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"missing_code_404",
|
||||
404,
|
||||
{
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"empty_code",
|
||||
400,
|
||||
{
|
||||
"error_code": "",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"wrong_type_code",
|
||||
400,
|
||||
{
|
||||
"error_code": 123,
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"null_code",
|
||||
404,
|
||||
{
|
||||
"error_code": None,
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"non_json",
|
||||
404,
|
||||
f"not-json {_LOOKUP_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_invalid_or_missing_error_code_is_payload_free_http(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label # parametrize label for failure diagnosis only
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,response_body",
|
||||
[
|
||||
(
|
||||
"missing_function",
|
||||
{
|
||||
"server_extra": True,
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"null_function",
|
||||
{
|
||||
"function": None,
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"wrong_type_function",
|
||||
{
|
||||
"function": "not-an-object",
|
||||
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
},
|
||||
),
|
||||
(
|
||||
"invalid_function_shape",
|
||||
{
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _LOOKUP_FUNCTION_ID,
|
||||
# missing signature
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_malformed_success_is_payload_free_http(label: str, response_body: dict):
|
||||
del label
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_function_error_surface_omits_server_marker_name_and_id():
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
|
||||
"function_id": _LOOKUP_FUNCTION_ID,
|
||||
"name": _LOOKUP_CATALOG_NAME,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
_read_body(request)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME)
|
||||
|
||||
err = exc_info.value
|
||||
_assert_payload_free(err)
|
||||
assert getattr(err, "code", None) == "name_or_function_not_found"
|
||||
|
||||
|
||||
def test_no_direct_db_lookup_methods_and_no_deleted_keywords():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "lookup_function")
|
||||
assert not hasattr(db, "lookup_function_by_name")
|
||||
assert not hasattr(db, "lookup_function_by_id")
|
||||
assert not hasattr(db, "get_function")
|
||||
assert not hasattr(db.functions, "get_by_name")
|
||||
assert not hasattr(db.functions, "list")
|
||||
|
||||
for keyword in _DELETED_LOOKUP_KEYWORDS:
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.get(_LOOKUP_CATALOG_NAME, **{keyword: True})
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.get_by_id(_LOOKUP_FUNCTION_ID, **{keyword: True})
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_function_error_is_not_top_level_export():
|
||||
assert not hasattr(lancedb, "FunctionError")
|
||||
function_error = _function_error_cls()
|
||||
assert issubclass(function_error, RuntimeError)
|
||||
@@ -0,0 +1,398 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for Python first-class Function registration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable
|
||||
from unittest import mock
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb._udf as _udf_mod
|
||||
import lancedb.job
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_SOURCE_MARKER = "registration-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/registration-redact-token-xyz"
|
||||
_SECRET_ENV = "REGISTER_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.registration-example.com"
|
||||
_FUNCTION_NAME = "text.normalize"
|
||||
_FUNCTION_ID_RETRY = "fn.register-retry-1"
|
||||
_JOB_ID_RETRY = "job-register-retry-1"
|
||||
_JOB_ID_ASYNC = "job-register-async-1"
|
||||
_REGISTER_PATH = "/v1/functions/register"
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
|
||||
_DELETED_REGISTER_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
)
|
||||
|
||||
_SPEC_KEYS = {
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
}
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_register_normalize(text, limit):
|
||||
"""registration-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _definition_json(fn: object) -> dict[str, Any]:
|
||||
payload = _udf_mod._build_function_definition(fn)._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return json.loads(payload.decode("utf-8"))
|
||||
assert isinstance(payload, str)
|
||||
return json.loads(payload)
|
||||
|
||||
|
||||
def _expected_register_spec(name: str, fn: object) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"name": name,
|
||||
"definition": _definition_json(fn),
|
||||
"expected_current_function_id": None,
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_markers_absent_from_exception(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _SOURCE_MARKER not in text
|
||||
assert _SECRET_REFERENCE not in text
|
||||
|
||||
|
||||
def _assert_exact_register_spec(body: dict[str, Any], expected: dict[str, Any]) -> None:
|
||||
assert set(body) == _SPEC_KEYS
|
||||
assert body == expected
|
||||
assert body["format_version"] == 1
|
||||
assert body["expected_current_function_id"] is None
|
||||
assert _SOURCE_MARKER in json.dumps(body["definition"])
|
||||
assert any(
|
||||
capability.get("reference") == _SECRET_REFERENCE
|
||||
for capability in body["definition"]["capabilities"]
|
||||
)
|
||||
|
||||
|
||||
def test_sync_remote_register_retries_exact_wire_and_returns_job():
|
||||
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
|
||||
attempts: list[dict[str, Any]] = []
|
||||
describe_calls: list[dict[str, Any]] = []
|
||||
function_result_wire = {
|
||||
"kind": "function",
|
||||
"format_version": 1,
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _FUNCTION_ID_RETRY,
|
||||
"signature": expected_spec["definition"]["signature"],
|
||||
},
|
||||
}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _REGISTER_PATH:
|
||||
request_id = request.headers.get("x-request-id")
|
||||
attempts.append(
|
||||
{
|
||||
"request_id": request_id,
|
||||
"raw": raw,
|
||||
"body": json.loads(raw.decode("utf-8")),
|
||||
}
|
||||
)
|
||||
if len(attempts) == 1:
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"transient register failure")
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_RETRY}).encode("utf-8"))
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_RETRY
|
||||
describe_calls.append(body)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_RETRY,
|
||||
"job_state": "DONE",
|
||||
"job_type": "register_function",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"result": function_result_wire,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
package_calls = {"n": 0}
|
||||
original_package = _udf_mod._package_udf
|
||||
|
||||
def counting_package(fn: object):
|
||||
package_calls["n"] += 1
|
||||
return original_package(fn)
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
assert not hasattr(db, "register_function")
|
||||
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
|
||||
job = db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_RETRY
|
||||
waited = job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert package_calls["n"] == 1
|
||||
assert len(attempts) == 2
|
||||
first, second = attempts
|
||||
assert isinstance(first["request_id"], str) and first["request_id"]
|
||||
assert first["request_id"] == second["request_id"]
|
||||
assert first["raw"] == second["raw"]
|
||||
assert first["raw"]
|
||||
_assert_exact_register_spec(first["body"], expected_spec)
|
||||
_assert_exact_register_spec(second["body"], expected_spec)
|
||||
|
||||
assert len(describe_calls) == 1
|
||||
assert describe_calls[0]["job_id"] == _JOB_ID_RETRY
|
||||
assert type(waited) is lancedb.Function
|
||||
assert waited.id == _FUNCTION_ID_RETRY
|
||||
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
|
||||
assert waited.output_type == pa.string()
|
||||
assert waited.output_nullable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_register_returns_async_job_with_exact_spec():
|
||||
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REGISTER_PATH
|
||||
raw = _read_body(request)
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
assert not hasattr(db, "register_function")
|
||||
job = await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert seen.get("raw")
|
||||
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
|
||||
_assert_exact_register_spec(seen["body"], expected_spec)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
|
||||
|
||||
def test_sync_remote_register_http_error_omits_source_and_secret_markers():
|
||||
echoed = f"register failed with {_SOURCE_MARKER} and {_SECRET_REFERENCE}"
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
assert request.path == _REGISTER_PATH
|
||||
_read_body(request)
|
||||
request.send_response(400)
|
||||
request.end_headers()
|
||||
request.wfile.write(echoed.encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert received["n"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert err.status_code == 400
|
||||
_assert_markers_absent_from_exception(err)
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_http():
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.register("", packable_register_normalize)
|
||||
|
||||
assert received["n"] == 0
|
||||
|
||||
|
||||
def test_local_sync_register_not_implemented_without_table_mutation(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "register_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_register_not_implemented_without_table_mutation(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "register_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REGISTER_KEYWORDS)
|
||||
def test_register_rejects_deleted_overdesign_keywords_before_submission(keyword):
|
||||
received = {"n": 0}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
received["n"] += 1
|
||||
_read_body(request)
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"should not be reached")
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.register(
|
||||
_FUNCTION_NAME,
|
||||
packable_register_normalize,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert received["n"] == 0
|
||||
@@ -0,0 +1,719 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python conditional first-class Function name removal."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_REMOVE_PATH = "/v1/functions/remove"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_REMOVE_CATALOG_NAME = "text.normalize.remove-name"
|
||||
_REMOVE_FUNCTION_ID = "fn.exact.remove-handle"
|
||||
_REMOVE_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_REMOVE_DIAGNOSTIC_MARKER name=text.normalize.remove-name "
|
||||
"id=fn.exact.remove-handle"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_REMOVE_BODY_MARKER"
|
||||
_CONFLICTING_MESSAGE_CODE = "revoked_function"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# replace tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REMOVE_KEYWORDS = (
|
||||
"expected_current_function_id",
|
||||
"function_id",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"force",
|
||||
"if_exists",
|
||||
"revoke",
|
||||
"delete",
|
||||
)
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _REMOVE_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body() -> bytes:
|
||||
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
def _close_db(db: Any) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
inner = getattr(db, "_conn", None)
|
||||
if inner is not None:
|
||||
inner.close()
|
||||
return
|
||||
close = getattr(db, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _REMOVE_SERVER_MESSAGE_MARKER not in text
|
||||
assert _REMOVE_CATALOG_NAME not in text
|
||||
assert _REMOVE_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_remove_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _REMOVE_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _REMOVE_CATALOG_NAME not in repr(function)
|
||||
assert _REMOVE_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_exact_remove_request(
|
||||
request: http.server.BaseHTTPRequestHandler,
|
||||
raw: bytes,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
expected_id: str,
|
||||
) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REMOVE_PATH
|
||||
assert "?" not in request.path
|
||||
assert raw
|
||||
assert body == {
|
||||
"name": _REMOVE_CATALOG_NAME,
|
||||
"expected_current_function_id": expected_id,
|
||||
}
|
||||
assert set(body) == {"name", "expected_current_function_id"}
|
||||
assert "format_version" not in body
|
||||
assert "function_id" not in body
|
||||
assert "function" not in body
|
||||
assert "signature" not in body
|
||||
assert "job_id" not in body
|
||||
assert "idempotency_key" not in body
|
||||
assert "user_version" not in body
|
||||
assert "force" not in body
|
||||
assert "if_exists" not in body
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
|
||||
|
||||
def _assert_native_remove_method_present() -> None:
|
||||
assert hasattr(_native.Connection, "_remove_function_name")
|
||||
assert callable(getattr(_native.Connection, "_remove_function_name"))
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact name lookup; optionally continue for remove."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _REMOVE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["remove"] = counters.get("remove", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected remove")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
_assert_exact_remove_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
_assert_exact_remove_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_remove_function_name():
|
||||
_assert_native_remove_method_present()
|
||||
|
||||
|
||||
def test_sync_remote_remove_exact_body_path_request_id_returns_none():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
remove_attempts: list[dict[str, Any]] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
remove_attempts.append(
|
||||
{
|
||||
"request": request,
|
||||
"raw": payload,
|
||||
"body": body,
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
}
|
||||
)
|
||||
# Illegal body on 204 must be ignored; success is status-driven only.
|
||||
request.send_response(204)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
assert len(remove_attempts) == 1
|
||||
attempt = remove_attempts[0]
|
||||
_assert_exact_remove_request(
|
||||
attempt["request"],
|
||||
attempt["raw"],
|
||||
attempt["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
assert attempt["body"]["expected_current_function_id"] == current.id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_remove_exact_body_returns_none():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
seen["request"] = request
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
result = await db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
assert seen.get("raw")
|
||||
_assert_exact_remove_request(
|
||||
seen["request"],
|
||||
seen["raw"],
|
||||
seen["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
|
||||
|
||||
def test_after_remove_name_lookup_not_found_id_lookup_same_function():
|
||||
"""Catalog-pointer SDK sequence via a stateful fixture; not server atomicity."""
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {
|
||||
"lookup_name": 0,
|
||||
"lookup_id": 0,
|
||||
"remove": 0,
|
||||
}
|
||||
removed = {"yes": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
if "name" in body:
|
||||
counters["lookup_name"] += 1
|
||||
assert body == {"name": _REMOVE_CATALOG_NAME}
|
||||
if removed["yes"]:
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
counters["lookup_id"] += 1
|
||||
assert body == {"function_id": _REMOVE_FUNCTION_ID}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, raw, body, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
removed["yes"] = True
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(handler) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["lookup_id"] == 0
|
||||
assert counters["remove"] == 0
|
||||
|
||||
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
assert result is None
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["remove"] == 1
|
||||
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.get(_REMOVE_CATALOG_NAME)
|
||||
err = exc_info.value
|
||||
assert err.code == "name_or_function_not_found"
|
||||
_assert_payload_free(err)
|
||||
|
||||
by_id = db.functions.get_by_id(_REMOVE_FUNCTION_ID)
|
||||
|
||||
assert counters["lookup_name"] == 2
|
||||
assert counters["lookup_id"] == 1
|
||||
assert counters["remove"] == 1
|
||||
_assert_exact_remove_function(by_id)
|
||||
assert by_id.id == current.id
|
||||
assert by_id.parameters == current.parameters
|
||||
assert by_id.output_type == current.output_type
|
||||
assert by_id.output_nullable is current.output_nullable
|
||||
|
||||
|
||||
def test_explicit_name_conflict_is_function_error_payload_free():
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
body = {
|
||||
"error_code": "name_conflict",
|
||||
"message": (
|
||||
f"{_REMOVE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE}"
|
||||
),
|
||||
"name": _REMOVE_CATALOG_NAME,
|
||||
"function_id": _REMOVE_FUNCTION_ID,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(409)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_conflict"
|
||||
assert err.code != _CONFLICTING_MESSAGE_CODE
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"200_with_body",
|
||||
200,
|
||||
{
|
||||
"ok": True,
|
||||
"message": _REMOVE_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"job_id": "must-not-infer-job",
|
||||
},
|
||||
),
|
||||
(
|
||||
"202_empty",
|
||||
202,
|
||||
f"{_REMOVE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
("200_empty", 200, ""),
|
||||
],
|
||||
)
|
||||
def test_http_200_202_cannot_return_success(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REMOVE_PATH
|
||||
counters["remove"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_remove_request(
|
||||
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_remove_transport():
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.remove("", current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_current",
|
||||
[
|
||||
_REMOVE_FUNCTION_ID,
|
||||
{"id": _REMOVE_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_current_rejected_without_remove(bad_current):
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-current path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, bad_current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
def test_local_sync_remove_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
_close_db(db)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_remove_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_remove_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.remove(_REMOVE_CATALOG_NAME, current)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REMOVE_KEYWORDS)
|
||||
def test_remove_rejects_deleted_cas_retry_version_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.remove(
|
||||
_REMOVE_CATALOG_NAME,
|
||||
current,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
|
||||
|
||||
def test_no_direct_remove_methods_and_function_has_no_remove_facade_private():
|
||||
counters: dict[str, int] = {"lookup": 0, "remove": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "remove_function")
|
||||
assert not hasattr(db, "remove_function_name")
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
assert callable(getattr(db.functions, "remove", None))
|
||||
assert not hasattr(lancedb, "_SyncFunctions")
|
||||
assert not hasattr(lancedb, "_AsyncFunctions")
|
||||
assert type(db.functions).__name__.startswith("_")
|
||||
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["remove"] == 0
|
||||
@@ -0,0 +1,579 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python conditional first-class Function replacement."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from datetime import timedelta
|
||||
from typing import Any, Callable
|
||||
from unittest import mock
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb._udf as _udf_mod
|
||||
import lancedb.job
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb.exceptions import JobFailedError
|
||||
|
||||
_SOURCE_MARKER = "replace-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/replace-redact-token-xyz"
|
||||
_SECRET_ENV = "REPLACE_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.replace-example.com"
|
||||
_FUNCTION_NAME = "text.normalize"
|
||||
_CURRENT_FUNCTION_ID = "fn.replace-current-1"
|
||||
_REPLACED_FUNCTION_ID = "fn.replace-result-1"
|
||||
_JOB_ID_SYNC = "job-replace-sync-1"
|
||||
_JOB_ID_ASYNC = "job-replace-async-1"
|
||||
_JOB_ID_CONFLICT = "job-replace-conflict-1"
|
||||
_REGISTER_PATH = "/v1/functions/register"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_CONFLICTING_MESSAGE_CODE = "definition_validation_failure"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# job-result tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REPLACE_KEYWORDS = (
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"replace",
|
||||
"expected_current_function_id",
|
||||
"alias",
|
||||
"lineage",
|
||||
)
|
||||
|
||||
_SPEC_KEYS = {
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
}
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_replace_normalize(text, limit):
|
||||
"""replace-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _current_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _CURRENT_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _definition_json(fn: object) -> dict[str, Any]:
|
||||
payload = _udf_mod._build_function_definition(fn)._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return json.loads(payload.decode("utf-8"))
|
||||
assert isinstance(payload, str)
|
||||
return json.loads(payload)
|
||||
|
||||
|
||||
def _expected_replace_spec(name: str, current_id: str, fn: object) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"name": name,
|
||||
"definition": _definition_json(fn),
|
||||
"expected_current_function_id": current_id,
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _assert_exact_replace_spec(
|
||||
body: dict[str, Any], expected: dict[str, Any], current_id: str
|
||||
) -> None:
|
||||
assert set(body) == _SPEC_KEYS
|
||||
assert body == expected
|
||||
assert body["format_version"] == 1
|
||||
assert body["expected_current_function_id"] == current_id
|
||||
assert body["expected_current_function_id"] is not None
|
||||
assert _SOURCE_MARKER in json.dumps(body["definition"])
|
||||
assert any(
|
||||
capability.get("reference") == _SECRET_REFERENCE
|
||||
for capability in body["definition"]["capabilities"]
|
||||
)
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact lookup; optionally continue for register/describe."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _FUNCTION_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps({"function": _current_function_wire()}).encode("utf-8")
|
||||
)
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["register"] = counters.get("register", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected register")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_FUNCTION_NAME)
|
||||
assert type(current) is lancedb.Function
|
||||
assert current.id == _CURRENT_FUNCTION_ID
|
||||
assert not hasattr(current, "name")
|
||||
assert not hasattr(current, "replace")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_FUNCTION_NAME)
|
||||
assert type(current) is lancedb.Function
|
||||
assert current.id == _CURRENT_FUNCTION_ID
|
||||
assert not hasattr(current, "name")
|
||||
assert not hasattr(current, "replace")
|
||||
return current
|
||||
|
||||
|
||||
def test_sync_remote_replace_exact_body_one_package_job_and_function_result():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
|
||||
expected_spec = _expected_replace_spec(
|
||||
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
|
||||
)
|
||||
register_attempts: list[dict[str, Any]] = []
|
||||
function_result_wire = {
|
||||
"kind": "function",
|
||||
"format_version": 1,
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _REPLACED_FUNCTION_ID,
|
||||
"signature": expected_spec["definition"]["signature"],
|
||||
},
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
if request.path == _REGISTER_PATH:
|
||||
counters["register"] += 1
|
||||
register_attempts.append(
|
||||
{
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
"raw": payload,
|
||||
"body": json.loads(payload.decode("utf-8")),
|
||||
}
|
||||
)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_SYNC}).encode("utf-8"))
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
counters["describe"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_SYNC
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_SYNC,
|
||||
"job_state": "DONE",
|
||||
"job_type": "register_function",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"result": function_result_wire,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
package_calls = {"n": 0}
|
||||
original_package = _udf_mod._package_udf
|
||||
|
||||
def counting_package(fn: object):
|
||||
package_calls["n"] += 1
|
||||
return original_package(fn)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "replace_function")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
|
||||
job = db.functions.replace(
|
||||
_FUNCTION_NAME, current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
waited = job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert package_calls["n"] == 1
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert counters["describe"] == 1
|
||||
assert len(register_attempts) == 1
|
||||
attempt = register_attempts[0]
|
||||
assert isinstance(attempt["request_id"], str) and attempt["request_id"]
|
||||
assert attempt["raw"]
|
||||
_assert_exact_replace_spec(attempt["body"], expected_spec, current.id)
|
||||
assert type(waited) is lancedb.Function
|
||||
assert waited.id == _REPLACED_FUNCTION_ID
|
||||
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
|
||||
assert waited.output_type == pa.string()
|
||||
assert waited.output_nullable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_replace_exact_body_returns_async_job():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
expected_spec = _expected_replace_spec(
|
||||
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
|
||||
)
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REGISTER_PATH
|
||||
counters["register"] += 1
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "replace_function")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
job = await db.functions.replace(
|
||||
_FUNCTION_NAME, current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert seen.get("raw")
|
||||
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
|
||||
_assert_exact_replace_spec(seen["body"], expected_spec, current.id)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
|
||||
|
||||
def test_sync_remote_replace_failed_name_conflict_raises_job_failed_error_code():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
if request.path == _REGISTER_PATH:
|
||||
counters["register"] += 1
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps({"job_id": _JOB_ID_CONFLICT}).encode("utf-8")
|
||||
)
|
||||
return
|
||||
|
||||
assert request.path == _DESCRIBE_PATH
|
||||
counters["describe"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body["job_id"] == _JOB_ID_CONFLICT
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
"job_id": _JOB_ID_CONFLICT,
|
||||
"job_type": "register_function",
|
||||
"job_state": "FAILED",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
"failure": {
|
||||
"phase": "validate",
|
||||
"message": (
|
||||
f"looks like {_CONFLICTING_MESSAGE_CODE} during CAS"
|
||||
),
|
||||
"retryable": False,
|
||||
"error_code": "name_conflict",
|
||||
},
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
job = db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
assert type(job) is lancedb.job.Job
|
||||
with pytest.raises(JobFailedError) as exc_info:
|
||||
job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 1
|
||||
assert counters["describe"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, JobFailedError)
|
||||
assert err.error_code == "name_conflict"
|
||||
assert err.error_code != _CONFLICTING_MESSAGE_CODE
|
||||
|
||||
|
||||
def test_empty_name_rejects_before_register_transport():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(ValueError):
|
||||
db.functions.replace("", current, packable_replace_normalize)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_current",
|
||||
[
|
||||
_CURRENT_FUNCTION_ID,
|
||||
{"id": _CURRENT_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_current_rejected_without_register(bad_current):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-current path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.replace(
|
||||
_FUNCTION_NAME, bad_current, packable_replace_normalize
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
def test_local_sync_replace_not_implemented_without_table_mutation(tmp_path):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "replace_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_replace_not_implemented_without_table_mutation(tmp_path):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "replace_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REPLACE_KEYWORDS)
|
||||
def test_replace_rejects_deleted_cas_retry_version_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.replace(
|
||||
_FUNCTION_NAME,
|
||||
current,
|
||||
packable_replace_normalize,
|
||||
**{keyword: True},
|
||||
)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
|
||||
|
||||
def test_no_direct_replace_function_methods_and_function_has_no_replace():
|
||||
counters: dict[str, int] = {"lookup": 0, "register": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "replace_function")
|
||||
assert not hasattr(db, "register_function")
|
||||
assert not hasattr(current, "replace")
|
||||
assert not hasattr(current, "replace_function")
|
||||
assert callable(getattr(db.functions, "replace", None))
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["register"] == 0
|
||||
@@ -0,0 +1,728 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python exact first-class Function revocation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.errors import HttpError
|
||||
|
||||
_REVOKE_PATH = "/v1/functions/revoke"
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_REVOKE_CATALOG_NAME = "text.normalize.revoke-name"
|
||||
_REVOKE_FUNCTION_ID = "fn.exact.revoke-handle"
|
||||
_REVOKE_SERVER_MESSAGE_MARKER = (
|
||||
"SERVER_REVOKE_DIAGNOSTIC_MARKER id=fn.exact.revoke-handle "
|
||||
"name=text.normalize.revoke-name"
|
||||
)
|
||||
_SENSITIVE_BODY_MARKER = "SENSITIVE_REVOKE_BODY_MARKER"
|
||||
_CONFLICTING_MESSAGE_CODE = "revoked_function"
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
|
||||
# remove tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
|
||||
_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_DELETED_REVOKE_KEYWORDS = (
|
||||
"function_id",
|
||||
"name",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"version",
|
||||
"reason",
|
||||
"expiry",
|
||||
"force",
|
||||
"if_exists",
|
||||
"remove",
|
||||
"delete",
|
||||
)
|
||||
|
||||
|
||||
def _function_error_cls(*, required: bool = True):
|
||||
"""Resolve FunctionError from the live module (records RED when absent)."""
|
||||
from lancedb import exceptions as exc_mod
|
||||
|
||||
cls = getattr(exc_mod, "FunctionError", None)
|
||||
if cls is None:
|
||||
if required:
|
||||
pytest.fail("lancedb.exceptions.FunctionError is missing")
|
||||
return type("MissingFunctionError", (), {})
|
||||
return cls
|
||||
|
||||
|
||||
def _sample_function_wire() -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": _REVOKE_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _lookup_success_body() -> bytes:
|
||||
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
def _close_db(db: Any) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
inner = getattr(db, "_conn", None)
|
||||
if inner is not None:
|
||||
inner.close()
|
||||
return
|
||||
close = getattr(db, "close", None)
|
||||
if callable(close):
|
||||
close()
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
db = None
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
if db is not None:
|
||||
_close_db(db)
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_chain_text(exc: BaseException) -> str:
|
||||
parts: list[str] = []
|
||||
seen: set[int] = set()
|
||||
current: BaseException | None = exc
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(str(current))
|
||||
parts.append(repr(current))
|
||||
current = current.__cause__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _assert_payload_free(exc: BaseException) -> None:
|
||||
text = _exception_chain_text(exc)
|
||||
assert _REVOKE_SERVER_MESSAGE_MARKER not in text
|
||||
assert _REVOKE_CATALOG_NAME not in text
|
||||
assert _REVOKE_FUNCTION_ID not in text
|
||||
assert _SENSITIVE_BODY_MARKER not in text
|
||||
|
||||
|
||||
def _assert_exact_revoke_function(function: object) -> None:
|
||||
assert type(function) is lancedb.Function
|
||||
assert function.id == _REVOKE_FUNCTION_ID
|
||||
assert not hasattr(function, "name")
|
||||
assert function.parameters == (
|
||||
("text", pa.string()),
|
||||
("limit", pa.int32()),
|
||||
)
|
||||
assert function.output_type == pa.string()
|
||||
assert function.output_nullable is True
|
||||
assert _REVOKE_CATALOG_NAME not in repr(function)
|
||||
assert _REVOKE_CATALOG_NAME not in str(function)
|
||||
|
||||
|
||||
def _assert_exact_revoke_request(
|
||||
request: http.server.BaseHTTPRequestHandler,
|
||||
raw: bytes,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
expected_id: str,
|
||||
) -> None:
|
||||
assert request.command == "POST"
|
||||
assert request.path == _REVOKE_PATH
|
||||
assert "?" not in request.path
|
||||
assert "remove" not in request.path
|
||||
assert raw
|
||||
assert body == {"function_id": expected_id}
|
||||
assert set(body) == {"function_id"}
|
||||
assert "name" not in body
|
||||
assert "expected_current_function_id" not in body
|
||||
assert "format_version" not in body
|
||||
assert "function" not in body
|
||||
assert "signature" not in body
|
||||
assert "job_id" not in body
|
||||
assert "idempotency_key" not in body
|
||||
assert "user_version" not in body
|
||||
assert "reason" not in body
|
||||
assert "expiry" not in body
|
||||
assert "force" not in body
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
|
||||
|
||||
def _assert_native_revoke_method_present() -> None:
|
||||
assert hasattr(_native.Connection, "_revoke_function")
|
||||
assert callable(getattr(_native.Connection, "_revoke_function"))
|
||||
|
||||
|
||||
def _lookup_success_handler(
|
||||
counters: dict[str, int],
|
||||
*,
|
||||
after_lookup: (
|
||||
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
|
||||
) = None,
|
||||
):
|
||||
"""Serve exact name lookup; optionally continue for revoke."""
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
counters["lookup"] = counters.get("lookup", 0) + 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
assert body == {"name": _REVOKE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
if after_lookup is not None:
|
||||
after_lookup(request, raw)
|
||||
return
|
||||
counters["revoke"] = counters.get("revoke", 0) + 1
|
||||
request.send_response(500)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected revoke")
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _observe_current(db) -> lancedb.Function:
|
||||
current = db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
_assert_exact_revoke_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
async def _observe_current_async(db) -> lancedb.Function:
|
||||
current = await db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
_assert_exact_revoke_function(current)
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
return current
|
||||
|
||||
|
||||
def test_native_connection_exposes_private_revoke_function():
|
||||
_assert_native_revoke_method_present()
|
||||
|
||||
|
||||
def test_sync_remote_revoke_exact_body_path_request_id_returns_none():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
revoke_attempts: list[dict[str, Any]] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
revoke_attempts.append(
|
||||
{
|
||||
"request": request,
|
||||
"raw": payload,
|
||||
"body": body,
|
||||
"request_id": request.headers.get("x-request-id"),
|
||||
}
|
||||
)
|
||||
# Illegal body on 204 must be ignored; success is status-driven only.
|
||||
request.send_response(204)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(
|
||||
json.dumps(
|
||||
{
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"message": _REVOKE_SERVER_MESSAGE_MARKER,
|
||||
}
|
||||
).encode("utf-8")
|
||||
)
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "revoke_function")
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
result = db.functions.revoke(current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert len(revoke_attempts) == 1
|
||||
attempt = revoke_attempts[0]
|
||||
_assert_exact_revoke_request(
|
||||
attempt["request"],
|
||||
attempt["raw"],
|
||||
attempt["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
assert attempt["body"]["function_id"] == current.id
|
||||
_assert_exact_revoke_function(current)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_revoke_exact_body_returns_none():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
seen["request"] = request
|
||||
seen["raw"] = raw
|
||||
seen["body"] = json.loads(raw.decode("utf-8"))
|
||||
seen["request_id"] = request.headers.get("x-request-id")
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
assert not hasattr(db, "revoke_function")
|
||||
current = await _observe_current_async(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
result = await db.functions.revoke(current)
|
||||
|
||||
assert result is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert seen.get("raw")
|
||||
_assert_exact_revoke_request(
|
||||
seen["request"],
|
||||
seen["raw"],
|
||||
seen["body"],
|
||||
expected_id=current.id,
|
||||
)
|
||||
|
||||
|
||||
def test_repeated_remote_revoke_204_both_return_none():
|
||||
"""Two logical calls each receiving 204 both succeed (Python outcome only)."""
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
revoke_request_ids: list[str] = []
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, body, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request_id = request.headers.get("x-request-id")
|
||||
assert isinstance(request_id, str) and request_id
|
||||
revoke_request_ids.append(request_id)
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
first = db.functions.revoke(current)
|
||||
second = db.functions.revoke(current)
|
||||
|
||||
assert first is None
|
||||
assert second is None
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 2
|
||||
assert len(revoke_request_ids) == 2
|
||||
_assert_exact_revoke_function(current)
|
||||
|
||||
|
||||
def test_after_revoke_name_and_id_lookup_still_return_function():
|
||||
"""Revoke does not unlink names; SDK-visible sequence only, not Sophon proof."""
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {
|
||||
"lookup_name": 0,
|
||||
"lookup_id": 0,
|
||||
"revoke": 0,
|
||||
}
|
||||
revoked = {"yes": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
if request.path == _LOOKUP_PATH:
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
if "name" in body:
|
||||
counters["lookup_name"] += 1
|
||||
assert body == {"name": _REVOKE_CATALOG_NAME}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
counters["lookup_id"] += 1
|
||||
assert body == {"function_id": _REVOKE_FUNCTION_ID}
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(_lookup_success_body())
|
||||
return
|
||||
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
body = json.loads(raw.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, raw, body, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
revoked["yes"] = True
|
||||
request.send_response(204)
|
||||
request.end_headers()
|
||||
|
||||
with _mock_remote_db(handler) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["lookup_id"] == 0
|
||||
assert counters["revoke"] == 0
|
||||
assert not revoked["yes"]
|
||||
|
||||
result = db.functions.revoke(current)
|
||||
assert result is None
|
||||
assert counters["lookup_name"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
assert revoked["yes"]
|
||||
|
||||
by_name = db.functions.get(_REVOKE_CATALOG_NAME)
|
||||
by_id = db.functions.get_by_id(_REVOKE_FUNCTION_ID)
|
||||
|
||||
assert counters["lookup_name"] == 2
|
||||
assert counters["lookup_id"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
_assert_exact_revoke_function(by_name)
|
||||
_assert_exact_revoke_function(by_id)
|
||||
assert by_name.id == current.id
|
||||
assert by_id.id == current.id
|
||||
assert by_name.parameters == current.parameters
|
||||
assert by_id.parameters == current.parameters
|
||||
assert by_name.output_type == current.output_type
|
||||
assert by_id.output_type == current.output_type
|
||||
assert by_name.output_nullable is current.output_nullable
|
||||
assert by_id.output_nullable is current.output_nullable
|
||||
|
||||
|
||||
def test_explicit_name_or_function_not_found_is_function_error_payload_free():
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
body = {
|
||||
"error_code": "name_or_function_not_found",
|
||||
"message": (
|
||||
f"{_REVOKE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE} "
|
||||
"name_conflict"
|
||||
),
|
||||
"name": _REVOKE_CATALOG_NAME,
|
||||
"function_id": _REVOKE_FUNCTION_ID,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(404)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(body).encode("utf-8"))
|
||||
|
||||
function_error = _function_error_cls()
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(function_error) as exc_info:
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, function_error)
|
||||
assert err.code == "name_or_function_not_found"
|
||||
assert err.code != _CONFLICTING_MESSAGE_CODE
|
||||
assert err.code != "name_conflict"
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,status,response_body",
|
||||
[
|
||||
(
|
||||
"200_with_body",
|
||||
200,
|
||||
{
|
||||
"ok": True,
|
||||
"message": _REVOKE_SERVER_MESSAGE_MARKER,
|
||||
_SENSITIVE_BODY_MARKER: True,
|
||||
"job_id": "must-not-infer-job",
|
||||
},
|
||||
),
|
||||
(
|
||||
"202_empty",
|
||||
202,
|
||||
f"{_REVOKE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
|
||||
),
|
||||
("200_empty", 200, ""),
|
||||
],
|
||||
)
|
||||
def test_http_200_202_cannot_return_success(
|
||||
label: str, status: int, response_body: object
|
||||
):
|
||||
del label
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
def after_lookup(
|
||||
request: http.server.BaseHTTPRequestHandler, payload: bytes
|
||||
) -> None:
|
||||
assert request.path == _REVOKE_PATH
|
||||
counters["revoke"] += 1
|
||||
parsed = json.loads(payload.decode("utf-8"))
|
||||
_assert_exact_revoke_request(
|
||||
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
|
||||
)
|
||||
request.send_response(status)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
if isinstance(response_body, str):
|
||||
request.wfile.write(response_body.encode("utf-8"))
|
||||
else:
|
||||
request.wfile.write(json.dumps(response_body).encode("utf-8"))
|
||||
|
||||
with _mock_remote_db(
|
||||
_lookup_success_handler(counters, after_lookup=after_lookup)
|
||||
) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(HttpError) as exc_info:
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 1
|
||||
err = exc_info.value
|
||||
assert isinstance(err, HttpError)
|
||||
assert not isinstance(err, _function_error_cls(required=False))
|
||||
_assert_payload_free(err)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_function",
|
||||
[
|
||||
_REVOKE_FUNCTION_ID,
|
||||
{"id": _REVOKE_FUNCTION_ID},
|
||||
object(),
|
||||
123,
|
||||
],
|
||||
)
|
||||
def test_raw_id_or_arbitrary_function_rejected_without_revoke(bad_function):
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
# Observe a real handle separately so the bad-function path is isolated.
|
||||
_ = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.revoke(bad_function)
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
|
||||
def test_local_sync_revoke_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
before = db.list_tables().tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "revoke_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
db.functions.revoke(current)
|
||||
|
||||
assert db.list_tables().tables == before
|
||||
_close_db(db)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_revoke_not_implemented_without_table_mutation(tmp_path):
|
||||
_assert_native_revoke_method_present()
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
|
||||
current = _observe_current(remote_db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
assert type(current) is lancedb.Function
|
||||
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
before = (await db.list_tables()).tables
|
||||
assert before == []
|
||||
assert not hasattr(db, "revoke_function")
|
||||
|
||||
with pytest.raises(NotImplementedError):
|
||||
await db.functions.revoke(current)
|
||||
|
||||
assert (await db.list_tables()).tables == before
|
||||
db.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("keyword", _DELETED_REVOKE_KEYWORDS)
|
||||
def test_revoke_rejects_overdesigned_kwargs(keyword):
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
with pytest.raises(TypeError):
|
||||
db.functions.revoke(current, **{keyword: True})
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
|
||||
|
||||
def test_no_direct_revoke_methods_and_function_has_no_revoke_facade_private():
|
||||
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
|
||||
|
||||
with _mock_remote_db(_lookup_success_handler(counters)) as db:
|
||||
current = _observe_current(db)
|
||||
assert not hasattr(db, "revoke_function")
|
||||
assert not hasattr(current, "remove")
|
||||
assert not hasattr(current, "delete")
|
||||
assert not hasattr(current, "revoke")
|
||||
assert callable(getattr(db.functions, "revoke", None))
|
||||
assert not hasattr(lancedb, "_SyncFunctions")
|
||||
assert not hasattr(lancedb, "_AsyncFunctions")
|
||||
assert type(db.functions).__name__.startswith("_")
|
||||
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
|
||||
|
||||
assert counters["lookup"] == 1
|
||||
assert counters["revoke"] == 0
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,899 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python ``table.add_generated_column`` (FF-032).
|
||||
|
||||
Public user shape under test:
|
||||
|
||||
job = table.add_generated_column(
|
||||
"normalized_text",
|
||||
normalize(text=col("text")),
|
||||
)
|
||||
job.wait()
|
||||
|
||||
These tests exercise the live worktree PyO3 extension and public sync/async
|
||||
wrappers. While the public methods and hidden native bridge are absent they
|
||||
fail against that extension; once present they freeze the public contract
|
||||
below. They must not fake success paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import inspect
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb.job
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.expr import col
|
||||
from lancedb.remote.table import RemoteTable
|
||||
from lancedb.table import AsyncTable, LanceTable, Table
|
||||
|
||||
_LOOKUP_PATH = "/v1/functions/lookup"
|
||||
_JOB_DESCRIBE_PATH = "/v1/jobs/describe"
|
||||
_TABLE_NAME = "articles"
|
||||
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
|
||||
_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/generated_columns/create/"
|
||||
_BRANCHES_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/branches/create/"
|
||||
_BRANCHES_LIST_PATH = f"/v1/table/{_TABLE_NAME}/branches/list/"
|
||||
|
||||
_CATALOG_NAME = "text.normalize"
|
||||
_FUNCTION_ID = "fn.exact.normalize.gen-col"
|
||||
_JOB_ID_SYNC = "job-create-gen-col-sync-1"
|
||||
_JOB_ID_ASYNC = "job-create-gen-col-async-1"
|
||||
_JOB_ID_BRANCH = "job-create-gen-col-branch-1"
|
||||
_SOURCE_TABLE_VERSION = 42
|
||||
_TEXT_FIELD_ID = 7
|
||||
_BRANCH_NAME = "exp"
|
||||
_BRANCH_SOURCE_VERSION = 9
|
||||
_BRANCH_TEXT_FIELD_ID = 11
|
||||
|
||||
_DESCRIBE_BODY_MARKER = "SENSITIVE_DESCRIBE_BODY_MARKER_gen_col_xyz"
|
||||
_CREATE_RESPONSE_MARKER = "SENSITIVE_CREATE_RESPONSE_MARKER_gen_col_xyz"
|
||||
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_gen_col_xyz"
|
||||
|
||||
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_FORBIDDEN_PUBLIC_NAMES = (
|
||||
"FunctionCall",
|
||||
"BoundFunctionCall",
|
||||
"AuthoredFunctionCall",
|
||||
"CreateGeneratedColumnRequest",
|
||||
"CreateGeneratedColumnJobSpec",
|
||||
"GeneratedColumnBindingSnapshot",
|
||||
"GeneratedColumnCreateRequest",
|
||||
"geneva",
|
||||
"GenevaFunction",
|
||||
"VirtualColumnDefinition",
|
||||
)
|
||||
|
||||
_FORBIDDEN_METHOD_KWARGS = (
|
||||
"source_table_version",
|
||||
"version",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"output",
|
||||
"output_type",
|
||||
"output_nullable",
|
||||
"nullable",
|
||||
"spec",
|
||||
"retry_key",
|
||||
"idempotency_key",
|
||||
"request",
|
||||
"envelope",
|
||||
"table_ref",
|
||||
"branch",
|
||||
)
|
||||
|
||||
|
||||
def _sample_function_wire(
|
||||
*,
|
||||
function_id: str = _FUNCTION_ID,
|
||||
parameters: list[dict[str, str]] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"format_version": 1,
|
||||
"id": function_id,
|
||||
"signature": {
|
||||
"parameters": parameters
|
||||
or [
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _text_schema_fields(
|
||||
*, arrow_type: str = "string", nullable: bool = True
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"fields": [
|
||||
{
|
||||
"name": "text",
|
||||
"type": {"type": arrow_type},
|
||||
"nullable": nullable,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _describe_body(
|
||||
*,
|
||||
version: int = _SOURCE_TABLE_VERSION,
|
||||
field_ids: list[int] | None = None,
|
||||
arrow_type: str = "string",
|
||||
include_marker: bool = True,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"version": version,
|
||||
"schema": _text_schema_fields(arrow_type=arrow_type),
|
||||
"field_ids": field_ids if field_ids is not None else [_TEXT_FIELD_ID],
|
||||
}
|
||||
if include_marker:
|
||||
body["server_diagnostic"] = _DESCRIBE_BODY_MARKER
|
||||
return body
|
||||
|
||||
|
||||
def _create_gen_column_done_body(job_id: str) -> dict[str, Any]:
|
||||
# DONE with omitted result: create_gen_column projects JobResult::None.
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"job_state": "DONE",
|
||||
"job_type": "create_gen_column",
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
}
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _json_response(
|
||||
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
|
||||
) -> None:
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(payload)
|
||||
|
||||
|
||||
def _lookup_function(db: Any) -> lancedb.Function:
|
||||
return db.functions.get(_CATALOG_NAME)
|
||||
|
||||
|
||||
class _RequestLog:
|
||||
"""Track lookup/describe/create after setup; setup traffic is excluded."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.lookup: list[dict[str, Any]] = []
|
||||
self.describe: list[dict[str, Any]] = []
|
||||
self.create: list[dict[str, Any]] = []
|
||||
self.other_table: list[str] = []
|
||||
self.recording = False
|
||||
|
||||
def start(self) -> None:
|
||||
# Drop setup's explicit Function lookup and open_table describe so
|
||||
# operation accounting cannot be polluted by fixture traffic.
|
||||
self.lookup.clear()
|
||||
self.describe.clear()
|
||||
self.create.clear()
|
||||
self.other_table.clear()
|
||||
self.recording = True
|
||||
|
||||
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
|
||||
if not self.recording:
|
||||
return
|
||||
if path == _LOOKUP_PATH:
|
||||
self.lookup.append(body or {})
|
||||
elif path == _DESCRIBE_PATH:
|
||||
self.describe.append(body or {})
|
||||
elif path == _CREATE_PATH:
|
||||
self.create.append(body or {})
|
||||
elif path.startswith(f"/v1/table/{_TABLE_NAME}/"):
|
||||
self.other_table.append(path)
|
||||
|
||||
|
||||
def _assert_no_operation_traffic(log: _RequestLog) -> None:
|
||||
assert log.lookup == []
|
||||
assert log.describe == []
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
|
||||
|
||||
def _assert_exact_public_signature(method: Any) -> None:
|
||||
"""Freeze ``(self, column_name, call)`` with no varargs/kwargs escape hatches."""
|
||||
params = list(inspect.signature(method).parameters.values())
|
||||
assert [p.name for p in params] == ["self", "column_name", "call"]
|
||||
for param in params:
|
||||
assert param.kind in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
)
|
||||
assert param.default is inspect.Parameter.empty
|
||||
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
|
||||
assert param.kind is not inspect.Parameter.VAR_KEYWORD
|
||||
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
|
||||
|
||||
|
||||
def _open_table_and_function(
|
||||
*,
|
||||
describe_body: dict[str, Any] | None = None,
|
||||
on_create: Callable[[dict[str, Any], http.server.BaseHTTPRequestHandler], None]
|
||||
| None = None,
|
||||
job_id: str = _JOB_ID_SYNC,
|
||||
support_branch_create: bool = False,
|
||||
function_wire: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open remote table + immutable Function; return (db, table, function, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding_describe = describe_body or _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": _text_schema_fields(),
|
||||
}
|
||||
state = {"opened": False}
|
||||
wire = function_wire or _sample_function_wire()
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": wire})
|
||||
return
|
||||
|
||||
if request.path == _JOB_DESCRIBE_PATH:
|
||||
assert body["job_id"] == job_id
|
||||
_json_response(request, _create_gen_column_done_body(job_id))
|
||||
return
|
||||
|
||||
if support_branch_create and request.path == _BRANCHES_CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {})
|
||||
return
|
||||
|
||||
if support_branch_create and request.path == _BRANCHES_LIST_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(
|
||||
request,
|
||||
{
|
||||
"branches": {
|
||||
_BRANCH_NAME: {
|
||||
"parentBranch": None,
|
||||
"parentVersion": 1,
|
||||
"createAt": 1,
|
||||
"manifestSize": 1,
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
# First describe seeds open_table; later ones are binding snapshots.
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
if on_create is not None:
|
||||
on_create(body, request)
|
||||
return
|
||||
_json_response(
|
||||
request,
|
||||
{
|
||||
"job_id": job_id,
|
||||
"server_extra": {"marker": _CREATE_RESPONSE_MARKER},
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if request.path.startswith(f"/v1/table/{_TABLE_NAME}/"):
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db(handler)
|
||||
db = cm.__enter__()
|
||||
function = _lookup_function(db)
|
||||
table = db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, RemoteTable)
|
||||
# open_table consumed the seed describe; binding/create accounting starts now.
|
||||
# Setup's one explicit lookup is cleared here and must not pollute counts.
|
||||
log.start()
|
||||
return db, table, function, log, cm
|
||||
|
||||
|
||||
def _assert_exact_create_envelope(
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
source_table_version: int,
|
||||
column_name: str,
|
||||
field_id: int,
|
||||
branch: str | None = None,
|
||||
) -> None:
|
||||
expected_keys = {"source_table_version", "spec"}
|
||||
if branch is not None:
|
||||
expected_keys.add("branch")
|
||||
assert set(body) == expected_keys
|
||||
assert body["source_table_version"] == source_table_version
|
||||
assert "table_ref" not in body
|
||||
if branch is None:
|
||||
assert "branch" not in body
|
||||
else:
|
||||
assert body["branch"] == branch
|
||||
|
||||
spec = body["spec"]
|
||||
assert set(spec) == {"format_version", "column_name", "function_call"}
|
||||
assert spec["format_version"] == 1
|
||||
assert spec["column_name"] == column_name
|
||||
for forbidden in (
|
||||
"table_ref",
|
||||
"source_table_version",
|
||||
"version",
|
||||
"output",
|
||||
"output_type",
|
||||
"output_field_id",
|
||||
"dependency_epoch",
|
||||
"materialized_epoch",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"name",
|
||||
"handle",
|
||||
"artifact",
|
||||
"geneva",
|
||||
):
|
||||
assert forbidden not in spec
|
||||
|
||||
call = spec["function_call"]
|
||||
assert set(call) == {"function_id", "arguments"}
|
||||
assert call["function_id"] == _FUNCTION_ID
|
||||
assert len(call["arguments"]) == 1
|
||||
binding = call["arguments"][0]
|
||||
assert binding["parameter"] == "text"
|
||||
value = binding["value"]
|
||||
assert value["kind"] == "field"
|
||||
assert value["field_id"] == field_id
|
||||
assert value["data_type_ipc"] == _UTF8_TYPE_IPC_B64
|
||||
assert "name" not in value
|
||||
assert "column_name" not in value
|
||||
assert "text" not in value
|
||||
# Serialized call must not late-bind by column name anywhere relevant.
|
||||
dumped = json.dumps(call)
|
||||
assert '"column_name"' not in dumped
|
||||
assert "normalized_text" not in dumped
|
||||
|
||||
|
||||
def test_public_and_native_add_generated_column_seams_must_exist():
|
||||
"""Public sync/async methods and the private native bridge must exist."""
|
||||
assert hasattr(_native.Table, "_add_generated_column"), (
|
||||
"native private bridge Table._add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(AsyncTable, "add_generated_column"), (
|
||||
"AsyncTable.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(Table, "add_generated_column"), (
|
||||
"Table.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(LanceTable, "add_generated_column"), (
|
||||
"LanceTable.add_generated_column is missing"
|
||||
)
|
||||
assert hasattr(RemoteTable, "add_generated_column"), (
|
||||
"RemoteTable.add_generated_column is missing"
|
||||
)
|
||||
|
||||
# Once present, freeze the exact public positional surface.
|
||||
_assert_exact_public_signature(Table.add_generated_column)
|
||||
_assert_exact_public_signature(LanceTable.add_generated_column)
|
||||
_assert_exact_public_signature(RemoteTable.add_generated_column)
|
||||
_assert_exact_public_signature(AsyncTable.add_generated_column)
|
||||
|
||||
|
||||
def test_sync_remote_add_generated_column_returns_job_without_eager_wrapper_mutation():
|
||||
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
|
||||
try:
|
||||
# Capture public wrapper state before the operation window.
|
||||
schema_before = table.schema
|
||||
version_before = table.version
|
||||
log.start()
|
||||
|
||||
call = normalize(text=col("text"))
|
||||
# Exact public argument order from the frozen user example.
|
||||
job = table.add_generated_column(
|
||||
"normalized_text",
|
||||
call,
|
||||
)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
|
||||
# Exact success path stops after submit: one binding describe, one create,
|
||||
# and no catalog re-lookup. Do not wait yet.
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
|
||||
# Public schema/version through the existing wrapper must still reflect
|
||||
# the pre-submit table: generated column is not published by Job accept.
|
||||
# Access both before wait so eager wrapper cache invalidation / refresh /
|
||||
# version advancement is observable.
|
||||
schema_after = table.schema
|
||||
assert "normalized_text" not in schema_after.names
|
||||
assert schema_after == schema_before
|
||||
# Schema must be served from the existing wrapper cache — no extra
|
||||
# describe beyond the one binding snapshot.
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
|
||||
version_after = table.version
|
||||
assert version_after == version_before
|
||||
# Public Remote ``version`` always describes once by design; that probe
|
||||
# must not drag a schema-cache miss, create, or catalog lookup with it.
|
||||
assert len(log.describe) == 2
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
|
||||
waited = job.wait()
|
||||
assert waited is None
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_add_generated_column_returns_async_job_and_wait_none():
|
||||
log = _RequestLog()
|
||||
state = {"opened": False}
|
||||
binding_describe = _describe_body()
|
||||
open_describe = {"version": 1, "schema": _text_schema_fields()}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
return
|
||||
if request.path == _JOB_DESCRIBE_PATH:
|
||||
assert body["job_id"] == _JOB_ID_ASYNC
|
||||
_json_response(request, _create_gen_column_done_body(_JOB_ID_ASYNC))
|
||||
return
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"job_id": _JOB_ID_ASYNC})
|
||||
return
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
normalize = await db.functions.get(_CATALOG_NAME)
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
log.start()
|
||||
call = normalize(text=col("text"))
|
||||
job = await table.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.AsyncJob
|
||||
assert job.id == _JOB_ID_ASYNC
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
waited = await job.wait()
|
||||
assert waited is None
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.create) == 1
|
||||
|
||||
|
||||
def test_remote_add_generated_column_one_describe_one_create_exact_envelope():
|
||||
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
job = table.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_SYNC
|
||||
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.other_table == []
|
||||
_assert_exact_create_envelope(
|
||||
log.create[0],
|
||||
source_table_version=_SOURCE_TABLE_VERSION,
|
||||
column_name="normalized_text",
|
||||
field_id=_TEXT_FIELD_ID,
|
||||
)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_remote_branch_add_generated_column_includes_exact_branch_identity():
|
||||
branch_describe = _describe_body(
|
||||
version=_BRANCH_SOURCE_VERSION,
|
||||
field_ids=[_BRANCH_TEXT_FIELD_ID],
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function(
|
||||
describe_body=branch_describe,
|
||||
job_id=_JOB_ID_BRANCH,
|
||||
support_branch_create=True,
|
||||
)
|
||||
try:
|
||||
branched = table.branches.create(_BRANCH_NAME)
|
||||
assert isinstance(branched, RemoteTable)
|
||||
assert branched.current_branch() == _BRANCH_NAME
|
||||
log.start()
|
||||
|
||||
call = normalize(text=col("text"))
|
||||
job = branched.add_generated_column("normalized_text", call)
|
||||
assert type(job) is lancedb.job.Job
|
||||
assert job.id == _JOB_ID_BRANCH
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert len(log.create) == 1
|
||||
assert log.describe[0].get("branch") == _BRANCH_NAME
|
||||
_assert_exact_create_envelope(
|
||||
log.create[0],
|
||||
source_table_version=_BRANCH_SOURCE_VERSION,
|
||||
column_name="normalized_text",
|
||||
field_id=_BRANCH_TEXT_FIELD_ID,
|
||||
branch=_BRANCH_NAME,
|
||||
)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_empty_column_name_fails_locally_with_zero_table_requests():
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
# Authored call owns a real literal so payload-free failure is not vacuous.
|
||||
call = normalize(text=_LITERAL_PAYLOAD_SENTINEL)
|
||||
with pytest.raises((ValueError, TypeError)) as raised:
|
||||
table.add_generated_column("", call)
|
||||
text = _exception_text(raised.value)
|
||||
lowered = text.lower()
|
||||
assert "column" in lowered or "empty" in lowered or "non-empty" in lowered
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("column_ref", "expected_token"),
|
||||
[
|
||||
("missing_text", "missing_text"),
|
||||
("Text", "Text"), # exact-case mismatch against schema field "text"
|
||||
],
|
||||
)
|
||||
def test_missing_or_case_mismatch_column_one_describe_zero_create(
|
||||
column_ref: str, expected_token: str
|
||||
):
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
call = normalize(text=col(column_ref))
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert expected_token in text
|
||||
assert "text" in text # parameter name from the Function signature
|
||||
assert "missing" in text.lower() or "field" in text.lower()
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_type_mismatch_one_describe_zero_create_identifies_parameter():
|
||||
db, table, normalize, log, cm = _open_table_and_function(
|
||||
describe_body=_describe_body(arrow_type="int32"),
|
||||
)
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "text" in text
|
||||
assert "type" in text.lower() or "mismatch" in text.lower()
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_literal_payload_stays_out_of_field_binding_failure():
|
||||
"""Authored call owns a real literal; later field binding fails payload-free."""
|
||||
wire = _sample_function_wire(
|
||||
parameters=[
|
||||
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
{"name": "prefix", "data_type_ipc": _UTF8_TYPE_IPC_B64},
|
||||
]
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function(function_wire=wire)
|
||||
try:
|
||||
call = normalize(text=col("missing_text"), prefix=_LITERAL_PAYLOAD_SENTINEL)
|
||||
with pytest.raises(ValueError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "missing_text" in text
|
||||
assert _LITERAL_PAYLOAD_SENTINEL not in text
|
||||
assert _DESCRIBE_BODY_MARKER not in text
|
||||
assert _CREATE_RESPONSE_MARKER not in text
|
||||
assert len(log.lookup) == 0
|
||||
assert len(log.describe) == 1
|
||||
assert log.create == []
|
||||
assert log.other_table == []
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_closed_async_table_fails_with_zero_operation_requests():
|
||||
log = _RequestLog()
|
||||
state = {"opened": False}
|
||||
binding_describe = _describe_body()
|
||||
open_describe = {"version": 1, "schema": _text_schema_fields()}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _LOOKUP_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
return
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding_describe)
|
||||
return
|
||||
if request.path == _CREATE_PATH:
|
||||
log.note(request.path, body)
|
||||
_json_response(request, {"job_id": _JOB_ID_ASYNC})
|
||||
return
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
async with _mock_remote_db_async(handler) as db:
|
||||
normalize = await db.functions.get(_CATALOG_NAME)
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
call = normalize(text=col("text"))
|
||||
# Public close only — do not mutate private implementation fields.
|
||||
table.close()
|
||||
log.start()
|
||||
try:
|
||||
await table.add_generated_column("normalized_text", call)
|
||||
except AttributeError:
|
||||
# Method missing: re-raise so the failure names the public seam.
|
||||
raise
|
||||
except Exception as exc:
|
||||
text = _exception_text(exc)
|
||||
assert "closed" in text.lower()
|
||||
else:
|
||||
pytest.fail("closed AsyncTable must fail before transport")
|
||||
_assert_no_operation_traffic(log)
|
||||
|
||||
|
||||
def test_rejects_non_authored_call_before_any_operation_request():
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
bad_values = (
|
||||
normalize, # exact Function handle itself
|
||||
{"text": "x"},
|
||||
col("text"), # direct query Expr
|
||||
object(),
|
||||
)
|
||||
for bad in bad_values:
|
||||
with pytest.raises(TypeError):
|
||||
table.add_generated_column("normalized_text", bad)
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_native_valid_call_returns_not_supported_without_mutation(tmp_path):
|
||||
# Immutable Function handle is connection-free; obtain it via remote lookup.
|
||||
def lookup_only(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.path == _LOOKUP_PATH
|
||||
_read_body(request)
|
||||
_json_response(request, {"function": _sample_function_wire()})
|
||||
|
||||
with _mock_remote_db(lookup_only) as remote_db:
|
||||
normalize = _lookup_function(remote_db)
|
||||
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(_TABLE_NAME, [{"text": "Hello"}, {"text": "World"}])
|
||||
assert isinstance(table, LanceTable)
|
||||
version_before = table.version
|
||||
schema_before = table.schema
|
||||
rows_before = table.to_arrow().to_pylist()
|
||||
call = normalize(text=col("text"))
|
||||
|
||||
with pytest.raises(NotImplementedError) as raised:
|
||||
table.add_generated_column("normalized_text", call)
|
||||
text = _exception_text(raised.value)
|
||||
assert "not supported" in text.lower() or "submit_create_generated_column" in text
|
||||
assert "add_columns" not in text.lower()
|
||||
|
||||
assert table.version == version_before
|
||||
assert table.schema == schema_before
|
||||
assert "normalized_text" not in table.schema.names
|
||||
assert table.to_arrow().to_pylist() == rows_before
|
||||
|
||||
|
||||
def test_public_surface_is_minimal_and_private_call_stays_opaque():
|
||||
for name in _FORBIDDEN_PUBLIC_NAMES:
|
||||
assert name not in getattr(lancedb, "__all__", [])
|
||||
assert not hasattr(lancedb, name)
|
||||
|
||||
assert not hasattr(lancedb, "_FunctionCall")
|
||||
authored_type = getattr(_native, "_FunctionCall", None)
|
||||
assert authored_type is not None
|
||||
with pytest.raises(TypeError):
|
||||
authored_type()
|
||||
|
||||
# When the public method exists, reject overdesign kwargs and keep the frozen
|
||||
# positional surface: (self, column_name, call).
|
||||
if hasattr(Table, "add_generated_column"):
|
||||
_assert_exact_public_signature(Table.add_generated_column)
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
assert (
|
||||
keyword not in inspect.signature(Table.add_generated_column).parameters
|
||||
)
|
||||
db, table, normalize, log, cm = _open_table_and_function()
|
||||
try:
|
||||
call = normalize(text=col("text"))
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
with pytest.raises(TypeError):
|
||||
table.add_generated_column(
|
||||
"normalized_text",
|
||||
call,
|
||||
**{keyword: object()},
|
||||
)
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
if hasattr(LanceTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(LanceTable.add_generated_column)
|
||||
if hasattr(RemoteTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(RemoteTable.add_generated_column)
|
||||
if hasattr(AsyncTable, "add_generated_column"):
|
||||
_assert_exact_public_signature(AsyncTable.add_generated_column)
|
||||
for keyword in _FORBIDDEN_METHOD_KWARGS:
|
||||
assert (
|
||||
keyword
|
||||
not in inspect.signature(AsyncTable.add_generated_column).parameters
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,672 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""Contract tests for Python ``table.generated_column_status`` (B3d2).
|
||||
|
||||
Public user shape under test:
|
||||
|
||||
status = table.generated_column_status("complete_col") # "complete" | "incomplete"
|
||||
|
||||
These tests exercise the live worktree PyO3 extension and public sync/async
|
||||
wrappers. While the public methods and hidden native bridge are absent they
|
||||
fail against that extension; once present they freeze the public contract
|
||||
below. They must not fake success paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import http.server
|
||||
import inspect
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Any, Callable, Literal, get_type_hints
|
||||
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
import lancedb.table
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb.remote.table import RemoteTable
|
||||
from lancedb.table import AsyncTable, LanceTable, Table
|
||||
|
||||
_TABLE_NAME = "articles"
|
||||
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
|
||||
|
||||
_ORDINARY_FIELD_ID = 1
|
||||
_COMPLETE_FIELD_ID = 5
|
||||
_INCOMPLETE_FIELD_ID = 7
|
||||
_STABLE_FIELD_IDS = [_ORDINARY_FIELD_ID, _COMPLETE_FIELD_ID, _INCOMPLETE_FIELD_ID]
|
||||
|
||||
_STATUS_FUNCTION_ID = "fn.exact.status.projection"
|
||||
_METADATA_KEY = "lancedb::generated_column"
|
||||
_RAW_METADATA_MARKER = "SENSITIVE_STATUS_METADATA_MARKER_b3d2_py_9f2e"
|
||||
|
||||
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
|
||||
_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
|
||||
_EXPECTED_RETURN = Literal["complete", "incomplete"]
|
||||
|
||||
_FORBIDDEN_PUBLIC_NAMES = (
|
||||
"GeneratedColumnStatus",
|
||||
"GeneratedColumnDefinition",
|
||||
"GeneratedColumnBindingSnapshot",
|
||||
"GeneratedColumnBindingEntry",
|
||||
)
|
||||
|
||||
_FORBIDDEN_BRIDGE_KWARGS = (
|
||||
"epoch",
|
||||
"dependency_epoch",
|
||||
"materialized_epoch",
|
||||
"function_id",
|
||||
"field_id",
|
||||
"field_ids",
|
||||
"version",
|
||||
"branch",
|
||||
"wait",
|
||||
"job",
|
||||
"request",
|
||||
"backend",
|
||||
)
|
||||
|
||||
|
||||
def _definition_metadata_json(
|
||||
output_field_id: int,
|
||||
dependency_epoch: int,
|
||||
materialized_epoch: int,
|
||||
*,
|
||||
text_field_id: int = _ORDINARY_FIELD_ID,
|
||||
) -> str:
|
||||
"""Exact JSON stored under Arrow field metadata ``lancedb::generated_column``."""
|
||||
return json.dumps(
|
||||
{
|
||||
"format_version": 1,
|
||||
"output_field_id": output_field_id,
|
||||
"function_call": {
|
||||
"function_id": _STATUS_FUNCTION_ID,
|
||||
"arguments": [
|
||||
{
|
||||
"parameter": "text",
|
||||
"value": {
|
||||
"kind": "field",
|
||||
"field_id": text_field_id,
|
||||
"data_type_ipc": _UTF8_TYPE_IPC_B64,
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
"dependency_epoch": dependency_epoch,
|
||||
"materialized_epoch": materialized_epoch,
|
||||
},
|
||||
separators=(",", ":"),
|
||||
)
|
||||
|
||||
|
||||
def _field(
|
||||
name: str,
|
||||
*,
|
||||
arrow_type: str = "string",
|
||||
nullable: bool = True,
|
||||
metadata: dict[str, str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"name": name,
|
||||
"type": {"type": arrow_type},
|
||||
"nullable": nullable,
|
||||
}
|
||||
if metadata is not None:
|
||||
body["metadata"] = metadata
|
||||
return body
|
||||
|
||||
|
||||
def _status_schema_fields(
|
||||
*,
|
||||
complete_meta: str | None = None,
|
||||
incomplete_meta: str | None = None,
|
||||
bad_name: str | None = None,
|
||||
bad_meta: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
fields = [
|
||||
_field("ordinary", arrow_type="string"),
|
||||
_field(
|
||||
"complete_col",
|
||||
arrow_type="int32",
|
||||
metadata={
|
||||
_METADATA_KEY: complete_meta
|
||||
if complete_meta is not None
|
||||
else _definition_metadata_json(_COMPLETE_FIELD_ID, 3, 3)
|
||||
},
|
||||
),
|
||||
_field(
|
||||
"incomplete_col",
|
||||
arrow_type="int32",
|
||||
metadata={
|
||||
_METADATA_KEY: incomplete_meta
|
||||
if incomplete_meta is not None
|
||||
else _definition_metadata_json(_INCOMPLETE_FIELD_ID, 4, 1)
|
||||
},
|
||||
),
|
||||
]
|
||||
if bad_name is not None and bad_meta is not None:
|
||||
fields.append(
|
||||
_field(
|
||||
bad_name,
|
||||
arrow_type="int32",
|
||||
metadata={_METADATA_KEY: bad_meta},
|
||||
)
|
||||
)
|
||||
return {"fields": fields}
|
||||
|
||||
|
||||
def _describe_body(
|
||||
*,
|
||||
version: int = 11,
|
||||
field_ids: list[int] | None = _STABLE_FIELD_IDS,
|
||||
schema: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
body: dict[str, Any] = {
|
||||
"version": version,
|
||||
"schema": schema if schema is not None else _status_schema_fields(),
|
||||
}
|
||||
if field_ids is not None:
|
||||
body["field_ids"] = field_ids
|
||||
return body
|
||||
|
||||
|
||||
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
if content_len <= 0:
|
||||
return b""
|
||||
return request.rfile.read(content_len)
|
||||
|
||||
|
||||
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
|
||||
class _Handler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
handler(self)
|
||||
|
||||
def do_POST(self):
|
||||
handler(self)
|
||||
|
||||
def log_message(self, format, *args): # noqa: A003
|
||||
return
|
||||
|
||||
return _Handler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _mock_remote_db(handler) -> Iterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
|
||||
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
|
||||
port = server.server_address[1]
|
||||
thread = threading.Thread(target=server.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{port}",
|
||||
client_config={
|
||||
"retry_config": {
|
||||
"retries": 2,
|
||||
"backoff_factor": 0.0,
|
||||
"backoff_jitter": 0.0,
|
||||
},
|
||||
"timeout_config": {"connect_timeout": 1},
|
||||
},
|
||||
)
|
||||
yield db
|
||||
finally:
|
||||
server.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def _exception_text(exc: BaseException) -> str:
|
||||
parts = [str(exc), repr(exc)]
|
||||
current: BaseException | None = exc
|
||||
seen: set[int] = set()
|
||||
while current is not None and id(current) not in seen:
|
||||
seen.add(id(current))
|
||||
parts.append(f"{type(current).__name__}: {current}")
|
||||
current = current.__cause__ or current.__context__
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _json_response(
|
||||
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
|
||||
) -> None:
|
||||
payload = json.dumps(body).encode("utf-8")
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(payload)
|
||||
|
||||
|
||||
class _RequestLog:
|
||||
"""Track post-open describe and any non-describe operation traffic."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.describe: list[dict[str, Any]] = []
|
||||
self.other: list[str] = []
|
||||
self.recording = False
|
||||
|
||||
def start(self) -> None:
|
||||
self.describe.clear()
|
||||
self.other.clear()
|
||||
self.recording = True
|
||||
|
||||
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
|
||||
if not self.recording:
|
||||
return
|
||||
if path == _DESCRIBE_PATH:
|
||||
self.describe.append(body or {})
|
||||
else:
|
||||
self.other.append(path)
|
||||
|
||||
|
||||
def _assert_no_operation_traffic(log: _RequestLog) -> None:
|
||||
assert log.describe == []
|
||||
assert log.other == []
|
||||
|
||||
|
||||
def _assert_one_status_describe(log: _RequestLog) -> None:
|
||||
assert len(log.describe) == 1, f"expected one status describe, got {log.describe!r}"
|
||||
assert log.other == [], f"unexpected non-describe traffic: {log.other!r}"
|
||||
|
||||
|
||||
def _assert_exact_public_signature(method: Any) -> None:
|
||||
"""Freeze ``(self, column_name)`` with no varargs/kwargs/keyword-only escape."""
|
||||
params = list(inspect.signature(method).parameters.values())
|
||||
assert [p.name for p in params] == ["self", "column_name"]
|
||||
for param in params:
|
||||
assert param.kind in (
|
||||
inspect.Parameter.POSITIONAL_ONLY,
|
||||
inspect.Parameter.POSITIONAL_OR_KEYWORD,
|
||||
)
|
||||
assert param.default is inspect.Parameter.empty
|
||||
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
|
||||
assert param.kind is not inspect.Parameter.VAR_KEYWORD
|
||||
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
|
||||
|
||||
|
||||
def _assert_status_string(value: Any, expected: str) -> None:
|
||||
assert value == expected
|
||||
assert type(value) is str
|
||||
assert value in ("complete", "incomplete")
|
||||
|
||||
|
||||
def _open_remote_table(
|
||||
*,
|
||||
status_describe: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open sync RemoteTable; return (table, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding = status_describe if status_describe is not None else _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
|
||||
}
|
||||
state = {"opened": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding)
|
||||
return
|
||||
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db(handler)
|
||||
db = cm.__enter__()
|
||||
table = db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, RemoteTable)
|
||||
log.start()
|
||||
return table, log, cm
|
||||
|
||||
|
||||
async def _open_remote_table_async(
|
||||
*,
|
||||
status_describe: dict[str, Any] | None = None,
|
||||
):
|
||||
"""Open async table under a live mock server; return (table, log, cm)."""
|
||||
log = _RequestLog()
|
||||
binding = status_describe if status_describe is not None else _describe_body()
|
||||
open_describe = {
|
||||
"version": 1,
|
||||
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
|
||||
}
|
||||
state = {"opened": False}
|
||||
|
||||
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
|
||||
assert request.command == "POST"
|
||||
raw = _read_body(request)
|
||||
body = json.loads(raw.decode("utf-8")) if raw else {}
|
||||
|
||||
if request.path == _DESCRIBE_PATH:
|
||||
if not state["opened"]:
|
||||
state["opened"] = True
|
||||
_json_response(request, open_describe)
|
||||
return
|
||||
log.note(request.path, body)
|
||||
_json_response(request, binding)
|
||||
return
|
||||
|
||||
log.note(request.path, body)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
request.wfile.write(b"unexpected path")
|
||||
|
||||
cm = _mock_remote_db_async(handler)
|
||||
db = await cm.__aenter__()
|
||||
table = await db.open_table(_TABLE_NAME)
|
||||
assert isinstance(table, AsyncTable)
|
||||
log.start()
|
||||
return table, log, cm
|
||||
|
||||
|
||||
def test_no_public_generated_column_status_resource_exported():
|
||||
"""Baseline: no public status class/enum/resource is exported."""
|
||||
for mod in (lancedb, lancedb.table, _native):
|
||||
for name in _FORBIDDEN_PUBLIC_NAMES:
|
||||
assert not hasattr(mod, name), f"{mod.__name__}.{name} must not be public"
|
||||
|
||||
|
||||
def test_public_surface_signatures_annotations_and_hidden_bridge():
|
||||
"""Four public methods + hidden native bridge must exist with frozen shape."""
|
||||
assert hasattr(_native.Table, "_generated_column_status"), (
|
||||
"native private bridge Table._generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(Table, "generated_column_status"), (
|
||||
"Table.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(LanceTable, "generated_column_status"), (
|
||||
"LanceTable.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(RemoteTable, "generated_column_status"), (
|
||||
"RemoteTable.generated_column_status is missing"
|
||||
)
|
||||
assert hasattr(AsyncTable, "generated_column_status"), (
|
||||
"AsyncTable.generated_column_status is missing"
|
||||
)
|
||||
|
||||
bridge = _native.Table._generated_column_status
|
||||
_assert_exact_public_signature(bridge)
|
||||
for keyword in _FORBIDDEN_BRIDGE_KWARGS:
|
||||
assert keyword not in inspect.signature(bridge).parameters
|
||||
|
||||
for method in (
|
||||
Table.generated_column_status,
|
||||
LanceTable.generated_column_status,
|
||||
RemoteTable.generated_column_status,
|
||||
):
|
||||
_assert_exact_public_signature(method)
|
||||
assert not inspect.iscoroutinefunction(method)
|
||||
assert get_type_hints(method)["return"] == _EXPECTED_RETURN
|
||||
|
||||
async_method = AsyncTable.generated_column_status
|
||||
_assert_exact_public_signature(async_method)
|
||||
assert inspect.iscoroutinefunction(async_method)
|
||||
assert get_type_hints(async_method)["return"] == _EXPECTED_RETURN
|
||||
|
||||
|
||||
def test_sync_remote_complete_and_incomplete_one_describe_each():
|
||||
table, log, cm = _open_remote_table()
|
||||
try:
|
||||
complete = table.generated_column_status("complete_col")
|
||||
_assert_status_string(complete, "complete")
|
||||
_assert_one_status_describe(log)
|
||||
|
||||
log.start()
|
||||
incomplete = table.generated_column_status("incomplete_col")
|
||||
_assert_status_string(incomplete, "incomplete")
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_complete_and_incomplete_one_describe_each():
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
complete = await table.generated_column_status("complete_col")
|
||||
_assert_status_string(complete, "complete")
|
||||
_assert_one_status_describe(log)
|
||||
|
||||
log.start()
|
||||
incomplete = await table.generated_column_status("incomplete_col")
|
||||
_assert_status_string(incomplete, "incomplete")
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("column_name", "status_describe", "expected_exc"),
|
||||
[
|
||||
(
|
||||
"missing",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"Complete_Col",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"ordinary",
|
||||
_describe_body(),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
complete_meta=_definition_metadata_json(
|
||||
_COMPLETE_FIELD_ID + 1, 3, 3
|
||||
)
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"gen_bad",
|
||||
_describe_body(
|
||||
field_ids=[*_STABLE_FIELD_IDS, 9],
|
||||
schema=_status_schema_fields(
|
||||
bad_name="gen_bad",
|
||||
bad_meta=(
|
||||
'{"format_version":1,"output_field_id":9,'
|
||||
f'"function_call":{_RAW_METADATA_MARKER},'
|
||||
'"dependency_epoch":1,"materialized_epoch":1}'
|
||||
),
|
||||
),
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
complete_meta=_definition_metadata_json(
|
||||
_COMPLETE_FIELD_ID, 1, 1
|
||||
).replace('"format_version":1', '"format_version":2')
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"incomplete_col",
|
||||
_describe_body(
|
||||
schema=_status_schema_fields(
|
||||
incomplete_meta=_definition_metadata_json(
|
||||
_INCOMPLETE_FIELD_ID, 1, 2
|
||||
)
|
||||
)
|
||||
),
|
||||
ValueError,
|
||||
),
|
||||
(
|
||||
"complete_col",
|
||||
_describe_body(field_ids=None),
|
||||
NotImplementedError,
|
||||
),
|
||||
],
|
||||
ids=[
|
||||
"missing",
|
||||
"case_mismatch",
|
||||
"ordinary",
|
||||
"output_id_mismatch",
|
||||
"malformed_metadata",
|
||||
"unknown_format_version",
|
||||
"reversed_epochs",
|
||||
"old_server_missing_field_ids",
|
||||
],
|
||||
)
|
||||
def test_remote_fail_closed_matrix_one_describe(
|
||||
column_name: str,
|
||||
status_describe: dict[str, Any],
|
||||
expected_exc: type[BaseException],
|
||||
):
|
||||
table, log, cm = _open_remote_table(status_describe=status_describe)
|
||||
try:
|
||||
with pytest.raises(expected_exc) as raised:
|
||||
table.generated_column_status(column_name)
|
||||
text = _exception_text(raised.value)
|
||||
assert _RAW_METADATA_MARKER not in text
|
||||
_assert_one_status_describe(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_sync_empty_name_zero_post_open_requests():
|
||||
table, log, cm = _open_remote_table()
|
||||
try:
|
||||
with pytest.raises(ValueError):
|
||||
table.generated_column_status("")
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
cm.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_empty_name_zero_post_open_requests():
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
with pytest.raises(ValueError):
|
||||
await table.generated_column_status("")
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_closed_status_empty_validation_wins_and_nonempty_closed():
|
||||
"""Publicly closed AsyncTable: empty validates first; nonempty is closed."""
|
||||
table, log, cm = await _open_remote_table_async()
|
||||
try:
|
||||
table.close()
|
||||
|
||||
log.start()
|
||||
try:
|
||||
await table.generated_column_status("complete_col")
|
||||
except AttributeError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
text = _exception_text(exc)
|
||||
assert "closed" in text.lower()
|
||||
else:
|
||||
pytest.fail("closed AsyncTable must fail before transport")
|
||||
_assert_no_operation_traffic(log)
|
||||
|
||||
log.start()
|
||||
with pytest.raises(ValueError) as raised:
|
||||
await table.generated_column_status("")
|
||||
text = _exception_text(raised.value)
|
||||
assert "closed" not in text.lower()
|
||||
_assert_no_operation_traffic(log)
|
||||
finally:
|
||||
await cm.__aexit__(None, None, None)
|
||||
|
||||
|
||||
def test_local_sync_ordinary_column_fails_without_side_effects(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table(
|
||||
"ordinary_only",
|
||||
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
|
||||
)
|
||||
assert isinstance(table, LanceTable)
|
||||
version_before = table.version
|
||||
schema_before = table.schema
|
||||
data_before = table.to_arrow()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
table.generated_column_status("ordinary")
|
||||
|
||||
assert table.version == version_before
|
||||
assert table.schema == schema_before
|
||||
assert table.to_arrow().equals(data_before)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_async_ordinary_column_fails_without_side_effects(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
table = await db.create_table(
|
||||
"ordinary_only_async",
|
||||
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
|
||||
)
|
||||
assert isinstance(table, AsyncTable)
|
||||
version_before = await table.version()
|
||||
schema_before = await table.schema()
|
||||
data_before = await table.to_arrow()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
await table.generated_column_status("ordinary")
|
||||
|
||||
assert await table.version() == version_before
|
||||
assert await table.schema() == schema_before
|
||||
assert (await table.to_arrow()).equals(data_before)
|
||||
@@ -0,0 +1,291 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for the local @udf declaration surface."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
import types
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import Function, Job, udf
|
||||
from lancedb._udf import _get_udf_config
|
||||
|
||||
_REMOVED_AUTHORING_KNOBS = (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _decorate(fn, **overrides):
|
||||
kwargs = {
|
||||
"inputs": {"x": pa.int32()},
|
||||
"output": pa.int64(),
|
||||
"python": "3.12",
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return udf(**kwargs)(fn)
|
||||
|
||||
|
||||
def test_udf_top_level_export_and_identity_metadata_behavior():
|
||||
assert "udf" in lancedb.__all__
|
||||
assert udf is lancedb.udf
|
||||
assert isinstance(importlib.import_module("lancedb._udf"), types.ModuleType)
|
||||
assert not isinstance(lancedb.udf, types.ModuleType)
|
||||
|
||||
def add(x, y=1):
|
||||
"""Add locally."""
|
||||
return x + y
|
||||
|
||||
original = add
|
||||
decorated = _decorate(
|
||||
add,
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
)
|
||||
|
||||
assert decorated is original
|
||||
assert decorated.__name__ == "add"
|
||||
assert decorated.__doc__ == "Add locally."
|
||||
assert str(inspect.signature(decorated)) == "(x, y=1)"
|
||||
assert decorated(2) == 3
|
||||
assert decorated(2, 5) == 7
|
||||
assert decorated(x=4, y=6) == 10
|
||||
|
||||
|
||||
def test_udf_config_snapshot_order_defaults_and_immutability():
|
||||
inputs = {"z": pa.string(), "a": pa.int32()}
|
||||
packages = ["pkg-b==2", "pkg-a==1"]
|
||||
|
||||
def combine(z, a):
|
||||
return f"{z}:{a}"
|
||||
|
||||
decorated = udf(
|
||||
inputs=inputs,
|
||||
output=pa.string(),
|
||||
python="3.11",
|
||||
packages=packages,
|
||||
output_nullable=False,
|
||||
)(combine)
|
||||
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
|
||||
assert isinstance(config.inputs, tuple)
|
||||
assert config.output == pa.string()
|
||||
assert config.output_nullable is False
|
||||
assert config.python == "3.11"
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
assert isinstance(config.packages, tuple)
|
||||
|
||||
inputs["extra"] = pa.bool_()
|
||||
del inputs["z"]
|
||||
packages.append("pkg-c==3")
|
||||
packages[0] = "mutated==0"
|
||||
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
|
||||
for attr in ("inputs", "output", "output_nullable", "python", "packages"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(config, attr, None)
|
||||
|
||||
def defaults_only(x):
|
||||
return x
|
||||
|
||||
defaulted = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
)(defaults_only)
|
||||
default_config = _get_udf_config(defaulted)
|
||||
assert default_config.packages == ()
|
||||
assert default_config.output_nullable is True
|
||||
|
||||
|
||||
def test_udf_accepts_lambda_and_closure_for_local_declaration():
|
||||
ambient = "ambient-secret-value-xyz"
|
||||
|
||||
lam = udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda n: n + 1)
|
||||
assert lam(3) == 4
|
||||
assert _get_udf_config(lam).inputs == (("n", pa.int32()),)
|
||||
|
||||
def factory(offset):
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
packages=["demo==0.1"],
|
||||
)
|
||||
def closed(n):
|
||||
return n + offset + len(ambient)
|
||||
|
||||
return closed
|
||||
|
||||
closed = factory(10)
|
||||
assert closed(2) == 12 + len(ambient)
|
||||
assert _get_udf_config(closed).packages == ("demo==0.1",)
|
||||
|
||||
|
||||
def test_udf_declaration_defers_signature_and_implementation_packaging():
|
||||
"""Declaration must not validate callable signature or embed implementation."""
|
||||
|
||||
def local_add(left, right=1):
|
||||
return left + right
|
||||
|
||||
decorated = udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(local_add)
|
||||
|
||||
assert decorated is local_add
|
||||
assert str(inspect.signature(decorated)) == "(left, right=1)"
|
||||
assert decorated(2) == 3
|
||||
assert decorated(2, 5) == 7
|
||||
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("x", pa.int32()), ("y", pa.int32()))
|
||||
for attr in (
|
||||
"source",
|
||||
"module",
|
||||
"callable",
|
||||
"function",
|
||||
"implementation",
|
||||
"bundle",
|
||||
"artifact",
|
||||
"digest",
|
||||
):
|
||||
assert not hasattr(config, attr)
|
||||
|
||||
|
||||
def test_udf_lookup_double_decoration_and_non_function_target():
|
||||
def plain(x):
|
||||
return x
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
_get_udf_config(plain)
|
||||
|
||||
decorated = _decorate(plain)
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
_decorate(decorated)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(object())
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(42)
|
||||
|
||||
|
||||
def test_udf_config_validation_errors():
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
udf({"x": pa.int32()}, pa.int32(), "3.12")(target)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs=[("x", pa.int32())])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs={1: pa.int32()})
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, inputs={"": pa.int32()})
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, inputs={"x": "int32"})
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output="int64")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, python=3.12)
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, python="")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, packages="pkg==1")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, packages=["pkg==1", ""])
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
_decorate(target, packages=["pkg==1", "pkg==1"])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, packages=["pkg==1", 2])
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output_nullable=1)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, output_nullable="true")
|
||||
|
||||
|
||||
def test_udf_rejects_removed_overdesign_and_has_no_durable_side_effects():
|
||||
params = inspect.signature(udf).parameters
|
||||
for name in _REMOVED_AUTHORING_KNOBS:
|
||||
assert name not in params
|
||||
|
||||
def score(x):
|
||||
"""score body marker unique-xyz."""
|
||||
ambient = "ambient-secret-value-xyz"
|
||||
return f"{ambient}:{x}"
|
||||
|
||||
decorated = _decorate(
|
||||
score,
|
||||
packages=["score==1.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
config = _get_udf_config(decorated)
|
||||
text = repr(config).lower()
|
||||
|
||||
assert "score body marker unique-xyz" not in text
|
||||
assert "ambient-secret-value-xyz" not in text
|
||||
for token in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"functionversion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
):
|
||||
assert token not in text
|
||||
|
||||
for attr in _REMOVED_AUTHORING_KNOBS:
|
||||
assert not hasattr(config, attr)
|
||||
|
||||
assert not isinstance(decorated, Function)
|
||||
assert not isinstance(decorated, Job)
|
||||
for attr in ("id", "function_id", "job", "job_id", "registration"):
|
||||
assert not hasattr(decorated, attr)
|
||||
@@ -0,0 +1,490 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for local FunctionCapability authoring and @udf capabilities."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import Function, FunctionCapability, Job, udf
|
||||
from lancedb._udf import _get_udf_config, _package_udf
|
||||
|
||||
_SECRET_REFERENCE = "secret://team/capability-redact-token-xyz"
|
||||
_SECRET_ENV = "API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.example.com"
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"authorization",
|
||||
"authorized",
|
||||
"value",
|
||||
"plaintext",
|
||||
"plaintext_secret",
|
||||
"secret_value",
|
||||
"job",
|
||||
"job_id",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"geneva",
|
||||
)
|
||||
|
||||
|
||||
def _decorate(fn, **overrides):
|
||||
kwargs = {
|
||||
"inputs": {"x": pa.int32()},
|
||||
"output": pa.int64(),
|
||||
"python": "3.12",
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return udf(**kwargs)(fn)
|
||||
|
||||
|
||||
def _network(origin: str = _NETWORK_ORIGIN) -> FunctionCapability:
|
||||
return FunctionCapability.network(origin)
|
||||
|
||||
|
||||
def _secret(
|
||||
reference: str = _SECRET_REFERENCE,
|
||||
*,
|
||||
environment_variable: str = _SECRET_ENV,
|
||||
) -> FunctionCapability:
|
||||
return FunctionCapability.secret(
|
||||
reference,
|
||||
environment_variable=environment_variable,
|
||||
)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_without_capabilities(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
],
|
||||
)
|
||||
def packable_with_capabilities(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
def test_function_capability_export_factories_projection_equality_immutability():
|
||||
assert "FunctionCapability" in lancedb.__all__
|
||||
assert FunctionCapability is lancedb.FunctionCapability
|
||||
|
||||
network = _network()
|
||||
secret = _secret()
|
||||
|
||||
assert network.kind == "network"
|
||||
assert network.origin == _NETWORK_ORIGIN
|
||||
assert network.reference is None
|
||||
assert network.environment_variable is None
|
||||
|
||||
assert secret.kind == "secret"
|
||||
assert secret.reference == _SECRET_REFERENCE
|
||||
assert secret.environment_variable == _SECRET_ENV
|
||||
assert secret.origin is None
|
||||
|
||||
assert network == FunctionCapability.network(_NETWORK_ORIGIN)
|
||||
assert secret == FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
)
|
||||
assert network != secret
|
||||
assert network != FunctionCapability.network("https://other.example.com")
|
||||
assert secret != FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable="OTHER_TOKEN",
|
||||
)
|
||||
|
||||
public_attrs = ("kind", "origin", "reference", "environment_variable")
|
||||
internal_slots = ("_kind", "_origin", "_reference", "_environment_variable")
|
||||
immutable_attrs = public_attrs + internal_slots
|
||||
|
||||
for attr in public_attrs:
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(network, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(secret, attr, None)
|
||||
|
||||
for attr in immutable_attrs:
|
||||
# Fresh instances per attempt so a RED slot mutation cannot corrupt
|
||||
# shared fixtures used by later assertions in this test.
|
||||
fresh_network = _network("https://fresh-immutability.example.com")
|
||||
fresh_secret = _secret(
|
||||
"secret://team/fresh-immutability-token",
|
||||
environment_variable="FRESH_IMMUTABILITY_TOKEN",
|
||||
)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(fresh_network, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(fresh_secret, attr, None)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(fresh_network, attr)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(fresh_secret, attr)
|
||||
|
||||
retained_origin = "https://config-retain.example.com"
|
||||
retained_reference = "secret://team/config-retain-token"
|
||||
retained_env = "CONFIG_RETAIN_TOKEN"
|
||||
retained_network = FunctionCapability.network(retained_origin)
|
||||
retained_secret = FunctionCapability.secret(
|
||||
retained_reference,
|
||||
environment_variable=retained_env,
|
||||
)
|
||||
expected_capabilities = (
|
||||
FunctionCapability.network(retained_origin),
|
||||
FunctionCapability.secret(
|
||||
retained_reference,
|
||||
environment_variable=retained_env,
|
||||
),
|
||||
)
|
||||
|
||||
def retain_target(x):
|
||||
return x
|
||||
|
||||
retained = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
capabilities=[retained_network, retained_secret],
|
||||
)(retain_target)
|
||||
retained_config = _get_udf_config(retained)
|
||||
assert retained_config.capabilities == expected_capabilities
|
||||
|
||||
for attr in immutable_attrs:
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(retained_network, attr, "mutated")
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(retained_secret, attr, "mutated")
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(retained_network, attr)
|
||||
with pytest.raises(AttributeError):
|
||||
delattr(retained_secret, attr)
|
||||
|
||||
assert retained_config.capabilities == expected_capabilities
|
||||
assert retained_config.capabilities[0] is retained_network
|
||||
assert retained_config.capabilities[1] is retained_secret
|
||||
assert retained_config.capabilities[0].kind == "network"
|
||||
assert retained_config.capabilities[0].origin == retained_origin
|
||||
assert retained_config.capabilities[0].reference is None
|
||||
assert retained_config.capabilities[0].environment_variable is None
|
||||
assert retained_config.capabilities[1].kind == "secret"
|
||||
assert retained_config.capabilities[1].reference == retained_reference
|
||||
assert retained_config.capabilities[1].environment_variable == retained_env
|
||||
assert retained_config.capabilities[1].origin is None
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability()
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability( # type: ignore[call-arg]
|
||||
kind="network",
|
||||
origin=_NETWORK_ORIGIN,
|
||||
)
|
||||
|
||||
assert not isinstance(network, Function)
|
||||
assert not isinstance(secret, Function)
|
||||
assert not isinstance(network, Job)
|
||||
assert not isinstance(secret, Job)
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(network, attr)
|
||||
assert not hasattr(secret, attr)
|
||||
|
||||
|
||||
def test_function_capability_validation_and_secret_redaction():
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network(None) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network(123) # type: ignore[arg-type]
|
||||
with pytest.raises(ValueError):
|
||||
FunctionCapability.network("")
|
||||
|
||||
# Backend authorization owns URL/scheme policy; non-empty is enough here.
|
||||
loose = FunctionCapability.network("example.com")
|
||||
assert loose.kind == "network"
|
||||
assert loose.origin == "example.com"
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[misc]
|
||||
_SECRET_REFERENCE,
|
||||
_SECRET_ENV,
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(None, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(123, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=None) # type: ignore[arg-type]
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=1) # type: ignore[arg-type]
|
||||
|
||||
with pytest.raises(ValueError) as empty_ref:
|
||||
FunctionCapability.secret("", environment_variable=_SECRET_ENV)
|
||||
assert _SECRET_REFERENCE not in str(empty_ref.value)
|
||||
assert _SECRET_REFERENCE not in repr(empty_ref.value)
|
||||
|
||||
with pytest.raises(ValueError) as empty_env:
|
||||
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable="")
|
||||
assert _SECRET_REFERENCE not in str(empty_env.value)
|
||||
assert _SECRET_REFERENCE not in repr(empty_env.value)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
value="super-secret",
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
plaintext_secret="super-secret",
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
environment={_SECRET_ENV: "super-secret"},
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.secret( # type: ignore[call-arg]
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
headers={"Authorization": "Bearer super-secret"},
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
FunctionCapability.network( # type: ignore[call-arg]
|
||||
_NETWORK_ORIGIN,
|
||||
headers={"X-Trace": "1"},
|
||||
)
|
||||
|
||||
secret = _secret()
|
||||
assert not hasattr(secret, "value")
|
||||
assert not hasattr(secret, "plaintext")
|
||||
assert not hasattr(secret, "plaintext_secret")
|
||||
assert not hasattr(secret, "secret_value")
|
||||
|
||||
secret_text = repr(secret)
|
||||
assert "secret" in secret_text.lower()
|
||||
assert _SECRET_ENV in secret_text
|
||||
assert _SECRET_REFERENCE not in secret_text
|
||||
assert "super-secret" not in secret_text
|
||||
|
||||
network_text = repr(_network())
|
||||
assert "network" in network_text.lower()
|
||||
assert _NETWORK_ORIGIN in network_text
|
||||
|
||||
|
||||
def test_udf_capabilities_ordered_immutable_config_default_and_validation():
|
||||
params = inspect.signature(udf).parameters
|
||||
assert "capabilities" in params
|
||||
assert params["capabilities"].kind is inspect.Parameter.KEYWORD_ONLY
|
||||
assert params["capabilities"].default == ()
|
||||
|
||||
def identity_target(x):
|
||||
"""capabilities identity marker."""
|
||||
return x + 1
|
||||
|
||||
original = identity_target
|
||||
decorated = _decorate(identity_target)
|
||||
assert decorated is original
|
||||
assert decorated.__name__ == "identity_target"
|
||||
assert decorated.__doc__ == "capabilities identity marker."
|
||||
assert decorated(2) == 3
|
||||
assert _get_udf_config(decorated).capabilities == ()
|
||||
|
||||
first = _network("https://b.example.com")
|
||||
second = _network("https://a.example.com")
|
||||
third = _network("https://b.example.com")
|
||||
secret = _secret()
|
||||
capabilities = [first, second, third, secret]
|
||||
|
||||
def combine(x):
|
||||
return x
|
||||
|
||||
with_caps = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
packages=["pkg-b==2", "pkg-a==1"],
|
||||
capabilities=capabilities,
|
||||
)(combine)
|
||||
config = _get_udf_config(with_caps)
|
||||
assert config.capabilities == (first, second, third, secret)
|
||||
assert isinstance(config.capabilities, tuple)
|
||||
assert config.packages == ("pkg-b==2", "pkg-a==1")
|
||||
assert config.inputs == (("x", pa.int32()),)
|
||||
|
||||
capabilities.append(_network("https://mutated.example.com"))
|
||||
capabilities[0] = _network("https://replaced.example.com")
|
||||
assert config.capabilities == (first, second, third, secret)
|
||||
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(config, "capabilities", ())
|
||||
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, capabilities="https://api.example.com")
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
_decorate(target, capabilities=b"https://api.example.com")
|
||||
|
||||
class _BadCapability:
|
||||
def __repr__(self) -> str:
|
||||
return "unique-bad-capability-repr-xyz"
|
||||
|
||||
with pytest.raises(TypeError) as bad_item:
|
||||
_decorate(target, capabilities=[_BadCapability()])
|
||||
assert "unique-bad-capability-repr-xyz" not in str(bad_item.value)
|
||||
assert "unique-bad-capability-repr-xyz" not in repr(bad_item.value)
|
||||
|
||||
with pytest.raises(TypeError) as bad_mixed:
|
||||
_decorate(
|
||||
target,
|
||||
capabilities=[_network(), "unique-bad-capability-string-xyz"],
|
||||
)
|
||||
assert "unique-bad-capability-string-xyz" not in str(bad_mixed.value)
|
||||
assert "unique-bad-capability-string-xyz" not in repr(bad_mixed.value)
|
||||
|
||||
|
||||
def test_udf_capabilities_rejects_function_capability_subclass_before_property_access():
|
||||
marker = "unique-hostile-capability-subclass-marker-xyz"
|
||||
|
||||
class _HostileFunctionCapability(FunctionCapability):
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
hostile = object.__new__(_HostileFunctionCapability)
|
||||
assert isinstance(hostile, FunctionCapability)
|
||||
assert type(hostile) is not FunctionCapability
|
||||
|
||||
def target(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_decorate(target, capabilities=[hostile])
|
||||
assert marker not in str(exc_info.value)
|
||||
assert marker not in repr(exc_info.value)
|
||||
assert _SECRET_REFERENCE not in str(exc_info.value)
|
||||
assert _SECRET_REFERENCE not in repr(exc_info.value)
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
def test_package_udf_preserves_capabilities_and_redacts_secret_reference():
|
||||
packaged = _package_udf(packable_with_capabilities)
|
||||
config = packaged.config
|
||||
|
||||
assert packaged.config is _get_udf_config(packable_with_capabilities)
|
||||
assert config.capabilities == (
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
)
|
||||
assert config.capabilities[0].kind == "network"
|
||||
assert config.capabilities[0].origin == _NETWORK_ORIGIN
|
||||
assert config.capabilities[1].kind == "secret"
|
||||
assert config.capabilities[1].reference == _SECRET_REFERENCE
|
||||
assert config.capabilities[1].environment_variable == _SECRET_ENV
|
||||
assert config.packages == ("pkg-a==1",)
|
||||
assert config.python == "3.12"
|
||||
assert config.output_nullable is False
|
||||
|
||||
nested = (
|
||||
f"{packaged!r}\n{config!r}\n{config.capabilities!r}\n{config.capabilities[1]!r}"
|
||||
)
|
||||
assert _SECRET_REFERENCE not in nested
|
||||
assert _SECRET_ENV in repr(config.capabilities[1])
|
||||
|
||||
|
||||
def test_capabilities_are_additive_to_existing_declaration_and_packaging():
|
||||
def score(x):
|
||||
return x
|
||||
|
||||
decorated = _decorate(
|
||||
score,
|
||||
packages=["score==1.0"],
|
||||
output_nullable=True,
|
||||
)
|
||||
config = _get_udf_config(decorated)
|
||||
assert config.inputs == (("x", pa.int32()),)
|
||||
assert config.output == pa.int64()
|
||||
assert config.output_nullable is True
|
||||
assert config.python == "3.12"
|
||||
assert config.packages == ("score==1.0",)
|
||||
assert config.capabilities == ()
|
||||
assert decorated is score
|
||||
assert decorated(4) == 4
|
||||
|
||||
packaged = _package_udf(packable_without_capabilities)
|
||||
assert packaged.config is _get_udf_config(packable_without_capabilities)
|
||||
assert packaged.callable_name == "packable_without_capabilities"
|
||||
assert packaged.config.capabilities == ()
|
||||
assert packaged.config.packages == ("pkg-a==1",)
|
||||
assert packaged.config.output_nullable is False
|
||||
assert packable_without_capabilities(1) == 2
|
||||
|
||||
params = inspect.signature(udf).parameters
|
||||
for name in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
):
|
||||
assert name not in params
|
||||
@@ -0,0 +1,506 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for the private UDF -> FunctionDefinition bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb import FunctionCapability, udf
|
||||
from lancedb import _lancedb as _native
|
||||
from lancedb import _udf as _udf_mod
|
||||
|
||||
_SOURCE_MARKER = "bridge-source-marker-unique-xyz"
|
||||
_SECRET_REFERENCE = "secret://team/bridge-redact-token-xyz"
|
||||
_SECRET_ENV = "BRIDGE_API_TOKEN"
|
||||
_NETWORK_ORIGIN = "https://api.bridge-example.com"
|
||||
_NETWORK_ORIGIN_B = "https://other.bridge-example.com"
|
||||
|
||||
_FORBIDDEN_WIRE_KEYS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionId",
|
||||
"catalog",
|
||||
"catalog_name",
|
||||
"version",
|
||||
"function_version",
|
||||
"FunctionVersion",
|
||||
"lineage",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"digest",
|
||||
"artifact",
|
||||
"artifact_digest",
|
||||
"storage",
|
||||
"storage_location",
|
||||
"location",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"nullPolicy",
|
||||
"timestamp",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"worker",
|
||||
"scheduler",
|
||||
"attempt",
|
||||
"attempt_id",
|
||||
"replica",
|
||||
"placement",
|
||||
"job",
|
||||
"job_id",
|
||||
"retry_key",
|
||||
"registration",
|
||||
)
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"job",
|
||||
"job_id",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"null_handling",
|
||||
)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"text": pa.string(), "limit": pa.int32()},
|
||||
output=pa.string(),
|
||||
python="3.12",
|
||||
packages=["pkg-b==2", "pkg-a==1"],
|
||||
output_nullable=True,
|
||||
capabilities=[
|
||||
FunctionCapability.network(_NETWORK_ORIGIN),
|
||||
FunctionCapability.secret(
|
||||
_SECRET_REFERENCE,
|
||||
environment_variable=_SECRET_ENV,
|
||||
),
|
||||
FunctionCapability.network(_NETWORK_ORIGIN_B),
|
||||
],
|
||||
)
|
||||
def packable_bridge_normalize(text, limit):
|
||||
"""bridge-source-marker-unique-xyz."""
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _build_function_definition(fn: object):
|
||||
return _udf_mod._build_function_definition(fn)
|
||||
|
||||
|
||||
def _function_definition_type():
|
||||
return _native._FunctionDefinition
|
||||
|
||||
|
||||
def _new_function_definition(**kwargs):
|
||||
return _native._new_function_definition(**kwargs)
|
||||
|
||||
|
||||
def _json_bytes(definition) -> bytes:
|
||||
payload = definition._to_json()
|
||||
if isinstance(payload, bytes):
|
||||
return payload
|
||||
assert isinstance(payload, str)
|
||||
return payload.encode("utf-8")
|
||||
|
||||
|
||||
def _decode_type_ipc(encoded: str) -> pa.DataType:
|
||||
raw = base64.b64decode(encoded)
|
||||
reader = pa.ipc.open_file(io.BytesIO(raw))
|
||||
assert reader.num_record_batches == 0
|
||||
assert len(reader.schema) == 1
|
||||
return reader.schema.field(0).type
|
||||
|
||||
|
||||
def _assert_exact_object_keys(value: dict, expected: set[str], *, context: str) -> None:
|
||||
assert isinstance(value, dict), f"{context} must be an object"
|
||||
assert set(value) == expected, f"{context} keys must match exactly: {set(value)!r}"
|
||||
|
||||
|
||||
def _assert_forbidden_keys_absent(value: object, *, context: str) -> None:
|
||||
if isinstance(value, dict):
|
||||
for key in value:
|
||||
assert key not in _FORBIDDEN_WIRE_KEYS, (
|
||||
f"forbidden key {key!r} at {context}: {value!r}"
|
||||
)
|
||||
if key == "name" and context in {
|
||||
"definition",
|
||||
"signature",
|
||||
"signature.output",
|
||||
"implementation",
|
||||
}:
|
||||
raise AssertionError(
|
||||
f"catalog/function identity key `name` must be absent at {context}"
|
||||
)
|
||||
child_context = f"{context}.{key}"
|
||||
if key == "parameters" and context == "signature":
|
||||
child_context = "signature.parameters"
|
||||
_assert_forbidden_keys_absent(value[key], context=child_context)
|
||||
elif isinstance(value, list):
|
||||
for idx, item in enumerate(value):
|
||||
item_context = (
|
||||
f"signature.parameters[{idx}]"
|
||||
if context == "signature.parameters"
|
||||
else f"{context}[{idx}]"
|
||||
)
|
||||
if context == "signature.parameters":
|
||||
assert isinstance(item, dict)
|
||||
assert "name" in item
|
||||
for key in item:
|
||||
assert key not in _FORBIDDEN_WIRE_KEYS
|
||||
assert key != "catalog_name"
|
||||
_assert_forbidden_keys_absent(
|
||||
{k: v for k, v in item.items() if k != "name"},
|
||||
context=item_context,
|
||||
)
|
||||
else:
|
||||
_assert_forbidden_keys_absent(item, context=item_context)
|
||||
|
||||
|
||||
def _assert_sanitized_text(*parts: object) -> None:
|
||||
combined = "\n".join(str(part) for part in parts)
|
||||
lowered = combined.lower()
|
||||
assert _SOURCE_MARKER.lower() not in lowered
|
||||
assert _SECRET_REFERENCE.lower() not in lowered
|
||||
assert str(Path(__file__).resolve()).lower() not in lowered
|
||||
assert Path(__file__).resolve().as_posix().lower() not in lowered
|
||||
|
||||
|
||||
def _assert_clean_validation_error(exc_info) -> None:
|
||||
_assert_sanitized_text(exc_info.value, repr(exc_info.value))
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
def _valid_builder_kwargs(**overrides):
|
||||
kwargs = {
|
||||
"parameters": [("text", pa.string()), ("limit", pa.int32())],
|
||||
"output_type": pa.string(),
|
||||
"output_nullable": True,
|
||||
"module": "bridge_mod",
|
||||
"callable_name": "normalize",
|
||||
"source": (
|
||||
"def normalize(text, limit):\n"
|
||||
f" # {_SOURCE_MARKER}\n"
|
||||
" return text[:limit]\n"
|
||||
),
|
||||
"python": "3.12",
|
||||
"packages": ["pkg-b==2", "pkg-a==1"],
|
||||
"capabilities": [
|
||||
("network", _NETWORK_ORIGIN, None),
|
||||
("secret", _SECRET_REFERENCE, _SECRET_ENV),
|
||||
("network", _NETWORK_ORIGIN_B, None),
|
||||
],
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return kwargs
|
||||
|
||||
|
||||
def test_build_function_definition_private_native_immutability_and_export_surface():
|
||||
assert "_build_function_definition" not in getattr(lancedb, "__all__", [])
|
||||
assert "_FunctionDefinition" not in lancedb.__all__
|
||||
assert not hasattr(lancedb, "_FunctionDefinition")
|
||||
assert not hasattr(lancedb, "_build_function_definition")
|
||||
assert not hasattr(lancedb, "_new_function_definition")
|
||||
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
definition_type = _function_definition_type()
|
||||
assert type(definition) is definition_type
|
||||
assert definition_type.__module__ == "lancedb._lancedb"
|
||||
assert definition_type.__name__ == "_FunctionDefinition"
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
definition_type()
|
||||
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(definition, attr)
|
||||
|
||||
for attr in ("signature", "module", "source", "capabilities"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(definition, attr, None)
|
||||
|
||||
|
||||
def test_build_function_definition_json_wire_ordered_contract_without_identity():
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
encoded_a = _json_bytes(definition)
|
||||
encoded_b = _json_bytes(definition)
|
||||
assert encoded_a == encoded_b
|
||||
|
||||
wire = json.loads(encoded_a.decode("utf-8"))
|
||||
_assert_exact_object_keys(
|
||||
wire,
|
||||
{"format_version", "signature", "implementation", "capabilities"},
|
||||
context="definition",
|
||||
)
|
||||
assert wire["format_version"] == 1
|
||||
_assert_forbidden_keys_absent(wire, context="definition")
|
||||
|
||||
signature = wire["signature"]
|
||||
_assert_exact_object_keys(signature, {"parameters", "output"}, context="signature")
|
||||
parameters = signature["parameters"]
|
||||
assert [parameter["name"] for parameter in parameters] == ["text", "limit"]
|
||||
for parameter in parameters:
|
||||
_assert_exact_object_keys(
|
||||
parameter, {"name", "data_type_ipc"}, context="parameter"
|
||||
)
|
||||
assert isinstance(parameter["data_type_ipc"], str)
|
||||
assert parameter["data_type_ipc"]
|
||||
assert _decode_type_ipc(parameters[0]["data_type_ipc"]) == pa.string()
|
||||
assert _decode_type_ipc(parameters[1]["data_type_ipc"]) == pa.int32()
|
||||
|
||||
output = signature["output"]
|
||||
_assert_exact_object_keys(
|
||||
output, {"data_type_ipc", "nullable"}, context="signature.output"
|
||||
)
|
||||
assert output["nullable"] is True
|
||||
assert _decode_type_ipc(output["data_type_ipc"]) == pa.string()
|
||||
|
||||
implementation = wire["implementation"]
|
||||
_assert_exact_object_keys(
|
||||
implementation,
|
||||
{"kind", "module", "callable", "source", "python", "packages"},
|
||||
context="implementation",
|
||||
)
|
||||
assert implementation["kind"] == "python"
|
||||
assert implementation["module"] == __name__
|
||||
assert implementation["callable"] == "packable_bridge_normalize"
|
||||
assert implementation["source"] == Path(__file__).read_text(encoding="utf-8")
|
||||
assert _SOURCE_MARKER in implementation["source"]
|
||||
assert implementation["python"] == "3.12"
|
||||
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
|
||||
|
||||
capabilities = wire["capabilities"]
|
||||
assert capabilities == [
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN},
|
||||
{
|
||||
"kind": "secret",
|
||||
"reference": _SECRET_REFERENCE,
|
||||
"environment_variable": _SECRET_ENV,
|
||||
},
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
|
||||
]
|
||||
for capability in capabilities:
|
||||
assert "value" not in capability
|
||||
assert "plaintext" not in capability
|
||||
assert "plaintext_secret" not in capability
|
||||
assert "secret_value" not in capability
|
||||
|
||||
|
||||
def test_native_definition_repr_includes_safe_structure_and_redacts_sensitive_text():
|
||||
definition = _build_function_definition(packable_bridge_normalize)
|
||||
rendered = repr(definition)
|
||||
assert "_FunctionDefinition" in rendered or "FunctionDefinition" in rendered
|
||||
assert __name__ in rendered
|
||||
assert "packable_bridge_normalize" in rendered
|
||||
assert "3.12" in rendered
|
||||
_assert_sanitized_text(rendered)
|
||||
|
||||
|
||||
def test_new_function_definition_builder_preserves_normalized_wire():
|
||||
definition = _new_function_definition(**_valid_builder_kwargs())
|
||||
assert type(definition) is _function_definition_type()
|
||||
|
||||
encoded_a = _json_bytes(definition)
|
||||
encoded_b = _json_bytes(definition)
|
||||
assert encoded_a == encoded_b
|
||||
|
||||
wire = json.loads(encoded_a.decode("utf-8"))
|
||||
assert wire["format_version"] == 1
|
||||
assert [parameter["name"] for parameter in wire["signature"]["parameters"]] == [
|
||||
"text",
|
||||
"limit",
|
||||
]
|
||||
assert _decode_type_ipc(wire["signature"]["parameters"][0]["data_type_ipc"]) == (
|
||||
pa.string()
|
||||
)
|
||||
assert _decode_type_ipc(wire["signature"]["parameters"][1]["data_type_ipc"]) == (
|
||||
pa.int32()
|
||||
)
|
||||
assert wire["signature"]["output"]["nullable"] is True
|
||||
assert _decode_type_ipc(wire["signature"]["output"]["data_type_ipc"]) == pa.string()
|
||||
|
||||
implementation = wire["implementation"]
|
||||
assert implementation["kind"] == "python"
|
||||
assert implementation["module"] == "bridge_mod"
|
||||
assert implementation["callable"] == "normalize"
|
||||
assert implementation["source"] == _valid_builder_kwargs()["source"]
|
||||
assert _SOURCE_MARKER in implementation["source"]
|
||||
assert implementation["python"] == "3.12"
|
||||
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
|
||||
assert wire["capabilities"] == [
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN},
|
||||
{
|
||||
"kind": "secret",
|
||||
"reference": _SECRET_REFERENCE,
|
||||
"environment_variable": _SECRET_ENV,
|
||||
},
|
||||
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
|
||||
]
|
||||
_assert_forbidden_keys_absent(wire, context="definition")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": [("text", pa.string()), ("text", pa.int32())]},),
|
||||
({"parameters": [("", pa.string())]},),
|
||||
({"module": ""},),
|
||||
({"callable_name": ""},),
|
||||
({"source": ""},),
|
||||
({"python": ""},),
|
||||
({"packages": ["pkg-a==1", ""]},),
|
||||
({"packages": ["pkg-a==1", "pkg-a==1"]},),
|
||||
({"capabilities": [("filesystem", _NETWORK_ORIGIN, None)]},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN, _SECRET_ENV)]},),
|
||||
({"capabilities": [("secret", _SECRET_REFERENCE, None)]},),
|
||||
({"capabilities": [("secret", _SECRET_REFERENCE, "")]},),
|
||||
({"capabilities": [("network", "", None)]},),
|
||||
({"capabilities": [("secret", "", _SECRET_ENV)]},),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_strict_validation_rejections(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
def test_new_function_definition_validation_does_not_echo_secret_or_source_marker():
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(**_valid_builder_kwargs(module=""))
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(
|
||||
**_valid_builder_kwargs(packages=["pkg-a==1", "pkg-a==1"])
|
||||
)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_new_function_definition(
|
||||
**_valid_builder_kwargs(
|
||||
capabilities=[("secret", _SECRET_REFERENCE, None)],
|
||||
)
|
||||
)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": [("text", "not-a-datatype")]},),
|
||||
({"parameters": [(123, pa.string())]},),
|
||||
({"output_type": "not-a-datatype"},),
|
||||
({"output_type": None},),
|
||||
({"output_nullable": "yes"},),
|
||||
({"packages": "pkg-a==1"},),
|
||||
({"capabilities": "network"},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN)]},),
|
||||
({"capabilities": [("network", _NETWORK_ORIGIN, None, "extra")]},),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_wrong_pyarrow_and_shape_values_fail_closed(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
class _HostileRaisingIterable:
|
||||
def __iter__(self):
|
||||
raise RuntimeError(f"{_SECRET_REFERENCE} {_SOURCE_MARKER}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides",),
|
||||
[
|
||||
({"parameters": _HostileRaisingIterable()},),
|
||||
({"packages": _HostileRaisingIterable()},),
|
||||
({"capabilities": _HostileRaisingIterable()},),
|
||||
(
|
||||
{
|
||||
"capabilities": [
|
||||
("network", _NETWORK_ORIGIN, None),
|
||||
_HostileRaisingIterable(),
|
||||
("network", _NETWORK_ORIGIN_B, None),
|
||||
]
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_new_function_definition_hostile_iterable_iter_raises_fail_closed(overrides):
|
||||
kwargs = _valid_builder_kwargs(**overrides)
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_new_function_definition(**kwargs)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_bridge_capability_exact_type(x):
|
||||
return x + 1
|
||||
|
||||
|
||||
def test_build_function_definition_rejects_forged_function_capability_subclass():
|
||||
marker = f"{_SECRET_REFERENCE} {_SOURCE_MARKER}"
|
||||
|
||||
class _HostileFunctionCapability(FunctionCapability):
|
||||
@property
|
||||
def kind(self) -> str:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def origin(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def reference(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
@property
|
||||
def environment_variable(self) -> str | None:
|
||||
raise RuntimeError(marker)
|
||||
|
||||
hostile = object.__new__(_HostileFunctionCapability)
|
||||
assert isinstance(hostile, FunctionCapability)
|
||||
assert type(hostile) is not FunctionCapability
|
||||
|
||||
config_attr = _udf_mod._CONFIG_ATTR
|
||||
original = getattr(packable_bridge_capability_exact_type, config_attr)
|
||||
forged = _udf_mod._UdfConfig(
|
||||
inputs=original.inputs,
|
||||
output=original.output,
|
||||
output_nullable=original.output_nullable,
|
||||
python=original.python,
|
||||
packages=original.packages,
|
||||
capabilities=(hostile,),
|
||||
)
|
||||
setattr(packable_bridge_capability_exact_type, config_attr, forged)
|
||||
try:
|
||||
with pytest.raises((TypeError, ValueError)) as exc_info:
|
||||
_build_function_definition(packable_bridge_capability_exact_type)
|
||||
_assert_clean_validation_error(exc_info)
|
||||
assert marker not in str(exc_info.value)
|
||||
assert marker not in repr(exc_info.value)
|
||||
finally:
|
||||
setattr(packable_bridge_capability_exact_type, config_attr, original)
|
||||
@@ -0,0 +1,486 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
"""RED contract tests for private UDF packaging validation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
import json
|
||||
import sys
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
from lancedb import Function, Job, udf
|
||||
from lancedb._udf import _get_udf_config, _package_udf
|
||||
|
||||
_BODY_MARKER = "packaging body marker unique-xyz"
|
||||
_AMBIENT_SECRET = "ambient-secret-value-xyz"
|
||||
_BUILTIN_SHADOW_SECRET = "builtin-shadow-secret-xyz"
|
||||
_SOURCE_MISMATCH_SECRET = "source-mismatch-secret-xyz"
|
||||
_INVALID_UTF8_SECRET = "invalid-utf8-secret-xyz"
|
||||
|
||||
_OVERDESIGN_ATTRS = (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_handling",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"error_policy",
|
||||
"FunctionVersion",
|
||||
"function_version",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
"id",
|
||||
"function_id",
|
||||
"job",
|
||||
"job_id",
|
||||
"registration",
|
||||
"catalog",
|
||||
"retry_key",
|
||||
"source_path",
|
||||
"path",
|
||||
"function",
|
||||
)
|
||||
|
||||
_PACKAGING_CONSTANT = 41
|
||||
|
||||
|
||||
def _packaging_helper(value: int) -> int:
|
||||
return value + _PACKAGING_CONSTANT
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int64(),
|
||||
python="3.12",
|
||||
packages=["pkg-a==1"],
|
||||
output_nullable=False,
|
||||
)
|
||||
def packable_add(x):
|
||||
"""packaging body marker unique-xyz."""
|
||||
return _packaging_helper(x) + len(json.dumps({"k": 1}))
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def packable_kwonly(x, *, y=2):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def packable_rebind_target(x):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def uses_injected_ambient(x):
|
||||
return x + len(INJECTED_AMBIENT_GLOBAL) # noqa: F821
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def uses_shadowed_builtin_len(x):
|
||||
return x + len((1, 2, 3))
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def mismatch_names(left, right):
|
||||
return left + right
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"y": pa.int32(), "x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def mismatch_order(x, y):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32(), "y": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def positional_only(x, /, y):
|
||||
return x + y
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def varargs_fn(x, *args):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def kwargs_fn(x, **kwargs):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
async def async_fn(x):
|
||||
return x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
async def async_gen_fn(x):
|
||||
yield x
|
||||
|
||||
|
||||
@udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def generator_fn(x):
|
||||
yield x
|
||||
|
||||
|
||||
def _assert_sanitized_text(*parts: object, secret: str = _AMBIENT_SECRET) -> None:
|
||||
combined = "\n".join(str(part) for part in parts)
|
||||
lowered = combined.lower()
|
||||
assert _BODY_MARKER.lower() not in lowered
|
||||
assert secret.lower() not in lowered
|
||||
assert str(Path(__file__).resolve()).lower() not in lowered
|
||||
assert Path(__file__).resolve().as_posix().lower() not in lowered
|
||||
|
||||
|
||||
def _assert_packaging_rejection(exc_info, *, secret: str = _AMBIENT_SECRET) -> None:
|
||||
_assert_sanitized_text(exc_info.value, repr(exc_info.value), secret=secret)
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _temporary_imported_module(
|
||||
directory: Path, module_name: str, source: str
|
||||
) -> Iterator[tuple[Path, object]]:
|
||||
path = directory / f"{module_name}.py"
|
||||
path.write_text(source, encoding="utf-8")
|
||||
inserted = str(directory)
|
||||
sys.path.insert(0, inserted)
|
||||
try:
|
||||
sys.modules.pop(module_name, None)
|
||||
module = importlib.import_module(module_name)
|
||||
yield path, module
|
||||
finally:
|
||||
sys.modules.pop(module_name, None)
|
||||
try:
|
||||
sys.path.remove(inserted)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
|
||||
def _temp_udf_module_source(*, body: str, secret: str | None = None) -> str:
|
||||
secret_line = f"_SECRET = {secret!r}\n" if secret is not None else ""
|
||||
return (
|
||||
"import pyarrow as pa\n"
|
||||
"from lancedb import udf\n"
|
||||
f"{secret_line}\n"
|
||||
"@udf(\n"
|
||||
' inputs={"x": pa.int32()},\n'
|
||||
" output=pa.int32(),\n"
|
||||
' python="3.12",\n'
|
||||
")\n"
|
||||
"def temp_pack_target(x):\n"
|
||||
f" {body}\n"
|
||||
)
|
||||
|
||||
|
||||
def test_package_udf_success_snapshot_source_module_callable_config_and_repr():
|
||||
packaged = _package_udf(packable_add)
|
||||
source = Path(__file__).read_text(encoding="utf-8")
|
||||
|
||||
assert packaged.source == source
|
||||
assert packaged.module == __name__
|
||||
assert packaged.module != "__main__"
|
||||
assert packaged.callable_name == "packable_add"
|
||||
assert packable_add.__qualname__ == "packable_add"
|
||||
assert packaged.config is _get_udf_config(packable_add)
|
||||
assert packaged.config.inputs == (("x", pa.int32()),)
|
||||
assert packaged.config.output == pa.int64()
|
||||
assert packaged.config.output_nullable is False
|
||||
assert packaged.config.python == "3.12"
|
||||
assert packaged.config.packages == ("pkg-a==1",)
|
||||
|
||||
for attr in ("source", "module", "callable_name", "config"):
|
||||
with pytest.raises(AttributeError):
|
||||
setattr(packaged, attr, None)
|
||||
|
||||
text = repr(packaged)
|
||||
_assert_sanitized_text(text)
|
||||
assert _BODY_MARKER not in text
|
||||
|
||||
|
||||
def test_package_udf_allows_source_bound_import_constant_and_helper():
|
||||
packaged = _package_udf(packable_add)
|
||||
assert packaged.callable_name == "packable_add"
|
||||
assert "import json" in packaged.source
|
||||
assert "_PACKAGING_CONSTANT" in packaged.source
|
||||
assert "_packaging_helper" in packaged.source
|
||||
assert packable_add(1) == _packaging_helper(1) + len(json.dumps({"k": 1}))
|
||||
|
||||
|
||||
def test_package_udf_accepts_positional_or_keyword_and_keyword_only_defaults():
|
||||
packaged = _package_udf(packable_kwonly)
|
||||
assert packaged.callable_name == "packable_kwonly"
|
||||
assert packaged.config.inputs == (("x", pa.int32()), ("y", pa.int32()))
|
||||
assert str(inspect.signature(packable_kwonly)) == "(x, *, y=2)"
|
||||
assert packable_kwonly(3) == 5
|
||||
assert packable_kwonly(3, y=7) == 10
|
||||
|
||||
|
||||
def test_package_udf_rejects_lambda_and_closure():
|
||||
lam = udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda n: n + 1)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(lam)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
ambient = _AMBIENT_SECRET
|
||||
|
||||
def factory(offset):
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def closed(n):
|
||||
return n + offset + len(ambient)
|
||||
|
||||
return closed
|
||||
|
||||
closed = factory(10)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(closed)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
def outer():
|
||||
total = 0
|
||||
|
||||
@udf(
|
||||
inputs={"n": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)
|
||||
def nested(n):
|
||||
nonlocal total
|
||||
total += n
|
||||
return total
|
||||
|
||||
return nested
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(outer())
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_signature_mismatches_and_unsupported_parameter_kinds():
|
||||
for target in (
|
||||
mismatch_names,
|
||||
mismatch_order,
|
||||
positional_only,
|
||||
varargs_fn,
|
||||
kwargs_fn,
|
||||
):
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_async_and_generator_functions():
|
||||
for target in (async_fn, async_gen_fn, generator_fn):
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_dynamic_exec_source():
|
||||
namespace: dict[str, object] = {}
|
||||
exec(
|
||||
"def dynamic_pack_target(x):\n return x + 1\n",
|
||||
namespace,
|
||||
)
|
||||
dynamic = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(namespace["dynamic_pack_target"])
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(dynamic)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_undecorated_and_wrong_input_types():
|
||||
def plain(x):
|
||||
return x
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(plain)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(object())
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
with pytest.raises(TypeError) as exc_info:
|
||||
_package_udf(42)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
|
||||
|
||||
def test_package_udf_rejects_rebound_module_attribute():
|
||||
module = sys.modules[__name__]
|
||||
original = module.packable_rebind_target
|
||||
replacement = udf(
|
||||
inputs={"x": pa.int32()},
|
||||
output=pa.int32(),
|
||||
python="3.12",
|
||||
)(lambda x: x)
|
||||
module.packable_rebind_target = replacement
|
||||
try:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(original)
|
||||
_assert_packaging_rejection(exc_info)
|
||||
finally:
|
||||
module.packable_rebind_target = original
|
||||
|
||||
|
||||
def test_package_udf_rejects_injected_ambient_global():
|
||||
module = sys.modules[__name__]
|
||||
secret = _AMBIENT_SECRET
|
||||
module.INJECTED_AMBIENT_GLOBAL = secret
|
||||
try:
|
||||
assert uses_injected_ambient(3) == 3 + len(secret)
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(uses_injected_ambient)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
finally:
|
||||
delattr(module, "INJECTED_AMBIENT_GLOBAL")
|
||||
|
||||
|
||||
def test_package_udf_rejects_builtin_shadow_injection():
|
||||
module = sys.modules[__name__]
|
||||
secret = _BUILTIN_SHADOW_SECRET
|
||||
assert not hasattr(module, "len")
|
||||
module.len = secret
|
||||
try:
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(uses_shadowed_builtin_len)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
finally:
|
||||
delattr(module, "len")
|
||||
|
||||
|
||||
def test_package_udf_rejects_loaded_code_source_mismatch(tmp_path: Path):
|
||||
secret = _SOURCE_MISMATCH_SECRET
|
||||
module_name = "udf_pkg_source_mismatch_mod"
|
||||
original = _temp_udf_module_source(body="return x + 1")
|
||||
replacement = _temp_udf_module_source(
|
||||
body=f"return x + 99 # {secret}",
|
||||
secret=secret,
|
||||
)
|
||||
with _temporary_imported_module(tmp_path, module_name, original) as (
|
||||
path,
|
||||
module,
|
||||
):
|
||||
target = module.temp_pack_target
|
||||
assert target(1) == 2
|
||||
path.write_text(replacement, encoding="utf-8")
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
_assert_packaging_rejection(exc_info, secret=secret)
|
||||
err_text = f"{exc_info.value}\n{exc_info.value!r}"
|
||||
assert str(path.resolve()) not in err_text
|
||||
assert path.resolve().as_posix() not in err_text
|
||||
|
||||
|
||||
def test_package_udf_rejects_invalid_utf8_after_import(tmp_path: Path):
|
||||
secret = _INVALID_UTF8_SECRET
|
||||
module_name = "udf_pkg_invalid_utf8_mod"
|
||||
original = _temp_udf_module_source(body="return x + 1")
|
||||
with _temporary_imported_module(tmp_path, module_name, original) as (
|
||||
path,
|
||||
module,
|
||||
):
|
||||
target = module.temp_pack_target
|
||||
assert target(1) == 2
|
||||
path.write_bytes(secret.encode("utf-8") + b"\xff\xfe invalid-bytes")
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
_package_udf(target)
|
||||
assert type(exc_info.value) is ValueError
|
||||
assert exc_info.value.__cause__ is None
|
||||
assert exc_info.value.__context__ is None
|
||||
err_text = f"{exc_info.value}\n{exc_info.value!r}"
|
||||
assert secret not in err_text
|
||||
assert "b'" not in err_text
|
||||
assert r"\xff" not in err_text
|
||||
assert str(path.resolve()) not in err_text
|
||||
assert path.resolve().as_posix() not in err_text
|
||||
|
||||
|
||||
def test_package_udf_snapshot_has_no_durable_overdesign_and_is_not_function_or_job():
|
||||
packaged = _package_udf(packable_add)
|
||||
assert not isinstance(packaged, Function)
|
||||
assert not isinstance(packaged, Job)
|
||||
for attr in _OVERDESIGN_ATTRS:
|
||||
assert not hasattr(packaged, attr)
|
||||
|
||||
text = repr(packaged).lower()
|
||||
for token in (
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"on_error",
|
||||
"functionversion",
|
||||
"artifact",
|
||||
"digest",
|
||||
"geneva",
|
||||
"retry_key",
|
||||
):
|
||||
assert token not in text
|
||||
_assert_sanitized_text(text)
|
||||
@@ -2306,3 +2306,228 @@ def test_remote_connection_jobs_surface():
|
||||
assert job.status() == "failed"
|
||||
with pytest.raises(JobFailedError, match="worker died"):
|
||||
job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
|
||||
# Pinned Rust-canonical schema-only type IPC (base64). PyArrow's schema-only
|
||||
# FileWriter bytes are not byte-identical to the Arrow Rust FileWriter used by
|
||||
# the strict Function decoder, so these fixtures are derived from Rust serde.
|
||||
_FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
|
||||
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
|
||||
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
|
||||
)
|
||||
_FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64 = (
|
||||
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
|
||||
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
|
||||
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
|
||||
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
|
||||
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
|
||||
"AAAAAAAAAAIAAAABBUlJPVzE="
|
||||
)
|
||||
_FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID = "fn.exact.python-job-result"
|
||||
_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT = object()
|
||||
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL = object()
|
||||
|
||||
|
||||
def _first_class_function_job_result_function_wire():
|
||||
int32_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64
|
||||
utf8_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64
|
||||
return {
|
||||
"kind": "function",
|
||||
"format_version": 1,
|
||||
"function": {
|
||||
"format_version": 1,
|
||||
"id": _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID,
|
||||
"signature": {
|
||||
"parameters": [
|
||||
{"name": "x", "data_type_ipc": int32_ipc},
|
||||
{"name": "label", "data_type_ipc": utf8_ipc},
|
||||
],
|
||||
"output": {
|
||||
"data_type_ipc": int32_ipc,
|
||||
"nullable": True,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _first_class_function_job_result_none_wire():
|
||||
return {"kind": "none", "format_version": 1}
|
||||
|
||||
|
||||
def _first_class_function_job_result_describe_body(
|
||||
job_id, job_type, result=_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT
|
||||
):
|
||||
body = {
|
||||
"job_id": job_id,
|
||||
"job_state": "DONE",
|
||||
"job_type": job_type,
|
||||
"creation_ms": 1,
|
||||
"spec": {},
|
||||
}
|
||||
if result is _FIRST_CLASS_FUNCTION_JOB_RESULT_NULL:
|
||||
body["result"] = None
|
||||
elif result is not _FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT:
|
||||
body["result"] = result
|
||||
return body
|
||||
|
||||
|
||||
def _first_class_function_job_result_describe_handler(bodies_by_job_id):
|
||||
def handler(request):
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
body = request.rfile.read(content_len) if content_len > 0 else b""
|
||||
payload = json.loads(body) if body else {}
|
||||
if request.path != "/v1/jobs/describe":
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
return
|
||||
job_id = payload["job_id"]
|
||||
if job_id not in bodies_by_job_id:
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
return
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(bodies_by_job_id[job_id]).encode())
|
||||
|
||||
return handler
|
||||
|
||||
|
||||
def _assert_exact_first_class_function_job_result(function):
|
||||
assert isinstance(function, lancedb.Function)
|
||||
assert function is not None
|
||||
assert not isinstance(function, dict)
|
||||
assert function.id == _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID
|
||||
assert function.parameters == (("x", pa.int32()), ("label", pa.utf8()))
|
||||
assert function.output_type == pa.int32()
|
||||
assert function.output_nullable is True
|
||||
text = repr(function)
|
||||
assert "Function" in text
|
||||
assert _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID in text
|
||||
for token in ("definition", "source", "packages", "artifact", "digest", "secret"):
|
||||
assert token not in text.lower()
|
||||
|
||||
|
||||
def test_first_class_function_job_result_sync_wait_returns_exact_function():
|
||||
bodies = {
|
||||
"job-register": _first_class_function_job_result_describe_body(
|
||||
"job-register",
|
||||
"register_function",
|
||||
_first_class_function_job_result_function_wire(),
|
||||
)
|
||||
}
|
||||
with mock_lancedb_connection(
|
||||
_first_class_function_job_result_describe_handler(bodies)
|
||||
) as db:
|
||||
result = db.job("job-register").wait()
|
||||
_assert_exact_first_class_function_job_result(result)
|
||||
|
||||
timed_out = db.job("job-register").wait(timeout=timedelta(seconds=5))
|
||||
_assert_exact_first_class_function_job_result(timed_out)
|
||||
|
||||
with pytest.raises(TypeError):
|
||||
lancedb.Function()
|
||||
with pytest.raises(AttributeError):
|
||||
result.id = "mutated"
|
||||
with pytest.raises(AttributeError):
|
||||
result.parameters = ()
|
||||
with pytest.raises(AttributeError):
|
||||
result.output_type = pa.int64()
|
||||
with pytest.raises(AttributeError):
|
||||
result.output_nullable = False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_first_class_function_job_result_async_wait_returns_exact_function():
|
||||
bodies = {
|
||||
"job-register": _first_class_function_job_result_describe_body(
|
||||
"job-register",
|
||||
"register_function",
|
||||
_first_class_function_job_result_function_wire(),
|
||||
)
|
||||
}
|
||||
async with mock_lancedb_connection_async(
|
||||
_first_class_function_job_result_describe_handler(bodies)
|
||||
) as db:
|
||||
result = await db.job("job-register").wait()
|
||||
_assert_exact_first_class_function_job_result(result)
|
||||
|
||||
timed_out = await db.job("job-register").wait(timeout=timedelta(seconds=5))
|
||||
_assert_exact_first_class_function_job_result(timed_out)
|
||||
|
||||
|
||||
def test_first_class_function_job_result_no_result_wait_returns_none():
|
||||
bodies = {
|
||||
"job-index-absent": _first_class_function_job_result_describe_body(
|
||||
"job-index-absent", "create_index"
|
||||
),
|
||||
"job-index-explicit": _first_class_function_job_result_describe_body(
|
||||
"job-index-explicit",
|
||||
"create_index",
|
||||
_first_class_function_job_result_none_wire(),
|
||||
),
|
||||
}
|
||||
with mock_lancedb_connection(
|
||||
_first_class_function_job_result_describe_handler(bodies)
|
||||
) as db:
|
||||
assert db.job("job-index-absent").wait() is None
|
||||
assert db.job("job-index-explicit").wait(timeout=timedelta(seconds=5)) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_first_class_function_job_result_async_no_result_wait_returns_none():
|
||||
bodies = {
|
||||
"job-index-absent": _first_class_function_job_result_describe_body(
|
||||
"job-index-absent", "create_index"
|
||||
),
|
||||
"job-index-explicit": _first_class_function_job_result_describe_body(
|
||||
"job-index-explicit",
|
||||
"create_index",
|
||||
_first_class_function_job_result_none_wire(),
|
||||
),
|
||||
}
|
||||
async with mock_lancedb_connection_async(
|
||||
_first_class_function_job_result_describe_handler(bodies)
|
||||
) as db:
|
||||
assert await db.job("job-index-absent").wait() is None
|
||||
assert (
|
||||
await db.job("job-index-explicit").wait(timeout=timedelta(seconds=5))
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_first_class_function_job_result_get_job_result_projection():
|
||||
bodies = {
|
||||
"job-register": _first_class_function_job_result_describe_body(
|
||||
"job-register",
|
||||
"register_function",
|
||||
_first_class_function_job_result_function_wire(),
|
||||
),
|
||||
"job-absent": _first_class_function_job_result_describe_body(
|
||||
"job-absent", "create_index"
|
||||
),
|
||||
"job-null": _first_class_function_job_result_describe_body(
|
||||
"job-null",
|
||||
"create_index",
|
||||
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL,
|
||||
),
|
||||
"job-explicit-none": _first_class_function_job_result_describe_body(
|
||||
"job-explicit-none",
|
||||
"create_index",
|
||||
_first_class_function_job_result_none_wire(),
|
||||
),
|
||||
}
|
||||
with mock_lancedb_connection(
|
||||
_first_class_function_job_result_describe_handler(bodies)
|
||||
) as db:
|
||||
register_description = db.get_job("job-register")
|
||||
_assert_exact_first_class_function_job_result(register_description.result)
|
||||
|
||||
assert db.get_job("job-absent").result is None
|
||||
assert db.get_job("job-null").result is None
|
||||
assert db.get_job("job-explicit-none").result is None
|
||||
|
||||
@@ -23,6 +23,7 @@ use lancedb::{
|
||||
connection::NamespaceClientPushdownOperation,
|
||||
database::namespace::LanceNamespaceDatabase,
|
||||
database::{CreateTableMode, Database, ReadConsistency},
|
||||
function::{FunctionId, RegisterFunctionJobSpec},
|
||||
};
|
||||
use pyo3::{
|
||||
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
|
||||
@@ -589,6 +590,121 @@ impl Connection {
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/// Submit a first-class Function registration job.
|
||||
///
|
||||
/// Accepts the exact private [`crate::function::PyFunctionDefinition`] and
|
||||
/// builds [`RegisterFunctionJobSpec`] with `expected_current_function_id =
|
||||
/// None` (create-if-absent). Does not JSON round-trip the definition.
|
||||
pub fn _register_function<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
definition: Bound<'_, crate::function::PyFunctionDefinition>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let definition = definition.get().inner().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let spec = RegisterFunctionJobSpec::try_new(name, definition, None).infer_error()?;
|
||||
let job = inner.register_function(spec).await.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Submit a first-class Function conditional replace job.
|
||||
///
|
||||
/// Accepts the observed native [`crate::function::Function`] handle and the
|
||||
/// exact private [`crate::function::PyFunctionDefinition`], then builds
|
||||
/// [`RegisterFunctionJobSpec`] with `expected_current_function_id =
|
||||
/// Some(current.id)`. Reads only `current.inner().id().clone()`. Does not
|
||||
/// JSON round-trip the definition.
|
||||
pub fn _replace_function<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
current: Bound<'_, crate::function::Function>,
|
||||
definition: Bound<'_, crate::function::PyFunctionDefinition>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let definition = definition.get().inner().clone();
|
||||
let current_id = current.get().inner().id().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let spec = RegisterFunctionJobSpec::try_new(name, definition, Some(current_id))
|
||||
.infer_error()?;
|
||||
let job = inner.register_function(spec).await.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Look up the Function currently bound to a database-scoped name.
|
||||
///
|
||||
/// Wraps the exact Rust [`lancedb::function::Function`] once. Empty names
|
||||
/// fail as [`PyValueError`] before transport via the Rust connection.
|
||||
pub fn _lookup_function_by_name<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let function = inner.lookup_function_by_name(&name).await.infer_error()?;
|
||||
Ok(crate::function::Function::new(function))
|
||||
})
|
||||
}
|
||||
|
||||
/// Look up an immutable Function by exact opaque Function ID string.
|
||||
///
|
||||
/// Constructs [`FunctionId`] with [`FunctionId::try_new`] before dispatch so
|
||||
/// empty IDs fail as [`PyValueError`] before transport.
|
||||
pub fn _lookup_function_by_id<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
function_id: String,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let id = FunctionId::try_new(function_id).infer_error()?;
|
||||
let function = inner.lookup_function_by_id(&id).await.infer_error()?;
|
||||
Ok(crate::function::Function::new(function))
|
||||
})
|
||||
}
|
||||
|
||||
/// Conditionally remove a database-scoped Function catalog name.
|
||||
///
|
||||
/// Clones the observed native [`crate::function::Function`] once and
|
||||
/// delegates to Rust [`lancedb::Connection::remove_function_name`]. Empty
|
||||
/// names fail as [`PyValueError`] before transport via the Rust connection.
|
||||
pub fn _remove_function_name<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
name: String,
|
||||
current: Bound<'_, crate::function::Function>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let current = current.get().inner().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner
|
||||
.remove_function_name(&name, ¤t)
|
||||
.await
|
||||
.infer_error()?;
|
||||
// `()` maps to an empty Python tuple via IntoPyObject; return Option
|
||||
// so the async bridge yields exact Python None.
|
||||
Ok(None::<()>)
|
||||
})
|
||||
}
|
||||
|
||||
/// Revoke an exact immutable Function by administrator set-bit.
|
||||
///
|
||||
/// Clones the observed native [`crate::function::Function`] once and
|
||||
/// delegates to Rust [`lancedb::Connection::revoke_function`].
|
||||
pub fn _revoke_function<'py>(
|
||||
self_: PyRef<'py, Self>,
|
||||
function: Bound<'_, crate::function::Function>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let inner = self_.get_inner()?.clone();
|
||||
let function = function.get().inner().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.revoke_function(&function).await.infer_error()?;
|
||||
// `()` maps to an empty Python tuple via IntoPyObject; return Option
|
||||
// so the async bridge yields exact Python None.
|
||||
Ok(None::<()>)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
|
||||
+15
-2
@@ -102,11 +102,14 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
|
||||
err.setattr(intern!(py, "__cause__"), cause_err)?;
|
||||
Err(PyErr::from_value(err))
|
||||
}),
|
||||
LanceError::JobFailed { .. } => Python::attach(|py| {
|
||||
LanceError::JobFailed { failure, .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
.import(intern!(py, "lancedb.exceptions"))?
|
||||
.getattr(intern!(py, "JobFailedError"))?;
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
|
||||
// Structural projection only: failure.error_code.as_str().
|
||||
// Never infer a code from message, phase, retryable, or source.
|
||||
let error_code = failure.error_code.as_ref().map(|code| code.as_str());
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(), error_code))?))
|
||||
}),
|
||||
LanceError::JobCancelled { .. } => Python::attach(|py| {
|
||||
let cls = py
|
||||
@@ -114,6 +117,16 @@ impl<T> PythonErrorExt<T> for std::result::Result<T, LanceError> {
|
||||
.getattr(intern!(py, "JobCancelledError"))?;
|
||||
Err(PyErr::from_value(cls.call1((err.to_string(),))?))
|
||||
}),
|
||||
LanceError::Function { code, message } => Python::attach(|py| {
|
||||
let cls = py
|
||||
.import(intern!(py, "lancedb.exceptions"))?
|
||||
.getattr(intern!(py, "FunctionError"))?;
|
||||
// Structural projection only: code.as_str() + sanitized message.
|
||||
// Never infer a code from HTTP status or diagnostic text.
|
||||
Err(PyErr::from_value(
|
||||
cls.call1((message.as_str(), code.as_str()))?,
|
||||
))
|
||||
}),
|
||||
_ => self.runtime_error(),
|
||||
},
|
||||
}
|
||||
|
||||
+28
-1
@@ -10,7 +10,7 @@
|
||||
use std::ops::{Add, Div, Mul, Not, Sub};
|
||||
|
||||
use arrow::{datatypes::DataType, pyarrow::PyArrowType};
|
||||
use datafusion_common::ScalarValue;
|
||||
use datafusion_common::{Column, ScalarValue};
|
||||
use lancedb::expr::{
|
||||
DfExpr, col as ldb_col, contains, expr_cast, is_in, lit as df_lit, lower, upper,
|
||||
};
|
||||
@@ -27,6 +27,33 @@ use pyo3::{Bound, PyAny, PyResult, exceptions::PyValueError, prelude::*, pyfunct
|
||||
#[derive(Clone)]
|
||||
pub struct PyExpr(pub DfExpr);
|
||||
|
||||
/// Crate-private inspection result for Function call authoring (FF-028).
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) enum DirectExprView<'a> {
|
||||
/// Direct unqualified DataFusion Column; name is case-sensitive.
|
||||
UnqualifiedColumn(&'a str),
|
||||
/// Direct Literal scalar; Arrow type is owned by the scalar value.
|
||||
Literal(&'a ScalarValue),
|
||||
}
|
||||
|
||||
impl PyExpr {
|
||||
/// Inspect a direct Column/Literal node for Function call authoring.
|
||||
///
|
||||
/// Returns `None` for every other expression shape (arithmetic, cast,
|
||||
/// scalar function, predicate, alias, qualified column, etc.).
|
||||
pub(crate) fn as_direct_column_or_literal(&self) -> Option<DirectExprView<'_>> {
|
||||
match &self.0 {
|
||||
DfExpr::Column(Column {
|
||||
relation: None,
|
||||
name,
|
||||
..
|
||||
}) => Some(DirectExprView::UnqualifiedColumn(name.as_str())),
|
||||
DfExpr::Literal(value, _) => Some(DirectExprView::Literal(value)),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl PyExpr {
|
||||
// ── comparisons ──────────────────────────────────────────────────────────
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -289,7 +289,7 @@ struct IvfHnswFlatParams {
|
||||
target_partition_size: Option<u32>,
|
||||
}
|
||||
|
||||
#[pyclass(get_all)]
|
||||
#[pyclass(module = "lancedb._lancedb", get_all)]
|
||||
/// A description of an index currently configured on a column
|
||||
pub struct IndexConfig {
|
||||
/// The type of the index
|
||||
|
||||
+31
-4
@@ -3,6 +3,7 @@
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::function::Function;
|
||||
use crate::runtime::future_into_py;
|
||||
use pyo3::{Bound, PyAny, PyRef, PyResult, pyclass, pymethods};
|
||||
|
||||
@@ -21,6 +22,23 @@ impl Job {
|
||||
}
|
||||
}
|
||||
|
||||
/// Project a Rust [`lancedb::JobResult`] onto the Python success surface.
|
||||
///
|
||||
/// Delegates variant interpretation to [`lancedb::JobResult::into_function`]:
|
||||
/// no nested Function collapses to Python `None`; an exact Function becomes
|
||||
/// the corresponding [`Function`] handle.
|
||||
fn project_wait_result(result: lancedb::JobResult) -> Option<Function> {
|
||||
result.into_function().map(Function::new)
|
||||
}
|
||||
|
||||
/// Project a describe `result` onto Python `Optional[Function]`.
|
||||
///
|
||||
/// Rust `None`, `Some(JobResult::None)`, and JSON null all become Python
|
||||
/// `None`. Only `Some(JobResult::Function)` becomes a [`Function`] handle.
|
||||
fn project_description_result(result: Option<lancedb::JobResult>) -> Option<Function> {
|
||||
result.and_then(project_wait_result)
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl Job {
|
||||
#[getter]
|
||||
@@ -39,8 +57,8 @@ impl Job {
|
||||
pub fn wait(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.wait().await.infer_error()?;
|
||||
Ok(())
|
||||
let result = inner.wait().await.infer_error()?;
|
||||
Ok(project_wait_result(result))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -93,14 +111,16 @@ pub struct JobFailureInfo {
|
||||
phase: Option<String>,
|
||||
message: Option<String>,
|
||||
retryable: Option<bool>,
|
||||
/// Exact wire `error_code` string when Rust decoded one; never inferred.
|
||||
error_code: Option<String>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl JobFailureInfo {
|
||||
fn __repr__(&self) -> String {
|
||||
format!(
|
||||
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?})",
|
||||
self.phase, self.message, self.retryable
|
||||
"JobFailureInfo(phase={:?}, message={:?}, retryable={:?}, error_code={:?})",
|
||||
self.phase, self.message, self.retryable, self.error_code
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -115,6 +135,7 @@ pub struct JobDescription {
|
||||
creation_ms: i64,
|
||||
spec_json: Option<String>,
|
||||
failure: Option<JobFailureInfo>,
|
||||
result: Option<Function>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
@@ -139,7 +160,13 @@ impl From<lancedb::database::JobDescription> for JobDescription {
|
||||
phase: failure.phase,
|
||||
message: failure.message,
|
||||
retryable: failure.retryable,
|
||||
// Structural projection only: exact as_str(); never infer.
|
||||
error_code: failure
|
||||
.error_code
|
||||
.as_ref()
|
||||
.map(|code| code.as_str().to_string()),
|
||||
}),
|
||||
result: project_description_result(description.result),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -23,6 +23,7 @@ pub mod arrow;
|
||||
pub mod connection;
|
||||
pub mod error;
|
||||
pub mod expr;
|
||||
pub mod function;
|
||||
pub mod header;
|
||||
pub mod index;
|
||||
pub mod job;
|
||||
@@ -45,6 +46,9 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_class::<Connection>()?;
|
||||
m.add_class::<Session>()?;
|
||||
m.add_class::<Table>()?;
|
||||
m.add_class::<crate::function::Function>()?;
|
||||
m.add_class::<crate::function::PyFunctionDefinition>()?;
|
||||
m.add_class::<crate::function::AuthoredFunctionCall>()?;
|
||||
m.add_class::<crate::job::Job>()?;
|
||||
m.add_class::<crate::job::JobInfo>()?;
|
||||
m.add_class::<crate::job::JobDescription>()?;
|
||||
@@ -88,6 +92,10 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_function(wrap_pyfunction!(expr_col, m)?)?;
|
||||
m.add_function(wrap_pyfunction!(expr_lit, m)?)?;
|
||||
m.add_function(wrap_pyfunction!(expr_func, m)?)?;
|
||||
m.add_function(wrap_pyfunction!(
|
||||
crate::function::_new_function_definition,
|
||||
m
|
||||
)?)?;
|
||||
m.add("__version__", env!("CARGO_PKG_VERSION"))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -11,7 +11,7 @@ use pyo3::{PyResult, pyclass, pymethods};
|
||||
/// Sessions allow you to configure cache sizes for index and metadata caches,
|
||||
/// which can significantly impact memory use and performance. They can
|
||||
/// also be re-used across multiple connections to share the same cache state.
|
||||
#[pyclass(from_py_object)]
|
||||
#[pyclass(module = "lancedb._lancedb", from_py_object)]
|
||||
#[derive(Clone)]
|
||||
pub struct Session {
|
||||
pub(crate) inner: Arc<LanceSession>,
|
||||
|
||||
+142
-2
@@ -26,7 +26,7 @@ use lancedb::table::{
|
||||
use lancedb::tokenize as lancedb_tokenize;
|
||||
use pyo3::{
|
||||
Bound, FromPyObject, Py, PyAny, PyRef, PyResult, Python,
|
||||
exceptions::{PyRuntimeError, PyValueError},
|
||||
exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError},
|
||||
pyclass, pyfunction, pymethods,
|
||||
types::{IntoPyDict, PyAnyMethods, PyBytes, PyDict, PyDictMethods, PyList, PyListMethods},
|
||||
};
|
||||
@@ -579,7 +579,7 @@ impl PyBlobFile {
|
||||
}
|
||||
}
|
||||
|
||||
#[pyclass(get_all, from_py_object)]
|
||||
#[pyclass(module = "lancedb._lancedb", get_all, from_py_object)]
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct FtsToken {
|
||||
pub text: String,
|
||||
@@ -930,6 +930,146 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: bind an authored Function call once and submit create.
|
||||
///
|
||||
/// Private native path for Python ``table.add_generated_column``. Rejects an
|
||||
/// empty ``column_name`` before reading the table handle. Does not expose
|
||||
/// source version, stable field IDs, the operation spec, or request envelope.
|
||||
#[doc(hidden)]
|
||||
pub fn _add_generated_column<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
call: Bound<'_, crate::function::AuthoredFunctionCall>,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
let authored = call.get().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let (source_table_version, bound_call) =
|
||||
authored.bind_to_table(&inner).await.infer_error()?;
|
||||
let spec = lancedb::function::CreateGeneratedColumnJobSpec::try_new(
|
||||
column_name,
|
||||
authored.function(),
|
||||
bound_call,
|
||||
)
|
||||
.infer_error()?;
|
||||
let job = inner
|
||||
.submit_create_generated_column(source_table_version, spec)
|
||||
.await
|
||||
.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: project generated-column completeness for one column name.
|
||||
///
|
||||
/// Private native path for Python ``table.generated_column_status``. Rejects
|
||||
/// an empty ``column_name`` before reading the table handle. Maps only the
|
||||
/// known Rust status variants to ``"complete"`` / ``"incomplete"``.
|
||||
#[doc(hidden)]
|
||||
pub fn _generated_column_status<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let status = inner
|
||||
.generated_column_status(column_name)
|
||||
.await
|
||||
.infer_error()?;
|
||||
match status {
|
||||
lancedb::function::GeneratedColumnStatus::Complete => Ok("complete"),
|
||||
lancedb::function::GeneratedColumnStatus::Incomplete => Ok("incomplete"),
|
||||
_ => Err(PyNotImplementedError::new_err(
|
||||
"unsupported generated column status",
|
||||
)),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: load exact definition, resolve Function by ID, submit refresh.
|
||||
///
|
||||
/// Private native path for Python ``table.refresh_generated_column``. Rejects
|
||||
/// an empty ``column_name`` before reading the table handle. Does not expose
|
||||
/// source version, Function, field IDs, epochs, specs, or request envelope.
|
||||
#[doc(hidden)]
|
||||
pub fn _refresh_generated_column<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let (source_table_version, definition) = inner
|
||||
.generated_column_definition_snapshot(column_name)
|
||||
.await
|
||||
.infer_error()?;
|
||||
let function_id = definition.function_call().function_id().clone();
|
||||
let function = inner
|
||||
.resolve_function_for_generated_column(&function_id)
|
||||
.await
|
||||
.infer_error()?;
|
||||
let spec =
|
||||
lancedb::function::RefreshGeneratedColumnJobSpec::try_new(&function, definition)
|
||||
.infer_error()?;
|
||||
let job = inner
|
||||
.submit_refresh_generated_column(source_table_version, spec)
|
||||
.await
|
||||
.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
/// Hidden bridge: one binding snapshot, bind new call, submit change.
|
||||
///
|
||||
/// Private native path for Python ``table.alter_generated_column``. Rejects
|
||||
/// an empty ``column_name`` before reading the table handle. Fetches exactly
|
||||
/// one binding snapshot, loads the expected definition from that same
|
||||
/// object, binds the authored call against it, and submits change. Does not
|
||||
/// expose source version, Function handles, field IDs, epochs, specs, or
|
||||
/// request envelope.
|
||||
#[doc(hidden)]
|
||||
pub fn _alter_generated_column<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
column_name: String,
|
||||
new_call: Bound<'_, crate::function::AuthoredFunctionCall>,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
if column_name.is_empty() {
|
||||
return Err(PyValueError::new_err("column_name must be non-empty"));
|
||||
}
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
let authored = new_call.get().clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let snapshot = inner
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.infer_error()?;
|
||||
let expected_definition = snapshot
|
||||
.generated_column_definition(&column_name)
|
||||
.infer_error()?;
|
||||
let (source_table_version, bound_new_call) =
|
||||
authored.bind_against_snapshot(&snapshot).infer_error()?;
|
||||
let spec = lancedb::function::ChangeGeneratedColumnJobSpec::try_new(
|
||||
expected_definition,
|
||||
authored.function(),
|
||||
bound_new_call,
|
||||
)
|
||||
.infer_error()?;
|
||||
let job = inner
|
||||
.submit_change_generated_column(source_table_version, spec)
|
||||
.await
|
||||
.infer_error()?;
|
||||
Ok(crate::job::Job::new(job))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn drop_index(self_: PyRef<'_, Self>, index_name: String) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
|
||||
@@ -12,6 +12,7 @@ rust-version.workspace = true
|
||||
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
|
||||
[dependencies]
|
||||
ahash = { workspace = true }
|
||||
base64 = "0.22"
|
||||
arrow = { workspace = true }
|
||||
arrow-array = { workspace = true }
|
||||
arrow-buffer = { workspace = true }
|
||||
@@ -49,8 +50,6 @@ lance-namespace = { workspace = true }
|
||||
lance-namespace-impls = { workspace = true }
|
||||
metrics = { workspace = true, optional = true }
|
||||
metrics-util = { workspace = true, optional = true }
|
||||
# Pin the GooseFS SDK to the version required by Lance's OpenDAL dependency.
|
||||
goosefs-sdk = { version = "=0.1.9", optional = true }
|
||||
moka = { workspace = true }
|
||||
pin-project = { workspace = true }
|
||||
tokio = { version = "1.23", features = ["rt-multi-thread", "sync"] }
|
||||
@@ -136,7 +135,6 @@ azure = [
|
||||
]
|
||||
cos = ["lance/tencent", "lance-io/tencent"]
|
||||
goosefs = [
|
||||
"dep:goosefs-sdk",
|
||||
"lance/goosefs",
|
||||
"lance-io/goosefs",
|
||||
"lance-namespace-impls/dir-goosefs",
|
||||
|
||||
@@ -17,7 +17,7 @@ use arrow_array::builder::LargeBinaryBuilder;
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use lance::dataset::{BlobRangeRequest as LanceBlobRangeRequest, Dataset, WriteParams};
|
||||
use lance_arrow::FieldExt;
|
||||
use lance_file::version::LanceFileVersion;
|
||||
use lance_file::version::{ConcreteFileVersion, LanceFileVersion};
|
||||
use lance_io::object_store::ObjectStore;
|
||||
use object_store::path::Path;
|
||||
|
||||
@@ -333,8 +333,13 @@ pub(crate) fn ensure_blob_storage_version(schema: &Schema, params: &mut WritePar
|
||||
.data_storage_version
|
||||
.unwrap_or(LanceFileVersion::Stable)
|
||||
.resolve();
|
||||
if resolved < LanceFileVersion::V2_2 {
|
||||
params.data_storage_version = Some(LanceFileVersion::V2_2);
|
||||
// Exact formats deliberately have no Ord: capability is not implied by
|
||||
// release order. Enumerate every current concrete variant explicitly.
|
||||
match resolved {
|
||||
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1 => {
|
||||
params.data_storage_version = Some(LanceFileVersion::V2_2);
|
||||
}
|
||||
ConcreteFileVersion::V2_2 | ConcreteFileVersion::V2_3 => {}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -499,7 +504,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_2
|
||||
LanceFileVersion::V2_2.resolve()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -512,7 +517,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_2
|
||||
LanceFileVersion::V2_2.resolve()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -523,7 +528,10 @@ mod tests {
|
||||
..Default::default()
|
||||
};
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(params.data_storage_version.unwrap(), LanceFileVersion::V2_3);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_3.resolve()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -28,6 +28,7 @@ use crate::database::{
|
||||
};
|
||||
use crate::embeddings::{EmbeddingRegistry, MemoryRegistry};
|
||||
use crate::error::{Error, Result};
|
||||
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
|
||||
#[cfg(feature = "remote")]
|
||||
use crate::remote::{
|
||||
client::ClientConfig,
|
||||
@@ -550,6 +551,88 @@ impl Connection {
|
||||
self.internal.job_history(job_id).await
|
||||
}
|
||||
|
||||
/// Submit a first-class Function registration job.
|
||||
///
|
||||
/// Returns a [`crate::job::Job`] handle for the accepted server-side job.
|
||||
/// Only remote databases support registration; local databases return
|
||||
/// [`Error::NotSupported`].
|
||||
pub async fn register_function(
|
||||
&self,
|
||||
spec: RegisterFunctionJobSpec,
|
||||
) -> Result<crate::job::Job> {
|
||||
self.internal.register_function(spec).await
|
||||
}
|
||||
|
||||
/// Look up the Function currently bound to a database-scoped name.
|
||||
///
|
||||
/// The name is lookup indirection only and is never part of the returned
|
||||
/// [`Function`]. Empty names return [`Error::InvalidInput`] before backend
|
||||
/// dispatch. Only remote databases support enterprise catalog lookup;
|
||||
/// nonempty local lookups return [`Error::NotSupported`].
|
||||
pub async fn lookup_function_by_name(&self, name: impl AsRef<str>) -> Result<Function> {
|
||||
let name = name.as_ref();
|
||||
// Public nonempty invariant: validate before any Database backend sees
|
||||
// the call so local and remote Connections agree on InvalidInput.
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function lookup name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
self.internal.lookup_function_by_name(name).await
|
||||
}
|
||||
|
||||
/// Look up an immutable Function by exact opaque [`FunctionId`].
|
||||
///
|
||||
/// Exact-ID lookup is independent of later catalog name changes. Only
|
||||
/// remote databases support enterprise catalog lookup; local databases
|
||||
/// return [`Error::NotSupported`].
|
||||
pub async fn lookup_function_by_id(&self, function_id: &FunctionId) -> Result<Function> {
|
||||
self.internal.lookup_function_by_id(function_id).await
|
||||
}
|
||||
|
||||
/// Conditionally remove a database-scoped Function catalog name.
|
||||
///
|
||||
/// This is a direct synchronous catalog compare-and-swap (CAS), not a
|
||||
/// [`crate::job::Job`], not physical [`Function`] deletion, and not
|
||||
/// revocation. The caller supplies an observed immutable [`Function`]
|
||||
/// handle; only [`Function::id`] is authority for the CAS precondition.
|
||||
///
|
||||
/// Empty names return [`Error::InvalidInput`] before backend dispatch.
|
||||
/// Nonempty names on local/default backends return [`Error::NotSupported`].
|
||||
/// Remote backends complete only when the server reports durable CAS
|
||||
/// success for the `(name, current.id)` pair.
|
||||
pub async fn remove_function_name(
|
||||
&self,
|
||||
name: impl AsRef<str>,
|
||||
current: &Function,
|
||||
) -> Result<()> {
|
||||
let name = name.as_ref();
|
||||
// Public nonempty invariant: validate before any Database backend sees
|
||||
// the call so local and remote Connections agree on InvalidInput.
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function name removal name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
self.internal.remove_function_name(name, current).await
|
||||
}
|
||||
|
||||
/// Revoke an exact immutable [`Function`] by opaque id.
|
||||
///
|
||||
/// This is a direct synchronous administrator catalog set-bit, not a
|
||||
/// [`crate::job::Job`], not catalog name removal, not physical deletion,
|
||||
/// and not [`Function`] or generated-column mutation. The caller supplies
|
||||
/// an already-validated exact [`Function`] handle; only [`Function::id`]
|
||||
/// is sent on the wire.
|
||||
///
|
||||
/// Local/default backends return [`Error::NotSupported`]. Remote backends
|
||||
/// complete only when the server reports durable success for that exact
|
||||
/// id. Repeated logical calls that each receive success succeed; there is
|
||||
/// no client-side already-revoked branch.
|
||||
pub async fn revoke_function(&self, function: &Function) -> Result<()> {
|
||||
self.internal.revoke_function(function).await
|
||||
}
|
||||
|
||||
/// Drop a table in the database.
|
||||
///
|
||||
/// # Arguments
|
||||
|
||||
@@ -438,10 +438,9 @@ mod tests {
|
||||
.await
|
||||
.unwrap()
|
||||
.data_storage_format
|
||||
.lance_file_version()
|
||||
.unwrap();
|
||||
// Compare resolved versions since Stable/Next are aliases that resolve at storage time
|
||||
assert_eq!(storage_format.resolve(), data_storage_version.resolve());
|
||||
.lance_file_format();
|
||||
// Compare concrete stored format to the resolved requested alias.
|
||||
assert_eq!(storage_format, data_storage_version.resolve());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -30,12 +30,16 @@ use lance_namespace::models::{
|
||||
|
||||
use crate::data::scannable::Scannable;
|
||||
use crate::error::Result;
|
||||
use crate::function::{Function, FunctionId, RegisterFunctionJobSpec};
|
||||
use crate::table::{BaseTable, WriteOptions};
|
||||
|
||||
pub mod listing;
|
||||
pub mod namespace;
|
||||
pub(crate) mod read_freshness;
|
||||
|
||||
#[cfg(test)]
|
||||
mod create_table_generated_column_schema_admission_contract;
|
||||
|
||||
pub trait DatabaseOptions {
|
||||
fn serialize_into_map(&self, map: &mut HashMap<String, String>);
|
||||
}
|
||||
@@ -230,6 +234,12 @@ pub struct JobDescription {
|
||||
pub creation_ms: i64,
|
||||
/// The job-type-specific specification. Null when the server omits it.
|
||||
pub spec: serde_json::Value,
|
||||
/// Explicit success result from the describe envelope, when present.
|
||||
///
|
||||
/// Missing or JSON `null` wire `result` is [`None`]. An explicit
|
||||
/// [`crate::JobResult::None`] object is `Some(JobResult::None)`. An exact
|
||||
/// Function result is `Some(JobResult::Function(...))`.
|
||||
pub result: Option<crate::job::JobResult>,
|
||||
/// Why the job failed, when the job is failed and the server reports a
|
||||
/// reason.
|
||||
pub failure: Option<crate::error::JobFailure>,
|
||||
@@ -311,6 +321,64 @@ pub trait Database:
|
||||
async fn job_history(&self, _job_id: Option<&str>) -> Result<Vec<RecordBatch>> {
|
||||
job_op_not_supported("job_history")
|
||||
}
|
||||
/// Submit a first-class Function registration job.
|
||||
///
|
||||
/// Returns a [`crate::job::Job`] handle for the accepted server-side job.
|
||||
/// Local databases do not support registration.
|
||||
async fn register_function(&self, _spec: RegisterFunctionJobSpec) -> Result<crate::job::Job> {
|
||||
job_op_not_supported("register_function")
|
||||
}
|
||||
/// Look up the Function currently bound to a database-scoped name.
|
||||
///
|
||||
/// The name is lookup indirection only and is never part of the returned
|
||||
/// [`Function`]. Empty names return [`crate::Error::InvalidInput`] before
|
||||
/// the unsupported fallback so local and remote backends agree. Nonempty
|
||||
/// names on databases without enterprise catalog lookup return
|
||||
/// [`crate::Error::NotSupported`].
|
||||
async fn lookup_function_by_name(&self, name: &str) -> Result<Function> {
|
||||
// Public nonempty invariant on the Database trait seam itself:
|
||||
// Connection::database() exposes Arc<dyn Database>, so empty-name
|
||||
// rejection cannot rely solely on Connection prevalidation.
|
||||
if name.is_empty() {
|
||||
return Err(crate::error::Error::InvalidInput {
|
||||
message: "function lookup name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
job_op_not_supported("lookup_function_by_name")
|
||||
}
|
||||
/// Look up an immutable Function by exact opaque [`FunctionId`].
|
||||
///
|
||||
/// Exact-ID lookup is independent of later catalog name changes. Local
|
||||
/// databases do not support enterprise catalog lookup.
|
||||
async fn lookup_function_by_id(&self, _function_id: &FunctionId) -> Result<Function> {
|
||||
job_op_not_supported("lookup_function_by_id")
|
||||
}
|
||||
/// Conditionally remove a database-scoped Function catalog name.
|
||||
///
|
||||
/// Direct synchronous catalog CAS, not a Job and not physical Function
|
||||
/// deletion. Empty names return [`crate::Error::InvalidInput`] before the
|
||||
/// unsupported fallback so local and remote backends agree. Nonempty names
|
||||
/// on databases without enterprise catalog mutation return
|
||||
/// [`crate::Error::NotSupported`].
|
||||
async fn remove_function_name(&self, name: &str, _current: &Function) -> Result<()> {
|
||||
// Public nonempty invariant on the Database trait seam itself:
|
||||
// Connection::database() exposes Arc<dyn Database>, so empty-name
|
||||
// rejection cannot rely solely on Connection prevalidation.
|
||||
if name.is_empty() {
|
||||
return Err(crate::error::Error::InvalidInput {
|
||||
message: "function name removal name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
job_op_not_supported("remove_function_name")
|
||||
}
|
||||
/// Revoke an exact immutable [`Function`] by opaque id.
|
||||
///
|
||||
/// Direct synchronous administrator catalog set-bit, not a Job, not name
|
||||
/// removal, and not physical Function deletion. Databases without
|
||||
/// enterprise catalog mutation return [`crate::Error::NotSupported`].
|
||||
async fn revoke_function(&self, _function: &Function) -> Result<()> {
|
||||
job_op_not_supported("revoke_function")
|
||||
}
|
||||
/// Open a table in the database
|
||||
async fn open_table(&self, request: OpenTableRequest) -> Result<Arc<dyn BaseTable>>;
|
||||
/// Rename a table in the database
|
||||
|
||||
@@ -0,0 +1,788 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! RED runtime contract tests for create-table schema admission (B4g).
|
||||
//!
|
||||
//! Caller-authored Arrow field metadata under
|
||||
//! [`crate::function::GENERATED_COLUMN_METADATA_KEY`] must not enter table
|
||||
//! schema state through general-purpose `Database::create_table`. Generated
|
||||
//! definitions are Job-owned. This module proves the missing admission guard
|
||||
//! on Native listing, Native namespace, and Remote create paths.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema, SchemaRef};
|
||||
use tempfile::TempDir;
|
||||
|
||||
use crate::arrow::SendableRecordBatchStream;
|
||||
use crate::data::scannable::Scannable;
|
||||
use crate::database::listing::ListingDatabase;
|
||||
use crate::database::{CreateTableMode, CreateTableRequest, Database, TableNamesRequest};
|
||||
use crate::error::Error;
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
|
||||
};
|
||||
|
||||
const ID: &str = "id";
|
||||
const ORDINARY: &str = "ordinary";
|
||||
const GEN_OUT: &str = "gen_out";
|
||||
const ORDINARY_META_KEY: &str = "unit";
|
||||
const ORDINARY_META_VALUE: &str = "label";
|
||||
const FN_ID: &str = "fn.exact.b4g.create_table.literal";
|
||||
const MALFORMED_MARKER: &str = "SENSITIVE_B4G_CREATE_TABLE_METADATA_MARKER_9d2e_a7c1";
|
||||
|
||||
/// Counts [`Scannable::scan_as_stream`] calls. [`Scannable::schema`] is free.
|
||||
struct ObservableScannable {
|
||||
batch: RecordBatch,
|
||||
scan_calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
impl ObservableScannable {
|
||||
fn new(batch: RecordBatch, scan_calls: Arc<AtomicUsize>) -> Self {
|
||||
Self { batch, scan_calls }
|
||||
}
|
||||
}
|
||||
|
||||
impl Scannable for ObservableScannable {
|
||||
fn schema(&self) -> SchemaRef {
|
||||
self.batch.schema()
|
||||
}
|
||||
|
||||
fn scan_as_stream(&mut self) -> SendableRecordBatchStream {
|
||||
self.scan_calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.batch.scan_as_stream()
|
||||
}
|
||||
|
||||
fn num_rows(&self) -> Option<usize> {
|
||||
Some(self.batch.num_rows())
|
||||
}
|
||||
|
||||
fn rescannable(&self) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
fn literal_definition(output_field_id: i32) -> GeneratedColumnDefinition {
|
||||
let function = Function::new(
|
||||
FunctionId::try_new(FN_ID).unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("literal-only")])) as arrow_array::ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
GeneratedColumnDefinition::try_new(output_field_id, call, 1, 1).unwrap()
|
||||
}
|
||||
|
||||
fn valid_reserved_payload() -> String {
|
||||
literal_definition(1).to_metadata_json().unwrap()
|
||||
}
|
||||
|
||||
fn malformed_reserved_payload() -> String {
|
||||
format!(
|
||||
r#"{{"format_version":1,"output_field_id":1,"function_call":"{MALFORMED_MARKER}","dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
)
|
||||
}
|
||||
|
||||
fn batch_with_field_metadata(metadata: HashMap<String, String>) -> RecordBatch {
|
||||
let gen_field = Field::new(GEN_OUT, DataType::Int32, true).with_metadata(metadata);
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(ID, DataType::Int32, false),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
gen_field,
|
||||
]));
|
||||
RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
Arc::new(StringArray::from(vec![Some("seed")])),
|
||||
Arc::new(Int32Array::from(vec![10])),
|
||||
],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn reserved_batch(payload: &str) -> RecordBatch {
|
||||
batch_with_field_metadata(
|
||||
[(
|
||||
GENERATED_COLUMN_METADATA_KEY.to_string(),
|
||||
payload.to_string(),
|
||||
)]
|
||||
.into(),
|
||||
)
|
||||
}
|
||||
|
||||
fn ordinary_metadata_batch() -> RecordBatch {
|
||||
batch_with_field_metadata(
|
||||
[(
|
||||
ORDINARY_META_KEY.to_string(),
|
||||
ORDINARY_META_VALUE.to_string(),
|
||||
)]
|
||||
.into(),
|
||||
)
|
||||
}
|
||||
|
||||
fn plain_seed_batch() -> RecordBatch {
|
||||
batch_with_field_metadata(HashMap::new())
|
||||
}
|
||||
|
||||
fn assert_not_supported_redacted(err: &Error, label: &str, forbidden_substrings: &[&str]) {
|
||||
match err {
|
||||
Error::NotSupported { message } => {
|
||||
let rendered = format!("{err}\n{err:?}\n{message}");
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: leaked metadata wire key: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(FN_ID),
|
||||
"{label}: leaked Function ID: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(MALFORMED_MARKER),
|
||||
"{label}: leaked malformed marker: {rendered}"
|
||||
);
|
||||
for needle in forbidden_substrings {
|
||||
assert!(
|
||||
!rendered.contains(needle),
|
||||
"{label}: leaked forbidden substring `{needle}`: {rendered}"
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
message.to_lowercase().contains("generated")
|
||||
|| message.to_lowercase().contains("job"),
|
||||
"{label}: message must describe Job-owned generated-column boundary: {message}"
|
||||
);
|
||||
}
|
||||
other => panic!("{label}: expected Error::NotSupported, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
async fn listing_db() -> (TempDir, ListingDatabase) {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().to_str().unwrap();
|
||||
let request = crate::connection::ConnectRequest {
|
||||
uri: uri.to_string(),
|
||||
#[cfg(feature = "remote")]
|
||||
client_config: Default::default(),
|
||||
options: Default::default(),
|
||||
namespace_client_properties: Default::default(),
|
||||
manifest_enabled: false,
|
||||
read_consistency_interval: None,
|
||||
session: None,
|
||||
};
|
||||
let db = ListingDatabase::connect_with_options(&request)
|
||||
.await
|
||||
.unwrap();
|
||||
(tmp, db)
|
||||
}
|
||||
|
||||
fn listing_table_dir(tmp: &TempDir, name: &str) -> std::path::PathBuf {
|
||||
tmp.path().join(format!("{name}.lance"))
|
||||
}
|
||||
|
||||
async fn listing_create(
|
||||
db: &ListingDatabase,
|
||||
name: &str,
|
||||
data: Box<dyn Scannable>,
|
||||
mode: CreateTableMode,
|
||||
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
|
||||
db.create_table(CreateTableRequest {
|
||||
name: name.to_string(),
|
||||
namespace_path: vec![],
|
||||
data,
|
||||
mode,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn assert_listing_absent(db: &ListingDatabase, tmp: &TempDir, name: &str) {
|
||||
#[allow(deprecated)]
|
||||
let names = db.table_names(TableNamesRequest::default()).await.unwrap();
|
||||
assert!(
|
||||
!names.contains(&name.to_string()),
|
||||
"rejected create must leave no listed table `{name}`; got {names:?}"
|
||||
);
|
||||
assert!(
|
||||
!listing_table_dir(tmp, name).exists(),
|
||||
"rejected create must leave no storage directory for `{name}`"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_create_rejects_reserved_generated_column_metadata_before_scan() {
|
||||
let (tmp, db) = listing_db().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = listing_create(&db, "b4g_listing_create", data, CreateTableMode::Create)
|
||||
.await
|
||||
.expect_err("listing Create must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"listing Create reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(
|
||||
scan_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"rejection must occur before Scannable::scan_as_stream"
|
||||
);
|
||||
assert_listing_absent(&db, &tmp, "b4g_listing_create").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_overwrite_rejects_reserved_generated_column_metadata_and_preserves_table() {
|
||||
let (tmp, db) = listing_db().await;
|
||||
let seed = listing_create(
|
||||
&db,
|
||||
"b4g_listing_overwrite",
|
||||
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
|
||||
CreateTableMode::Create,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let version_before = seed.version().await.unwrap();
|
||||
let schema_before = seed.schema().await.unwrap();
|
||||
assert!(
|
||||
!schema_before
|
||||
.field_with_name(GEN_OUT)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
);
|
||||
|
||||
let payload = malformed_reserved_payload();
|
||||
assert!(payload.contains(MALFORMED_MARKER));
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = listing_create(
|
||||
&db,
|
||||
"b4g_listing_overwrite",
|
||||
data,
|
||||
CreateTableMode::Overwrite,
|
||||
)
|
||||
.await
|
||||
.expect_err("listing Overwrite must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"listing Overwrite reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
let reopened = db
|
||||
.open_table(crate::database::OpenTableRequest {
|
||||
name: "b4g_listing_overwrite".to_string(),
|
||||
namespace_path: vec![],
|
||||
index_cache_size: None,
|
||||
lance_read_params: None,
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
managed_versioning: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(reopened.version().await.unwrap(), version_before);
|
||||
let schema_after = reopened.schema().await.unwrap();
|
||||
assert_eq!(schema_after.as_ref(), schema_before.as_ref());
|
||||
assert!(
|
||||
!schema_after
|
||||
.field_with_name(GEN_OUT)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
);
|
||||
assert!(listing_table_dir(&tmp, "b4g_listing_overwrite").exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_exist_ok_absent_rejects_reserved_generated_column_metadata_before_scan() {
|
||||
let (tmp, db) = listing_db().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = listing_create(
|
||||
&db,
|
||||
"b4g_listing_exist_ok",
|
||||
data,
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
)
|
||||
.await
|
||||
.expect_err("listing ExistOk (absent) must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"listing ExistOk absent reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
assert_listing_absent(&db, &tmp, "b4g_listing_exist_ok").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn listing_ordinary_field_metadata_is_accepted_and_preserved() {
|
||||
let (_tmp, db) = listing_db().await;
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
ordinary_metadata_batch(),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let table = listing_create(&db, "b4g_listing_ordinary", data, CreateTableMode::Create)
|
||||
.await
|
||||
.expect("ordinary field metadata must remain accepted");
|
||||
assert!(
|
||||
scan_calls.load(Ordering::SeqCst) > 0,
|
||||
"successful create may consume the Scannable"
|
||||
);
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
let md = schema.field_with_name(GEN_OUT).unwrap().metadata();
|
||||
assert_eq!(
|
||||
md.get(ORDINARY_META_KEY).map(String::as_str),
|
||||
Some(ORDINARY_META_VALUE)
|
||||
);
|
||||
assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY));
|
||||
}
|
||||
|
||||
#[cfg(not(windows))] // directory namespace tests are unix-only in this crate
|
||||
mod namespace_admission {
|
||||
use super::*;
|
||||
use crate::connect_namespace;
|
||||
use lance_namespace::models::{CreateNamespaceRequest, DescribeTableRequest};
|
||||
|
||||
async fn namespace_conn() -> (TempDir, crate::Connection) {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let root = tmp.path().to_str().unwrap().to_string();
|
||||
let mut properties = HashMap::new();
|
||||
properties.insert("root".to_string(), root);
|
||||
let conn = connect_namespace("dir", properties)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
conn.create_namespace(CreateNamespaceRequest {
|
||||
id: Some(vec!["b4g_ns".into()]),
|
||||
..Default::default()
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
(tmp, conn)
|
||||
}
|
||||
|
||||
async fn assert_namespace_undeclared(conn: &crate::Connection, name: &str) {
|
||||
let names = conn
|
||||
.table_names()
|
||||
.namespace(vec!["b4g_ns".into()])
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(
|
||||
!names.contains(&name.to_string()),
|
||||
"rejected namespace create must leave no declared/listed table `{name}`; got {names:?}"
|
||||
);
|
||||
let ns = conn.namespace_client().await.unwrap();
|
||||
let describe = ns
|
||||
.describe_table(DescribeTableRequest {
|
||||
id: Some(vec!["b4g_ns".into(), name.into()]),
|
||||
..Default::default()
|
||||
})
|
||||
.await;
|
||||
assert!(
|
||||
describe.is_err(),
|
||||
"rejected namespace create must leave no describable table `{name}`"
|
||||
);
|
||||
}
|
||||
|
||||
async fn namespace_create(
|
||||
conn: &crate::Connection,
|
||||
name: &str,
|
||||
data: Box<dyn Scannable>,
|
||||
mode: CreateTableMode,
|
||||
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
|
||||
conn.database()
|
||||
.create_table(CreateTableRequest {
|
||||
name: name.to_string(),
|
||||
namespace_path: vec!["b4g_ns".into()],
|
||||
data,
|
||||
mode,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_create_rejects_reserved_before_declare_describe_or_storage() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(&conn, "b4g_ns_create", data, CreateTableMode::Create)
|
||||
.await
|
||||
.expect_err("namespace Create must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace Create reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
assert_namespace_undeclared(&conn, "b4g_ns_create").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_overwrite_rejects_reserved_before_declare_describe_or_storage() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let seed = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_overwrite",
|
||||
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
|
||||
CreateTableMode::Create,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let version_before = seed.version().await.unwrap();
|
||||
let schema_before = seed.schema().await.unwrap();
|
||||
|
||||
let payload = malformed_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(&conn, "b4g_ns_overwrite", data, CreateTableMode::Overwrite)
|
||||
.await
|
||||
.expect_err("namespace Overwrite must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace Overwrite reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
let reopened = conn
|
||||
.database()
|
||||
.open_table(crate::database::OpenTableRequest {
|
||||
name: "b4g_ns_overwrite".to_string(),
|
||||
namespace_path: vec!["b4g_ns".into()],
|
||||
index_cache_size: None,
|
||||
lance_read_params: None,
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
managed_versioning: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(reopened.version().await.unwrap(), version_before);
|
||||
assert_eq!(
|
||||
reopened.schema().await.unwrap().as_ref(),
|
||||
schema_before.as_ref()
|
||||
);
|
||||
assert!(
|
||||
!reopened
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name(GEN_OUT)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_exist_ok_absent_rejects_reserved_before_declare_describe_or_storage() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_exist_ok",
|
||||
data,
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
)
|
||||
.await
|
||||
.expect_err("namespace ExistOk (absent) must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace ExistOk absent reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
assert_namespace_undeclared(&conn, "b4g_ns_exist_ok").await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn namespace_exist_ok_existing_rejects_reserved_even_when_mode_would_ignore_data() {
|
||||
let (_tmp, conn) = namespace_conn().await;
|
||||
let seed = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_exist_ok_existing",
|
||||
Box::new(plain_seed_batch()) as Box<dyn Scannable>,
|
||||
CreateTableMode::Create,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let version_before = seed.version().await.unwrap();
|
||||
let schema_before = seed.schema().await.unwrap();
|
||||
|
||||
let payload = valid_reserved_payload();
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(&payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = namespace_create(
|
||||
&conn,
|
||||
"b4g_ns_exist_ok_existing",
|
||||
data,
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
)
|
||||
.await
|
||||
.expect_err(
|
||||
"namespace ExistOk must not accept reserved metadata merely because data is ignored",
|
||||
);
|
||||
assert_not_supported_redacted(
|
||||
&err,
|
||||
"namespace ExistOk existing reserved admission",
|
||||
&[payload.as_str()],
|
||||
);
|
||||
assert_eq!(scan_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
let reopened = conn
|
||||
.database()
|
||||
.open_table(crate::database::OpenTableRequest {
|
||||
name: "b4g_ns_exist_ok_existing".to_string(),
|
||||
namespace_path: vec!["b4g_ns".into()],
|
||||
index_cache_size: None,
|
||||
lance_read_params: None,
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
managed_versioning: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(reopened.version().await.unwrap(), version_before);
|
||||
assert_eq!(
|
||||
reopened.schema().await.unwrap().as_ref(),
|
||||
schema_before.as_ref()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "remote")]
|
||||
mod remote_admission {
|
||||
use super::*;
|
||||
use std::io::Cursor;
|
||||
|
||||
use arrow_ipc::reader::StreamReader;
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::Connection;
|
||||
use crate::remote::{ClientConfig, HeaderProvider};
|
||||
|
||||
#[derive(Debug)]
|
||||
struct CountingHeaderProvider {
|
||||
calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl HeaderProvider for CountingHeaderProvider {
|
||||
async fn get_headers(&self) -> crate::Result<HashMap<String, String>> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(HashMap::from([(
|
||||
"X-B4g-Test".to_string(),
|
||||
"must-not-be-requested".to_string(),
|
||||
)]))
|
||||
}
|
||||
}
|
||||
|
||||
fn counting_handler(
|
||||
calls: Arc<AtomicUsize>,
|
||||
) -> impl Fn(reqwest::Request) -> http::Response<String> + Clone + Send + Sync + 'static {
|
||||
move |_request| {
|
||||
calls.fetch_add(1, Ordering::SeqCst);
|
||||
http::Response::builder()
|
||||
.status(200)
|
||||
.body(String::new())
|
||||
.unwrap()
|
||||
}
|
||||
}
|
||||
|
||||
async fn remote_create(
|
||||
conn: &Connection,
|
||||
name: &str,
|
||||
data: Box<dyn Scannable>,
|
||||
mode: CreateTableMode,
|
||||
) -> crate::Result<Arc<dyn crate::table::BaseTable>> {
|
||||
// Direct Database trait path used by Connection::create_table.
|
||||
conn.database()
|
||||
.create_table(CreateTableRequest {
|
||||
name: name.to_string(),
|
||||
namespace_path: vec![],
|
||||
data,
|
||||
mode,
|
||||
write_options: Default::default(),
|
||||
location: None,
|
||||
namespace_client: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn assert_remote_rejects(
|
||||
mode: CreateTableMode,
|
||||
table_name: &str,
|
||||
payload: &str,
|
||||
label: &str,
|
||||
) {
|
||||
let handler_calls = Arc::new(AtomicUsize::new(0));
|
||||
let header_calls = Arc::new(AtomicUsize::new(0));
|
||||
let scan_calls = Arc::new(AtomicUsize::new(0));
|
||||
let config = ClientConfig {
|
||||
header_provider: Some(Arc::new(CountingHeaderProvider {
|
||||
calls: header_calls.clone(),
|
||||
}) as Arc<dyn HeaderProvider>),
|
||||
..Default::default()
|
||||
};
|
||||
let conn = Connection::new_with_handler_and_config(
|
||||
counting_handler(handler_calls.clone()),
|
||||
config,
|
||||
);
|
||||
let data = Box::new(ObservableScannable::new(
|
||||
reserved_batch(payload),
|
||||
scan_calls.clone(),
|
||||
)) as Box<dyn Scannable>;
|
||||
|
||||
let err = remote_create(&conn, table_name, data, mode)
|
||||
.await
|
||||
.expect_err("remote create must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(&err, label, &[payload]);
|
||||
assert_eq!(
|
||||
scan_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"{label}: rejection must occur before scan_as_stream"
|
||||
);
|
||||
assert_eq!(
|
||||
header_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"{label}: rejection must occur before header-provider invocation"
|
||||
);
|
||||
assert_eq!(
|
||||
handler_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"{label}: rejection must occur before HTTP handler"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_create_rejects_reserved_before_scan_headers_and_http() {
|
||||
assert_remote_rejects(
|
||||
CreateTableMode::Create,
|
||||
"b4g_remote_create",
|
||||
&valid_reserved_payload(),
|
||||
"remote Create reserved admission",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_overwrite_rejects_reserved_before_scan_headers_and_http() {
|
||||
assert_remote_rejects(
|
||||
CreateTableMode::Overwrite,
|
||||
"b4g_remote_overwrite",
|
||||
&malformed_reserved_payload(),
|
||||
"remote Overwrite reserved admission",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_exist_ok_rejects_reserved_before_scan_headers_and_http() {
|
||||
assert_remote_rejects(
|
||||
CreateTableMode::exist_ok(|req| req),
|
||||
"b4g_remote_exist_ok",
|
||||
&valid_reserved_payload(),
|
||||
"remote ExistOk reserved admission",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_ordinary_field_metadata_is_transmitted_unchanged() {
|
||||
let conn = Connection::new_with_handler(|request| {
|
||||
assert_eq!(request.method(), &reqwest::Method::POST);
|
||||
assert_eq!(
|
||||
request.url().path(),
|
||||
"/v1/table/b4g_remote_ordinary/create/"
|
||||
);
|
||||
let body = request
|
||||
.body()
|
||||
.and_then(|b| b.as_bytes())
|
||||
.expect("ordinary create must send an Arrow IPC body");
|
||||
let reader = StreamReader::try_new(Cursor::new(body), None).unwrap();
|
||||
let schema = reader.schema();
|
||||
let md = schema.field_with_name(GEN_OUT).unwrap().metadata();
|
||||
assert_eq!(
|
||||
md.get(ORDINARY_META_KEY).map(String::as_str),
|
||||
Some(ORDINARY_META_VALUE),
|
||||
"ordinary field metadata must be transmitted unchanged"
|
||||
);
|
||||
assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY));
|
||||
// Consume stream to completion for a well-formed IPC body.
|
||||
for batch in reader {
|
||||
batch.unwrap();
|
||||
}
|
||||
http::Response::builder()
|
||||
.status(200)
|
||||
.body(String::new())
|
||||
.unwrap()
|
||||
});
|
||||
|
||||
conn.create_table("b4g_remote_ordinary", ordinary_metadata_batch())
|
||||
.mode(CreateTableMode::Create)
|
||||
.execute()
|
||||
.await
|
||||
.expect("ordinary field metadata must remain accepted on remote create");
|
||||
}
|
||||
}
|
||||
@@ -23,6 +23,7 @@ use crate::connection::ConnectRequest;
|
||||
use crate::database::ReadConsistency;
|
||||
use crate::database::namespace::LanceNamespaceDatabase;
|
||||
use crate::error::{CreateDirSnafu, Error, Result};
|
||||
use crate::function::schema_admission::reject_caller_authored_generated_column_schema;
|
||||
use crate::io::object_store::MirroringObjectStoreWrapper;
|
||||
use crate::table::NativeTable;
|
||||
use crate::utils::validate_table_name;
|
||||
@@ -1038,6 +1039,10 @@ impl Database for ListingDatabase {
|
||||
}
|
||||
|
||||
async fn create_table(&self, request: CreateTableRequest) -> Result<Arc<dyn BaseTable>> {
|
||||
// Admit schema before namespace forwarding, URI/config work, or NativeTable::create.
|
||||
// Scannable::schema is free; must not call scan_as_stream yet.
|
||||
reject_caller_authored_generated_column_schema(request.data.schema().as_ref())?;
|
||||
|
||||
if !request.namespace_path.is_empty() {
|
||||
return self.namespace_database().create_table(request).await;
|
||||
}
|
||||
|
||||
@@ -34,6 +34,7 @@ use crate::database::read_freshness::{
|
||||
FreshnessBaselines, ReadFreshnessContextProvider, TableFreshness,
|
||||
};
|
||||
use crate::error::{Error, Result};
|
||||
use crate::function::schema_admission::reject_caller_authored_generated_column_schema;
|
||||
use crate::table::{NativeTable, map_namespace_lance_error};
|
||||
use lance::dataset::WriteMode;
|
||||
|
||||
@@ -349,6 +350,10 @@ impl Database for LanceNamespaceDatabase {
|
||||
}
|
||||
|
||||
async fn create_table(&self, request: DbCreateTableRequest) -> Result<Arc<dyn BaseTable>> {
|
||||
// Admit schema before any mode branch, describe, declare, or storage work.
|
||||
// Scannable::schema is free; must not call scan_as_stream yet.
|
||||
reject_caller_authored_generated_column_schema(request.data.schema().as_ref())?;
|
||||
|
||||
let mut table_id = request.namespace_path.clone();
|
||||
table_id.push(request.name.clone());
|
||||
let mut existing_table = None;
|
||||
|
||||
@@ -6,10 +6,91 @@ use std::sync::{Arc, PoisonError};
|
||||
|
||||
use arrow_schema::ArrowError;
|
||||
use datafusion_common::DataFusionError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
use snafu::Snafu;
|
||||
|
||||
pub(crate) type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||
|
||||
/// Stable Function error category (FF-006).
|
||||
///
|
||||
/// The known variants serialize to fixed JSON strings. Any other wire string
|
||||
/// decodes as [`Self::Unrecognized`] with the exact value preserved, and
|
||||
/// re-serializes unchanged. Category judgment is structural: do not infer a
|
||||
/// code from diagnostic message text, HTTP status, job phase, or retryability.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[non_exhaustive]
|
||||
pub enum FunctionErrorCode {
|
||||
/// Function definition failed validation.
|
||||
DefinitionValidationFailure,
|
||||
/// A named Function or Function reference was not found.
|
||||
NameOrFunctionNotFound,
|
||||
/// A Function name conflicts with an existing name.
|
||||
NameConflict,
|
||||
/// The requested runtime or capability is not supported.
|
||||
UnsupportedRuntimeOrCapability,
|
||||
/// The Function has been revoked and cannot be used.
|
||||
RevokedFunction,
|
||||
/// User-defined Function execution failed.
|
||||
UdfExecutionFailure,
|
||||
/// A generated column was not fully materialized.
|
||||
GeneratedColumnIncomplete,
|
||||
/// Input was stale or conflicted with the current state.
|
||||
StaleOrConflictingInput,
|
||||
/// A wire string this client version does not recognize.
|
||||
///
|
||||
/// The inner value is preserved exactly for forward compatibility.
|
||||
Unrecognized(String),
|
||||
}
|
||||
|
||||
impl FunctionErrorCode {
|
||||
/// The stable JSON / wire string for this code.
|
||||
pub fn as_str(&self) -> &str {
|
||||
match self {
|
||||
Self::DefinitionValidationFailure => "definition_validation_failure",
|
||||
Self::NameOrFunctionNotFound => "name_or_function_not_found",
|
||||
Self::NameConflict => "name_conflict",
|
||||
Self::UnsupportedRuntimeOrCapability => "unsupported_runtime_or_capability",
|
||||
Self::RevokedFunction => "revoked_function",
|
||||
Self::UdfExecutionFailure => "udf_execution_failure",
|
||||
Self::GeneratedColumnIncomplete => "generated_column_incomplete",
|
||||
Self::StaleOrConflictingInput => "stale_or_conflicting_input",
|
||||
Self::Unrecognized(raw) => raw.as_str(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(value: &str) -> Self {
|
||||
match value {
|
||||
"definition_validation_failure" => Self::DefinitionValidationFailure,
|
||||
"name_or_function_not_found" => Self::NameOrFunctionNotFound,
|
||||
"name_conflict" => Self::NameConflict,
|
||||
"unsupported_runtime_or_capability" => Self::UnsupportedRuntimeOrCapability,
|
||||
"revoked_function" => Self::RevokedFunction,
|
||||
"udf_execution_failure" => Self::UdfExecutionFailure,
|
||||
"generated_column_incomplete" => Self::GeneratedColumnIncomplete,
|
||||
"stale_or_conflicting_input" => Self::StaleOrConflictingInput,
|
||||
other => Self::Unrecognized(other.to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Display for FunctionErrorCode {
|
||||
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
|
||||
f.write_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
impl Serialize for FunctionErrorCode {
|
||||
fn serialize<S: Serializer>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error> {
|
||||
serializer.serialize_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for FunctionErrorCode {
|
||||
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> std::result::Result<Self, D::Error> {
|
||||
Ok(Self::from_wire(String::deserialize(deserializer)?.as_str()))
|
||||
}
|
||||
}
|
||||
|
||||
/// Why a job failed, to whatever precision the backend provides.
|
||||
///
|
||||
/// A job run in this process carries the error it failed with in [`Self::source`].
|
||||
@@ -18,6 +99,12 @@ pub(crate) type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||
/// backend does not supply it.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct JobFailure {
|
||||
/// Stable Function error category, when the backend supplied one.
|
||||
///
|
||||
/// Present only when copied from [`Error::Function`] or decoded from a
|
||||
/// remote `error_code` field. Never inferred from message, phase,
|
||||
/// retryable, HTTP status, or other diagnostics.
|
||||
pub error_code: Option<FunctionErrorCode>,
|
||||
/// The stage the job was in, when known.
|
||||
pub phase: Option<String>,
|
||||
/// A human-readable reason, when known.
|
||||
@@ -30,8 +117,16 @@ pub struct JobFailure {
|
||||
|
||||
impl JobFailure {
|
||||
/// A failure whose only known detail is the error that caused it.
|
||||
///
|
||||
/// When `source` is [`Error::Function`], [`Self::error_code`] is copied
|
||||
/// from that error. Other error kinds leave `error_code` as [`None`].
|
||||
pub(crate) fn from_source(source: Arc<Error>) -> Self {
|
||||
let error_code = match source.as_ref() {
|
||||
Error::Function { code, .. } => Some(code.clone()),
|
||||
_ => None,
|
||||
};
|
||||
Self {
|
||||
error_code,
|
||||
message: Some(source.to_string()),
|
||||
source: Some(source),
|
||||
..Default::default()
|
||||
@@ -92,6 +187,15 @@ pub enum Error {
|
||||
},
|
||||
#[snafu(display("Job{} was cancelled", job_id.as_ref().map(|id| format!(" {id}")).unwrap_or_default()))]
|
||||
JobCancelled { job_id: Option<String> },
|
||||
/// A first-class Function operation failed with a stable category.
|
||||
///
|
||||
/// [`Self::Function::code`] is the semantic category. [`Self::Function::message`]
|
||||
/// is diagnostic only and must not be used to recover or override the code.
|
||||
#[snafu(display("Function error ({code}): {message}"))]
|
||||
Function {
|
||||
code: FunctionErrorCode,
|
||||
message: String,
|
||||
},
|
||||
|
||||
// 3rd party / external errors
|
||||
#[snafu(display("object_store error: {source}"))]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,776 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Atomic generated-column binding snapshot projection (FF-029).
|
||||
//!
|
||||
//! This is an implementation projection for table call binding. It is not a
|
||||
//! catalog resource, Job, persistent model, wire payload, or table-version
|
||||
//! replacement.
|
||||
|
||||
use std::collections::HashSet;
|
||||
|
||||
use arrow_schema::FieldRef;
|
||||
|
||||
use super::{
|
||||
FunctionCall, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition, invalid_input,
|
||||
};
|
||||
use crate::Result;
|
||||
|
||||
/// One top-level field identity from a single table snapshot.
|
||||
///
|
||||
/// Pairs a non-negative Lance stable field ID with the exact Arrow field from
|
||||
/// that same snapshot. IDs are carried only here; they are never injected into
|
||||
/// Arrow field metadata.
|
||||
#[doc(hidden)]
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct GeneratedColumnBindingEntry {
|
||||
field_id: i32,
|
||||
field: FieldRef,
|
||||
}
|
||||
|
||||
impl GeneratedColumnBindingEntry {
|
||||
/// Stable Lance field ID for this top-level entry.
|
||||
pub fn field_id(&self) -> i32 {
|
||||
self.field_id
|
||||
}
|
||||
|
||||
/// Exact Arrow field from the same snapshot.
|
||||
pub fn field(&self) -> &FieldRef {
|
||||
&self.field
|
||||
}
|
||||
|
||||
/// Strict generated-column definition from this entry's Arrow metadata.
|
||||
///
|
||||
/// Reads only [`GENERATED_COLUMN_METADATA_KEY`] on the exact snapshot field
|
||||
/// and decodes through
|
||||
/// [`GeneratedColumnDefinition::from_metadata_json`] with
|
||||
/// [`Self::field_id`] as the expected output identity. The same-snapshot
|
||||
/// stable field ID is mandatory so decode rejects metadata whose embedded
|
||||
/// `output_field_id` does not match this entry; name/ordinal/hash fallbacks
|
||||
/// are not used.
|
||||
///
|
||||
/// Returns [`Ok`]`(`[`None`]`)` when the key is absent. Present but invalid
|
||||
/// metadata fails closed as [`crate::Error::InvalidInput`] with a short
|
||||
/// field-ID diagnostic that does not echo the raw metadata payload.
|
||||
pub(crate) fn generated_column_definition(&self) -> Result<Option<GeneratedColumnDefinition>> {
|
||||
let Some(raw) = self.field.metadata().get(GENERATED_COLUMN_METADATA_KEY) else {
|
||||
return Ok(None);
|
||||
};
|
||||
match GeneratedColumnDefinition::from_metadata_json(raw, self.field_id) {
|
||||
Ok(definition) => Ok(Some(definition)),
|
||||
Err(_) => Err(invalid_input(format!(
|
||||
"invalid generated-column metadata for field id {}",
|
||||
self.field_id
|
||||
))),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Atomic table snapshot projection for generated-column call binding.
|
||||
///
|
||||
/// Contains one table version and immutable top-level field entries in schema
|
||||
/// order. Construction validates field/ID count equality, non-negative unique
|
||||
/// IDs, and unique top-level names.
|
||||
#[doc(hidden)]
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct GeneratedColumnBindingSnapshot {
|
||||
version: u64,
|
||||
entries: Vec<GeneratedColumnBindingEntry>,
|
||||
}
|
||||
|
||||
impl GeneratedColumnBindingSnapshot {
|
||||
/// Build a binding snapshot from one version and ordered field/ID pairs.
|
||||
///
|
||||
/// `fields` and `field_ids` must have the same length. Every ID must be
|
||||
/// non-negative and unique. Top-level field names must be unique. Order is
|
||||
/// preserved exactly as provided.
|
||||
pub fn try_new(
|
||||
version: u64,
|
||||
fields: impl IntoIterator<Item = FieldRef>,
|
||||
field_ids: impl IntoIterator<Item = i32>,
|
||||
) -> Result<Self> {
|
||||
let fields: Vec<FieldRef> = fields.into_iter().collect();
|
||||
let field_ids: Vec<i32> = field_ids.into_iter().collect();
|
||||
if fields.len() != field_ids.len() {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot field count must equal field_ids count",
|
||||
));
|
||||
}
|
||||
|
||||
let mut seen_ids = HashSet::with_capacity(field_ids.len());
|
||||
let mut seen_names = HashSet::with_capacity(fields.len());
|
||||
let mut entries = Vec::with_capacity(fields.len());
|
||||
|
||||
for (field, field_id) in fields.into_iter().zip(field_ids) {
|
||||
if field_id < 0 {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot field IDs must be non-negative",
|
||||
));
|
||||
}
|
||||
if !seen_ids.insert(field_id) {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot field IDs must be unique",
|
||||
));
|
||||
}
|
||||
if !seen_names.insert(field.name().clone()) {
|
||||
return Err(invalid_input(
|
||||
"generated-column binding snapshot top-level field names must be unique",
|
||||
));
|
||||
}
|
||||
entries.push(GeneratedColumnBindingEntry { field_id, field });
|
||||
}
|
||||
|
||||
Ok(Self { version, entries })
|
||||
}
|
||||
|
||||
/// Table version for this snapshot.
|
||||
pub fn version(&self) -> u64 {
|
||||
self.version
|
||||
}
|
||||
|
||||
/// Top-level entries in schema order.
|
||||
pub fn entries(&self) -> &[GeneratedColumnBindingEntry] {
|
||||
&self.entries
|
||||
}
|
||||
|
||||
/// Exact case-sensitive top-level field name lookup.
|
||||
///
|
||||
/// A name containing `.` is a literal top-level field name, not a nested
|
||||
/// path. Lookup does not fold case or interpret dotted selectors.
|
||||
pub fn field(&self, name: &str) -> Option<&GeneratedColumnBindingEntry> {
|
||||
self.entries
|
||||
.iter()
|
||||
.find(|entry| entry.field().name() == name)
|
||||
}
|
||||
|
||||
/// Strict generated-column definition for one top-level column name.
|
||||
///
|
||||
/// Looks up the exact case-sensitive top-level name (`.` is literal, not a
|
||||
/// nested path), decodes through
|
||||
/// [`GeneratedColumnBindingEntry::generated_column_definition`] (preserving
|
||||
/// output stable-ID checking and raw-metadata redaction), then validates
|
||||
/// stored field arguments against this same snapshot via
|
||||
/// [`Self::validate_field_arguments`]. Returns the complete or incomplete
|
||||
/// definition unchanged. Does not perform table, catalog, network, or Job
|
||||
/// work and does not resolve a Function.
|
||||
///
|
||||
/// Returns [`crate::Error::InvalidInput`] for an empty name, a missing
|
||||
/// top-level field, an ordinary field without a valid generated-column
|
||||
/// definition, invalid metadata, or a field-argument identity/type
|
||||
/// mismatch against this snapshot.
|
||||
#[doc(hidden)]
|
||||
pub fn generated_column_definition(
|
||||
&self,
|
||||
column_name: impl AsRef<str>,
|
||||
) -> Result<GeneratedColumnDefinition> {
|
||||
let column_name = column_name.as_ref();
|
||||
if column_name.is_empty() {
|
||||
return Err(invalid_input("generated column name must not be empty"));
|
||||
}
|
||||
let Some(entry) = self.field(column_name) else {
|
||||
return Err(invalid_input(format!(
|
||||
"generated column '{column_name}' was not found in the table schema"
|
||||
)));
|
||||
};
|
||||
let Some(definition) = entry.generated_column_definition()? else {
|
||||
return Err(invalid_input(format!(
|
||||
"column '{column_name}' is not a generated column"
|
||||
)));
|
||||
};
|
||||
self.validate_field_arguments(definition.function_call())?;
|
||||
Ok(definition)
|
||||
}
|
||||
|
||||
/// Validate table-dependent field arguments of an already canonical call.
|
||||
///
|
||||
/// For every field argument, finds the snapshot entry by stable Lance field
|
||||
/// ID and requires exact Arrow [`arrow_schema::DataType`] equality. Literal
|
||||
/// arguments are table-independent and ignored. Missing field ID or type
|
||||
/// mismatch returns [`crate::Error::InvalidInput`] without modifying `call`
|
||||
/// or this snapshot.
|
||||
///
|
||||
/// This check is orthogonal to [`FunctionCall::validate_against`]: it does
|
||||
/// not perform catalog lookup, Function identity/signature validation, or
|
||||
/// table mutation.
|
||||
pub fn validate_field_arguments(&self, call: &FunctionCall) -> Result<()> {
|
||||
for (_parameter, argument) in call.arguments() {
|
||||
let Some(field_id) = argument.field_id() else {
|
||||
continue;
|
||||
};
|
||||
let Some(entry) = self.entry_by_field_id(field_id) else {
|
||||
return Err(invalid_input(format!(
|
||||
"generated-column binding snapshot missing field id {field_id}"
|
||||
)));
|
||||
};
|
||||
let expected = argument.data_type();
|
||||
let current = entry.field().data_type();
|
||||
if current != expected {
|
||||
return Err(invalid_input(format!(
|
||||
"generated-column binding snapshot field id {field_id} type mismatch: \
|
||||
expected {expected}, found {current}"
|
||||
)));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn entry_by_field_id(&self, field_id: i32) -> Option<&GeneratedColumnBindingEntry> {
|
||||
self.entries
|
||||
.iter()
|
||||
.find(|entry| entry.field_id() == field_id)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_schema::{DataType, Field};
|
||||
|
||||
use super::*;
|
||||
use crate::Error;
|
||||
|
||||
fn fields() -> Vec<FieldRef> {
|
||||
vec![
|
||||
Arc::new(Field::new("text", DataType::Utf8, true)),
|
||||
Arc::new(Field::new("Score", DataType::Int32, false)),
|
||||
Arc::new(Field::new("a.b", DataType::Utf8, true)),
|
||||
]
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn try_new_preserves_version_order_and_entry_data() {
|
||||
let snapshot =
|
||||
GeneratedColumnBindingSnapshot::try_new(11, fields(), vec![2, 4, 8]).unwrap();
|
||||
assert_eq!(snapshot.version(), 11);
|
||||
assert_eq!(snapshot.entries().len(), 3);
|
||||
assert_eq!(snapshot.entries()[0].field_id(), 2);
|
||||
assert_eq!(snapshot.entries()[0].field().name(), "text");
|
||||
assert_eq!(snapshot.entries()[0].field().data_type(), &DataType::Utf8);
|
||||
assert_eq!(snapshot.entries()[1].field_id(), 4);
|
||||
assert_eq!(snapshot.entries()[1].field().name(), "Score");
|
||||
assert_eq!(snapshot.entries()[2].field_id(), 8);
|
||||
assert_eq!(snapshot.entries()[2].field().name(), "a.b");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lookup_is_exact_case_sensitive_and_treats_dot_literally() {
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(1, fields(), vec![2, 4, 8]).unwrap();
|
||||
assert_eq!(snapshot.field("Score").unwrap().field_id(), 4);
|
||||
assert!(snapshot.field("score").is_none());
|
||||
assert!(snapshot.field("TEXT").is_none());
|
||||
assert!(snapshot.field("a").is_none());
|
||||
assert!(snapshot.field("b").is_none());
|
||||
assert_eq!(snapshot.field("a.b").unwrap().field_id(), 8);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn try_new_rejects_count_mismatch_negative_duplicate_ids_and_names() {
|
||||
let base = fields();
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, base.clone(), vec![1, 2]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, base.clone(), vec![1, 2, -3]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, base.clone(), vec![1, 2, 1]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
let duplicate_names = vec![
|
||||
Arc::new(Field::new("text", DataType::Utf8, true)),
|
||||
Arc::new(Field::new("text", DataType::Int32, false)),
|
||||
];
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, duplicate_names, vec![1, 2]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
}
|
||||
|
||||
fn sample_function() -> crate::function::Function {
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
let id = FunctionId::try_new("fn.exact.snapshot.lib").unwrap();
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![
|
||||
FunctionParameter::new("payload_arg", DataType::Utf8),
|
||||
FunctionParameter::new("metric_arg", DataType::Int32),
|
||||
],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap();
|
||||
Function::new(id, signature)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_field_arguments_value_cases() {
|
||||
use crate::function::{FunctionArgument, FunctionCall};
|
||||
use arrow_array::{ArrayRef, Int32Array};
|
||||
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(2, fields(), vec![2, 4, 8]).unwrap();
|
||||
let function = sample_function();
|
||||
|
||||
let valid = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(2, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_field(4, DataType::Int32).unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
snapshot.validate_field_arguments(&valid).unwrap();
|
||||
|
||||
let missing = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(99, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_field(4, DataType::Int32).unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
assert!(matches!(
|
||||
snapshot.validate_field_arguments(&missing),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
|
||||
// Same stable ID, different Arrow type: exact-type equality must reject.
|
||||
let type_mismatch = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
// ID 4 is Int32 in the snapshot.
|
||||
FunctionArgument::try_field(4, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_field(4, DataType::Int32).unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
let err = snapshot
|
||||
.validate_field_arguments(&type_mismatch)
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let message = err.to_string();
|
||||
assert!(message.contains('4'));
|
||||
assert!(message.contains("Utf8") && message.contains("Int32"));
|
||||
assert!(!message.contains("Score") && !message.contains("text"));
|
||||
|
||||
let mixed = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(2, DataType::Utf8).unwrap(),
|
||||
),
|
||||
(
|
||||
"metric_arg".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
snapshot.validate_field_arguments(&mixed).unwrap();
|
||||
|
||||
let literal_only_fn = {
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.snapshot.literal").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("constant_arg", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
};
|
||||
let literal_only =
|
||||
FunctionCall::try_new(
|
||||
&literal_only_fn,
|
||||
vec![(
|
||||
"constant_arg".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(Int32Array::from(vec![Some(9)])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
// Empty snapshot still accepts literal-only calls.
|
||||
let empty =
|
||||
GeneratedColumnBindingSnapshot::try_new(1, Vec::<FieldRef>::new(), vec![]).unwrap();
|
||||
empty.validate_field_arguments(&literal_only).unwrap();
|
||||
}
|
||||
|
||||
fn status_sample_function() -> crate::function::Function {
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.status.binding").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn status_sample_call() -> crate::function::FunctionCall {
|
||||
use crate::function::{FunctionArgument, FunctionCall};
|
||||
use arrow_array::{ArrayRef, StringArray};
|
||||
let function = status_sample_function();
|
||||
FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("ok")])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn definition_json(
|
||||
output_field_id: i32,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> String {
|
||||
use crate::function::GeneratedColumnDefinition;
|
||||
GeneratedColumnDefinition::try_new(
|
||||
output_field_id,
|
||||
status_sample_call(),
|
||||
dependency_epoch,
|
||||
materialized_epoch,
|
||||
)
|
||||
.unwrap()
|
||||
.to_metadata_json()
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn entry_with_metadata(
|
||||
name: &str,
|
||||
field_id: i32,
|
||||
metadata_json: Option<&str>,
|
||||
) -> GeneratedColumnBindingEntry {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
let field = if let Some(json) = metadata_json {
|
||||
Field::new(name, DataType::Int32, true).with_metadata(
|
||||
[(GENERATED_COLUMN_METADATA_KEY.to_string(), json.to_string())].into(),
|
||||
)
|
||||
} else {
|
||||
Field::new(name, DataType::Int32, true)
|
||||
};
|
||||
let snapshot =
|
||||
GeneratedColumnBindingSnapshot::try_new(1, vec![Arc::new(field)], vec![field_id])
|
||||
.unwrap();
|
||||
snapshot.entries()[0].clone()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_absent_returns_none() {
|
||||
let entry = entry_with_metadata("ordinary", 3, None);
|
||||
let got = entry.generated_column_definition().unwrap();
|
||||
assert!(got.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_decodes_complete_and_incomplete() {
|
||||
use crate::function::{GeneratedColumnDefinition, GeneratedColumnStatus};
|
||||
|
||||
let complete_json = definition_json(5, 3, 3);
|
||||
let complete_entry = entry_with_metadata("gen_complete", 5, Some(&complete_json));
|
||||
let complete = complete_entry
|
||||
.generated_column_definition()
|
||||
.unwrap()
|
||||
.expect("complete metadata present");
|
||||
assert_eq!(complete.output_field_id(), 5);
|
||||
assert_eq!(complete.dependency_epoch(), 3);
|
||||
assert_eq!(complete.materialized_epoch(), 3);
|
||||
assert_eq!(complete.status(), GeneratedColumnStatus::Complete);
|
||||
assert_eq!(
|
||||
complete,
|
||||
GeneratedColumnDefinition::from_metadata_json(&complete_json, 5).unwrap()
|
||||
);
|
||||
|
||||
let incomplete_json = definition_json(7, 4, 2);
|
||||
let incomplete_entry = entry_with_metadata("gen_incomplete", 7, Some(&incomplete_json));
|
||||
let incomplete = incomplete_entry
|
||||
.generated_column_definition()
|
||||
.unwrap()
|
||||
.expect("incomplete metadata present");
|
||||
assert_eq!(incomplete.output_field_id(), 7);
|
||||
assert_eq!(incomplete.dependency_epoch(), 4);
|
||||
assert_eq!(incomplete.materialized_epoch(), 2);
|
||||
assert_eq!(incomplete.status(), GeneratedColumnStatus::Incomplete);
|
||||
assert_eq!(
|
||||
incomplete,
|
||||
GeneratedColumnDefinition::from_metadata_json(&incomplete_json, 7).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_fail_closed_for_invalid_metadata() {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
|
||||
let field_id = 9i32;
|
||||
let valid = definition_json(field_id, 2, 2);
|
||||
let mut mismatched: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
mismatched["output_field_id"] = serde_json::json!(field_id + 1);
|
||||
|
||||
let mut unsupported: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
unsupported["format_version"] = serde_json::json!(2);
|
||||
|
||||
let mut reversed: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
reversed["dependency_epoch"] = serde_json::json!(1);
|
||||
reversed["materialized_epoch"] = serde_json::json!(2);
|
||||
|
||||
let malformed_json = "{not-json";
|
||||
let mut malformed_call: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
malformed_call["function_call"] = serde_json::json!("not-an-object");
|
||||
|
||||
for raw in [
|
||||
mismatched.to_string(),
|
||||
unsupported.to_string(),
|
||||
reversed.to_string(),
|
||||
malformed_json.to_string(),
|
||||
malformed_call.to_string(),
|
||||
] {
|
||||
let entry = entry_with_metadata("gen_bad", field_id, Some(&raw));
|
||||
assert!(
|
||||
entry
|
||||
.field()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY),
|
||||
"fixture must carry generated-column metadata"
|
||||
);
|
||||
let err = entry.generated_column_definition().unwrap_err();
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_column_definition_errors_omit_raw_metadata_marker() {
|
||||
const MARKER: &str = "SENSITIVE_STATUS_METADATA_MARKER_b3d1_9f2e";
|
||||
let raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":3,"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(raw.contains(MARKER));
|
||||
let entry = entry_with_metadata("gen_redact", 3, Some(&raw));
|
||||
let err = entry.generated_column_definition().unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let text = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!text.contains(MARKER),
|
||||
"status definition diagnostics must not echo raw metadata marker: {text}"
|
||||
);
|
||||
assert!(
|
||||
!text.contains(&raw),
|
||||
"status definition diagnostics must not echo raw metadata payload: {text}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Build a definition whose stored field argument matches `input_type`.
|
||||
/// Construction succeeds even when the snapshot field at `input_field_id`
|
||||
/// later has a different Arrow type; same-snapshot validation catches that.
|
||||
fn field_arg_definition(
|
||||
output_field_id: i32,
|
||||
input_field_id: i32,
|
||||
input_type: DataType,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput,
|
||||
FunctionParameter, FunctionSignature,
|
||||
};
|
||||
let function = Function::new(
|
||||
FunctionId::try_new("fn.exact.snapshot.field_arg").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("payload", input_type.clone())],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"payload".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, input_type).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
GeneratedColumnDefinition::try_new(
|
||||
output_field_id,
|
||||
call,
|
||||
dependency_epoch,
|
||||
materialized_epoch,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn snapshot_with_definition(
|
||||
version: u64,
|
||||
ordinary_name: &str,
|
||||
ordinary_id: i32,
|
||||
ordinary_type: DataType,
|
||||
gen_name: &str,
|
||||
gen_id: i32,
|
||||
definition: &GeneratedColumnDefinition,
|
||||
) -> GeneratedColumnBindingSnapshot {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
let gen_field = Field::new(gen_name, DataType::Int32, true).with_metadata(
|
||||
[(
|
||||
GENERATED_COLUMN_METADATA_KEY.to_string(),
|
||||
definition.to_metadata_json().unwrap(),
|
||||
)]
|
||||
.into(),
|
||||
);
|
||||
GeneratedColumnBindingSnapshot::try_new(
|
||||
version,
|
||||
vec![
|
||||
Arc::new(Field::new(ordinary_name, ordinary_type, true)),
|
||||
Arc::new(gen_field),
|
||||
],
|
||||
vec![ordinary_id, gen_id],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn assert_snapshot_definition_invalid_input(err: &Error, label: &str) {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"{label}: expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let rendered = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: diagnostic leaked metadata wire key: {rendered}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_returns_complete_and_incomplete() {
|
||||
use crate::function::GeneratedColumnStatus;
|
||||
|
||||
let complete = field_arg_definition(11, 3, DataType::Utf8, 4, 4);
|
||||
let snapshot =
|
||||
snapshot_with_definition(9, "text", 3, DataType::Utf8, "gen_out", 11, &complete);
|
||||
// High-level seam: name lookup + decode + same-snapshot field-arg check.
|
||||
// Callers keep using snapshot.version() for the FF-011 source pin.
|
||||
let got = snapshot.generated_column_definition("gen_out").unwrap();
|
||||
assert_eq!(got, complete);
|
||||
assert_eq!(got.status(), GeneratedColumnStatus::Complete);
|
||||
assert_eq!(snapshot.version(), 9);
|
||||
|
||||
let incomplete = field_arg_definition(11, 3, DataType::Utf8, 5, 2);
|
||||
let snapshot =
|
||||
snapshot_with_definition(10, "text", 3, DataType::Utf8, "gen_out", 11, &incomplete);
|
||||
let got = snapshot.generated_column_definition("gen_out").unwrap();
|
||||
assert_eq!(got, incomplete);
|
||||
assert_eq!(got.status(), GeneratedColumnStatus::Incomplete);
|
||||
|
||||
// Literal-only definitions remain valid (no field args to re-check).
|
||||
let literal = GeneratedColumnDefinition::try_new(13, status_sample_call(), 2, 2).unwrap();
|
||||
let snapshot = snapshot_with_definition(1, "text", 3, DataType::Utf8, "a.b", 13, &literal);
|
||||
assert_eq!(
|
||||
snapshot.generated_column_definition("a.b").unwrap(),
|
||||
literal
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_rejects_empty_missing_ordinary_and_case() {
|
||||
let definition = field_arg_definition(11, 3, DataType::Utf8, 1, 1);
|
||||
let snapshot =
|
||||
snapshot_with_definition(1, "ordinary", 3, DataType::Utf8, "gen_out", 11, &definition);
|
||||
|
||||
for name in ["", "missing", "Gen_Out", "GEN_OUT", "ordinary", "gen.out"] {
|
||||
let err = snapshot.generated_column_definition(name).unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, name);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_fail_closed_for_invalid_metadata() {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
|
||||
let field_id = 11i32;
|
||||
let valid = definition_json(field_id, 2, 2);
|
||||
let mut mismatched: serde_json::Value = serde_json::from_str(&valid).unwrap();
|
||||
mismatched["output_field_id"] = serde_json::json!(field_id + 1);
|
||||
|
||||
const MARKER: &str = "SENSITIVE_SNAPSHOT_DEF_MARKER_c8e4_1a90";
|
||||
let malformed = format!(
|
||||
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(malformed.contains(MARKER));
|
||||
|
||||
for (label, raw) in [
|
||||
("output_field_id mismatch", mismatched.to_string()),
|
||||
("malformed function_call", malformed.clone()),
|
||||
] {
|
||||
let field = Field::new("gen_out", DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), raw.clone())].into());
|
||||
let snapshot =
|
||||
GeneratedColumnBindingSnapshot::try_new(1, vec![Arc::new(field)], vec![field_id])
|
||||
.unwrap();
|
||||
let err = snapshot.generated_column_definition("gen_out").unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, label);
|
||||
let rendered = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!rendered.contains(MARKER) && !rendered.contains(&raw),
|
||||
"{label}: must not echo raw metadata: {rendered}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn snapshot_generated_column_definition_validates_field_args_against_same_snapshot() {
|
||||
// Missing stable input identity: fixture constructs cleanly; projection fails.
|
||||
let missing = field_arg_definition(11, 99_999, DataType::Utf8, 3, 3);
|
||||
let snapshot =
|
||||
snapshot_with_definition(2, "text", 3, DataType::Utf8, "gen_out", 11, &missing);
|
||||
let err = snapshot.generated_column_definition("gen_out").unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, "missing stored input field id");
|
||||
|
||||
// Type drift: stored argument type matches FunctionCall construction, not
|
||||
// the snapshot field at that id.
|
||||
let mistyped = field_arg_definition(11, 3, DataType::Int32, 4, 4);
|
||||
let snapshot =
|
||||
snapshot_with_definition(3, "text", 3, DataType::Utf8, "gen_out", 11, &mistyped);
|
||||
assert_eq!(
|
||||
snapshot.field("text").unwrap().field().data_type(),
|
||||
&DataType::Utf8
|
||||
);
|
||||
let err = snapshot.generated_column_definition("gen_out").unwrap_err();
|
||||
assert_snapshot_definition_invalid_input(&err, "stored input Arrow type mismatch");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable ChangeGeneratedColumnJobSpec change-generated-column Job
|
||||
//! operation input (FF-011).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not look up catalogs or
|
||||
//! tables, execute Jobs, stage artifacts, call Lance, derive candidate
|
||||
//! definitions, or mutate epochs.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{Function, FunctionCall, GeneratedColumnDefinition, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for changing a generated column (format
|
||||
/// version 1).
|
||||
///
|
||||
/// Semantic fields are exactly the expected [`GeneratedColumnDefinition`] CAS
|
||||
/// precondition and the new [`FunctionCall`]. Wire keys are exactly
|
||||
/// `format_version`, `expected_generated_column_definition`, and
|
||||
/// `new_function_call`.
|
||||
///
|
||||
/// Construction via [`Self::try_new`] validates only the new call against the
|
||||
/// new catalog [`Function`]. The expected definition is an opaque exact CAS
|
||||
/// precondition and is not validated against an old Function handle.
|
||||
/// Structural deserialize does not validate the new call either; execution
|
||||
/// consumers must call [`Self::validate_against`].
|
||||
///
|
||||
/// Both complete and incomplete expected definitions are accepted. Same-call
|
||||
/// change and new Functions whose output type or nullability differs from the
|
||||
/// old Function are valid. Status and output-type equality are not constructor
|
||||
/// or wire restrictions.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct ChangeGeneratedColumnJobSpec {
|
||||
expected_generated_column_definition: GeneratedColumnDefinition,
|
||||
new_function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl ChangeGeneratedColumnJobSpec {
|
||||
/// Create a change-generated-column Job operation input.
|
||||
///
|
||||
/// Requires [`FunctionCall::validate_against`] to succeed for
|
||||
/// `new_function_call` and `new_function` before returning (exact Function
|
||||
/// ID, parameter name/order, argument count, and Arrow type equality).
|
||||
///
|
||||
/// The `expected_definition` is stored as an opaque exact CAS
|
||||
/// precondition. Its nested call is not validated against any Function.
|
||||
pub fn try_new(
|
||||
expected_definition: GeneratedColumnDefinition,
|
||||
new_function: &Function,
|
||||
new_function_call: FunctionCall,
|
||||
) -> Result<Self> {
|
||||
new_function_call.validate_against(new_function)?;
|
||||
Ok(Self {
|
||||
expected_generated_column_definition: expected_definition,
|
||||
new_function_call,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Expected generated-column definition used as an exact CAS precondition.
|
||||
pub fn expected_generated_column_definition(&self) -> &GeneratedColumnDefinition {
|
||||
&self.expected_generated_column_definition
|
||||
}
|
||||
|
||||
/// New function call to apply.
|
||||
pub fn new_function_call(&self) -> &FunctionCall {
|
||||
&self.new_function_call
|
||||
}
|
||||
|
||||
/// Validate the new call against a catalog [`Function`].
|
||||
///
|
||||
/// Structural decode does not perform this check. Execution consumers must
|
||||
/// call this before using the new call. The expected definition remains an
|
||||
/// opaque CAS precondition and is not validated here.
|
||||
pub fn validate_against(&self, new_function: &Function) -> Result<()> {
|
||||
self.new_function_call.validate_against(new_function)
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> ChangeGeneratedColumnJobSpecWire {
|
||||
ChangeGeneratedColumnJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
expected_generated_column_definition: self.expected_generated_column_definition.clone(),
|
||||
new_function_call: self.new_function_call.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: ChangeGeneratedColumnJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported ChangeGeneratedColumnJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
Ok(Self {
|
||||
expected_generated_column_definition: wire.expected_generated_column_definition,
|
||||
new_function_call: wire.new_function_call,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for ChangeGeneratedColumnJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let expected = &self.expected_generated_column_definition;
|
||||
let old_call = expected.function_call();
|
||||
let new_call = &self.new_function_call;
|
||||
let old_field_ids: Vec<_> = old_call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
let new_field_ids: Vec<_> = new_call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
f.debug_struct("ChangeGeneratedColumnJobSpec")
|
||||
.field("output_field_id", &expected.output_field_id())
|
||||
.field("old_function_id", &old_call.function_id().as_str())
|
||||
.field("new_function_id", &new_call.function_id().as_str())
|
||||
.field("dependency_epoch", &expected.dependency_epoch())
|
||||
.field("materialized_epoch", &expected.materialized_epoch())
|
||||
.field("old_argument_count", &old_call.arguments().len())
|
||||
.field("new_argument_count", &new_call.arguments().len())
|
||||
.field("old_field_ids", &old_field_ids)
|
||||
.field("new_field_ids", &new_field_ids)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested GeneratedColumnDefinition / FunctionCall may
|
||||
// carry typed literal payloads on the trusted change-generated-column wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct ChangeGeneratedColumnJobSpecWire {
|
||||
format_version: u32,
|
||||
expected_generated_column_definition: GeneratedColumnDefinition,
|
||||
new_function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl Serialize for ChangeGeneratedColumnJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for ChangeGeneratedColumnJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = ChangeGeneratedColumnJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable CreateGeneratedColumnJobSpec create-generated-column Job operation
|
||||
//! input (FF-009).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not allocate output fields,
|
||||
//! construct [`super::GeneratedColumnDefinition`], mutate tables, or execute
|
||||
//! Jobs.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{Function, FunctionCall, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for creating a generated column (format
|
||||
/// version 1).
|
||||
///
|
||||
/// Semantic fields are exactly `column_name` and [`FunctionCall`]. Wire keys
|
||||
/// are exactly `format_version`, `column_name`, and `function_call`.
|
||||
///
|
||||
/// Construction via [`Self::try_new`] validates the call against a catalog
|
||||
/// [`Function`]. Structural deserialize does not; execution consumers must
|
||||
/// call [`Self::validate_against`].
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct CreateGeneratedColumnJobSpec {
|
||||
column_name: String,
|
||||
function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl CreateGeneratedColumnJobSpec {
|
||||
/// Create a create-generated-column Job operation input.
|
||||
///
|
||||
/// Rejects an empty `column_name`. Requires
|
||||
/// [`FunctionCall::validate_against`] to succeed for `function` before
|
||||
/// returning (exact Function ID, parameter name/order, argument count, and
|
||||
/// Arrow type equality).
|
||||
pub fn try_new(
|
||||
column_name: impl Into<String>,
|
||||
function: &Function,
|
||||
call: FunctionCall,
|
||||
) -> Result<Self> {
|
||||
let column_name = column_name.into();
|
||||
if column_name.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"CreateGeneratedColumnJobSpec column_name must be non-empty",
|
||||
));
|
||||
}
|
||||
call.validate_against(function)?;
|
||||
Ok(Self {
|
||||
column_name,
|
||||
function_call: call,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Target generated column name.
|
||||
pub fn column_name(&self) -> &str {
|
||||
&self.column_name
|
||||
}
|
||||
|
||||
/// Embedded function call.
|
||||
pub fn function_call(&self) -> &FunctionCall {
|
||||
&self.function_call
|
||||
}
|
||||
|
||||
/// Validate the embedded call against a catalog [`Function`].
|
||||
///
|
||||
/// Structural decode does not perform this check. Execution consumers must
|
||||
/// call this before using the call.
|
||||
pub fn validate_against(&self, function: &Function) -> Result<()> {
|
||||
self.function_call.validate_against(function)
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> CreateGeneratedColumnJobSpecWire {
|
||||
CreateGeneratedColumnJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
column_name: self.column_name.clone(),
|
||||
function_call: self.function_call.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: CreateGeneratedColumnJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported CreateGeneratedColumnJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
if wire.column_name.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"CreateGeneratedColumnJobSpec column_name must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
column_name: wire.column_name,
|
||||
function_call: wire.function_call,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for CreateGeneratedColumnJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let field_ids: Vec<_> = self
|
||||
.function_call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
f.debug_struct("CreateGeneratedColumnJobSpec")
|
||||
.field("column_name", &self.column_name)
|
||||
.field("function_id", &self.function_call.function_id().as_str())
|
||||
.field("argument_count", &self.function_call.arguments().len())
|
||||
.field("field_ids", &field_ids)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested FunctionCall may carry typed literal payloads on
|
||||
// the trusted create-generated-column wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct CreateGeneratedColumnJobSpecWire {
|
||||
format_version: u32,
|
||||
column_name: String,
|
||||
function_call: FunctionCall,
|
||||
}
|
||||
|
||||
impl Serialize for CreateGeneratedColumnJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for CreateGeneratedColumnJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = CreateGeneratedColumnJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,427 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable FunctionDefinition registration input (B1c / FF-007).
|
||||
//!
|
||||
//! These types are authoring/transport values only. They do not mint identity,
|
||||
//! store digests/artifacts, or execute Python.
|
||||
|
||||
use std::collections::HashSet;
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{FunctionSignature, SignatureWire, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Python implementation description for a [`FunctionDefinition`].
|
||||
///
|
||||
/// The source body is carried on the trusted registration wire but is omitted
|
||||
/// from [`Debug`] output.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct PythonFunctionDefinition {
|
||||
module: String,
|
||||
callable: String,
|
||||
source: String,
|
||||
python: String,
|
||||
packages: Vec<String>,
|
||||
}
|
||||
|
||||
impl PythonFunctionDefinition {
|
||||
/// Create a Python implementation description.
|
||||
///
|
||||
/// Rejects empty `module`, `callable`, `source`, `python`, or any empty
|
||||
/// package requirement, and rejects duplicate package requirement strings.
|
||||
pub fn try_new(
|
||||
module: impl Into<String>,
|
||||
callable: impl Into<String>,
|
||||
source: impl Into<String>,
|
||||
python: impl Into<String>,
|
||||
packages: Vec<String>,
|
||||
) -> Result<Self> {
|
||||
let module = module.into();
|
||||
let callable = callable.into();
|
||||
let source = source.into();
|
||||
let python = python.into();
|
||||
validate_python_fields(&module, &callable, &source, &python, &packages)?;
|
||||
Ok(Self {
|
||||
module,
|
||||
callable,
|
||||
source,
|
||||
python,
|
||||
packages,
|
||||
})
|
||||
}
|
||||
|
||||
/// Python module name.
|
||||
pub fn module(&self) -> &str {
|
||||
&self.module
|
||||
}
|
||||
|
||||
/// Callable name within the module.
|
||||
pub fn callable(&self) -> &str {
|
||||
&self.callable
|
||||
}
|
||||
|
||||
/// Source body submitted at the trusted registration boundary.
|
||||
pub fn source(&self) -> &str {
|
||||
&self.source
|
||||
}
|
||||
|
||||
/// Requested Python runtime version string.
|
||||
pub fn python(&self) -> &str {
|
||||
&self.python
|
||||
}
|
||||
|
||||
/// Ordered package requirements.
|
||||
pub fn packages(&self) -> &[String] {
|
||||
&self.packages
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for PythonFunctionDefinition {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("PythonFunctionDefinition")
|
||||
.field("module", &self.module)
|
||||
.field("callable", &self.callable)
|
||||
.field("source", &"<redacted>")
|
||||
.field("python", &self.python)
|
||||
.field("packages", &self.packages)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_python_fields(
|
||||
module: &str,
|
||||
callable: &str,
|
||||
source: &str,
|
||||
python: &str,
|
||||
packages: &[String],
|
||||
) -> Result<()> {
|
||||
if module.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition module must be non-empty",
|
||||
));
|
||||
}
|
||||
if callable.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition callable must be non-empty",
|
||||
));
|
||||
}
|
||||
if source.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition source must be non-empty",
|
||||
));
|
||||
}
|
||||
if python.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition python must be non-empty",
|
||||
));
|
||||
}
|
||||
let mut seen = HashSet::with_capacity(packages.len());
|
||||
for package in packages {
|
||||
if package.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition package must be non-empty",
|
||||
));
|
||||
}
|
||||
if !seen.insert(package.as_str()) {
|
||||
return Err(invalid_input(
|
||||
"PythonFunctionDefinition packages must not contain duplicates",
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Explicit capability grant attached to a [`FunctionDefinition`].
|
||||
///
|
||||
/// Secret references are carried on the trusted registration wire but are
|
||||
/// omitted from [`Debug`] output. Plaintext secret values are never part of
|
||||
/// this type.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct FunctionCapability {
|
||||
kind: FunctionCapabilityKind,
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
enum FunctionCapabilityKind {
|
||||
Network {
|
||||
origin: String,
|
||||
},
|
||||
Secret {
|
||||
reference: String,
|
||||
environment_variable: String,
|
||||
},
|
||||
}
|
||||
|
||||
impl FunctionCapability {
|
||||
/// Create a network capability for a non-empty origin.
|
||||
pub fn try_network(origin: impl Into<String>) -> Result<Self> {
|
||||
let origin = origin.into();
|
||||
if origin.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"FunctionCapability network origin must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
kind: FunctionCapabilityKind::Network { origin },
|
||||
})
|
||||
}
|
||||
|
||||
/// Create a secret capability for a non-empty reference and environment variable.
|
||||
///
|
||||
/// Errors name the fields and never echo the reference value.
|
||||
pub fn try_secret(
|
||||
reference: impl Into<String>,
|
||||
environment_variable: impl Into<String>,
|
||||
) -> Result<Self> {
|
||||
let reference = reference.into();
|
||||
let environment_variable = environment_variable.into();
|
||||
if reference.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"FunctionCapability secret reference must be non-empty",
|
||||
));
|
||||
}
|
||||
if environment_variable.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"FunctionCapability secret environment_variable must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
kind: FunctionCapabilityKind::Secret {
|
||||
reference,
|
||||
environment_variable,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
/// Network origin when this capability is a network grant.
|
||||
pub fn origin(&self) -> Option<&str> {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Network { origin } => Some(origin.as_str()),
|
||||
FunctionCapabilityKind::Secret { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Secret reference when this capability is a secret grant.
|
||||
pub fn reference(&self) -> Option<&str> {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Secret { reference, .. } => Some(reference.as_str()),
|
||||
FunctionCapabilityKind::Network { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Environment variable name when this capability is a secret grant.
|
||||
pub fn environment_variable(&self) -> Option<&str> {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Secret {
|
||||
environment_variable,
|
||||
..
|
||||
} => Some(environment_variable.as_str()),
|
||||
FunctionCapabilityKind::Network { .. } => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> CapabilityWire {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Network { origin } => CapabilityWire::Network {
|
||||
origin: origin.clone(),
|
||||
},
|
||||
FunctionCapabilityKind::Secret {
|
||||
reference,
|
||||
environment_variable,
|
||||
} => CapabilityWire::Secret {
|
||||
reference: reference.clone(),
|
||||
environment_variable: environment_variable.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: CapabilityWire) -> Result<Self> {
|
||||
match wire {
|
||||
CapabilityWire::Network { origin } => Self::try_network(origin),
|
||||
CapabilityWire::Secret {
|
||||
reference,
|
||||
environment_variable,
|
||||
} => Self::try_secret(reference, environment_variable),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for FunctionCapability {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match &self.kind {
|
||||
FunctionCapabilityKind::Network { origin } => f
|
||||
.debug_struct("FunctionCapability")
|
||||
.field("kind", &"network")
|
||||
.field("origin", origin)
|
||||
.finish(),
|
||||
FunctionCapabilityKind::Secret {
|
||||
environment_variable,
|
||||
..
|
||||
} => f
|
||||
.debug_struct("FunctionCapability")
|
||||
.field("kind", &"secret")
|
||||
.field("reference", &"<redacted>")
|
||||
.field("environment_variable", environment_variable)
|
||||
.finish(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Immutable registration input for a first-class Function (format version 1).
|
||||
///
|
||||
/// This value has no catalog identity. Source bodies and secret references are
|
||||
/// present on the trusted serde wire but omitted from [`Debug`].
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct FunctionDefinition {
|
||||
signature: FunctionSignature,
|
||||
python_definition: PythonFunctionDefinition,
|
||||
capabilities: Vec<FunctionCapability>,
|
||||
}
|
||||
|
||||
impl FunctionDefinition {
|
||||
/// Create a definition from a signature, Python implementation, and capabilities.
|
||||
///
|
||||
/// Emptiness and package uniqueness are enforced by the child constructors.
|
||||
pub fn try_new(
|
||||
signature: FunctionSignature,
|
||||
python_definition: PythonFunctionDefinition,
|
||||
capabilities: Vec<FunctionCapability>,
|
||||
) -> Result<Self> {
|
||||
Ok(Self {
|
||||
signature,
|
||||
python_definition,
|
||||
capabilities,
|
||||
})
|
||||
}
|
||||
|
||||
/// Function signature.
|
||||
pub fn signature(&self) -> &FunctionSignature {
|
||||
&self.signature
|
||||
}
|
||||
|
||||
/// Python implementation description.
|
||||
pub fn python_definition(&self) -> &PythonFunctionDefinition {
|
||||
&self.python_definition
|
||||
}
|
||||
|
||||
/// Ordered capability grants.
|
||||
pub fn capabilities(&self) -> &[FunctionCapability] {
|
||||
&self.capabilities
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> Result<FunctionDefinitionWire> {
|
||||
Ok(FunctionDefinitionWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
signature: self.signature.to_wire()?,
|
||||
implementation: ImplementationWire::Python {
|
||||
module: self.python_definition.module.clone(),
|
||||
callable: self.python_definition.callable.clone(),
|
||||
source: self.python_definition.source.clone(),
|
||||
python: self.python_definition.python.clone(),
|
||||
packages: self.python_definition.packages.clone(),
|
||||
},
|
||||
capabilities: self
|
||||
.capabilities
|
||||
.iter()
|
||||
.map(FunctionCapability::to_wire)
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
fn from_wire(wire: FunctionDefinitionWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported FunctionDefinition format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
let signature = FunctionSignature::from_wire(wire.signature)?;
|
||||
let python_definition = match wire.implementation {
|
||||
ImplementationWire::Python {
|
||||
module,
|
||||
callable,
|
||||
source,
|
||||
python,
|
||||
packages,
|
||||
} => PythonFunctionDefinition::try_new(module, callable, source, python, packages)?,
|
||||
};
|
||||
let capabilities = wire
|
||||
.capabilities
|
||||
.into_iter()
|
||||
.map(FunctionCapability::from_wire)
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
Self::try_new(signature, python_definition, capabilities)
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for FunctionDefinition {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("FunctionDefinition")
|
||||
.field("signature", &self.signature)
|
||||
.field("python_definition", &self.python_definition)
|
||||
.field("capabilities", &self.capabilities)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: wire payloads carry Python source and secret references.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct FunctionDefinitionWire {
|
||||
format_version: u32,
|
||||
signature: SignatureWire,
|
||||
implementation: ImplementationWire,
|
||||
capabilities: Vec<CapabilityWire>,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", deny_unknown_fields)]
|
||||
enum ImplementationWire {
|
||||
#[serde(rename = "python")]
|
||||
Python {
|
||||
module: String,
|
||||
callable: String,
|
||||
source: String,
|
||||
python: String,
|
||||
packages: Vec<String>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", deny_unknown_fields)]
|
||||
enum CapabilityWire {
|
||||
#[serde(rename = "network")]
|
||||
Network { origin: String },
|
||||
#[serde(rename = "secret")]
|
||||
Secret {
|
||||
reference: String,
|
||||
environment_variable: String,
|
||||
},
|
||||
}
|
||||
|
||||
impl Serialize for FunctionDefinition {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire()
|
||||
.map_err(serde::ser::Error::custom)?
|
||||
.serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for FunctionDefinition {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = FunctionDefinitionWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,315 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Pure crate-private generated-column invalidation planner (B4a).
|
||||
//!
|
||||
//! Plans column-wide dependency-epoch advances from a binding snapshot and a
|
||||
//! mutation impact. This module does not mutate tables, write metadata, or
|
||||
//! execute append/update/delete/merge paths. Native append and update consume
|
||||
//! the plan through the B4b / B4c runtime wiring.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use super::{GeneratedColumnBindingSnapshot, GeneratedColumnDefinition};
|
||||
use crate::Result;
|
||||
|
||||
/// Mutation impact considered by the crate-private invalidation planner.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum GeneratedColumnMutationImpact {
|
||||
/// Append or delete: whole-column coverage / row membership changed.
|
||||
RowSetChanged,
|
||||
/// Update of the listed stable field IDs (direct and transitive dependents).
|
||||
///
|
||||
/// Native update (B4c) constructs this impact. Native append (B4b) only
|
||||
/// constructs [`Self::RowSetChanged`].
|
||||
UpdatedFields(BTreeSet<i32>),
|
||||
}
|
||||
|
||||
/// One planned field-metadata replacement produced by the pure planner.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct PlannedGeneratedColumnMetadataUpdate {
|
||||
output_field_id: i32,
|
||||
metadata_json: String,
|
||||
}
|
||||
|
||||
impl PlannedGeneratedColumnMetadataUpdate {
|
||||
/// Stable output field ID whose metadata should be replaced.
|
||||
pub fn output_field_id(&self) -> i32 {
|
||||
self.output_field_id
|
||||
}
|
||||
|
||||
/// Canonical [`GeneratedColumnDefinition::to_metadata_json`] bytes.
|
||||
pub fn metadata_json(&self) -> &str {
|
||||
&self.metadata_json
|
||||
}
|
||||
}
|
||||
|
||||
/// Plan generated-column metadata replacements for `impact`.
|
||||
///
|
||||
/// Planning is pure: `snapshot` is never mutated. Every present
|
||||
/// `lancedb::generated_column` value is decoded and every decoded call's field
|
||||
/// arguments are validated against `snapshot` before impact is calculated.
|
||||
/// Decode, missing-field, type-mismatch, serialization, or overflow errors
|
||||
/// return no plan.
|
||||
///
|
||||
/// Impacted definitions advance `dependency_epoch` exactly once (checked
|
||||
/// arithmetic) while preserving `materialized_epoch`, output identity, and the
|
||||
/// embedded [`super::FunctionCall`]. Replacements are returned in snapshot
|
||||
/// schema order.
|
||||
pub fn plan_generated_column_invalidation(
|
||||
snapshot: &GeneratedColumnBindingSnapshot,
|
||||
impact: &GeneratedColumnMutationImpact,
|
||||
) -> Result<Vec<PlannedGeneratedColumnMetadataUpdate>> {
|
||||
let definitions = decode_and_validate_generated_columns(snapshot)?;
|
||||
let impacted = compute_impacted_output_ids(&definitions, impact);
|
||||
|
||||
let mut plan = Vec::new();
|
||||
for (output_field_id, definition) in &definitions {
|
||||
if !impacted.contains(output_field_id) {
|
||||
continue;
|
||||
}
|
||||
let mut next = definition.clone();
|
||||
next.invalidate()?;
|
||||
let metadata_json = next.to_metadata_json()?;
|
||||
plan.push(PlannedGeneratedColumnMetadataUpdate {
|
||||
output_field_id: *output_field_id,
|
||||
metadata_json,
|
||||
});
|
||||
}
|
||||
Ok(plan)
|
||||
}
|
||||
|
||||
/// Decode every present generated-column definition in schema order and
|
||||
/// validate field arguments against the same snapshot.
|
||||
fn decode_and_validate_generated_columns(
|
||||
snapshot: &GeneratedColumnBindingSnapshot,
|
||||
) -> Result<Vec<(i32, GeneratedColumnDefinition)>> {
|
||||
let mut definitions = Vec::new();
|
||||
for entry in snapshot.entries() {
|
||||
let Some(definition) = entry.generated_column_definition()? else {
|
||||
continue;
|
||||
};
|
||||
snapshot.validate_field_arguments(definition.function_call())?;
|
||||
definitions.push((entry.field_id(), definition));
|
||||
}
|
||||
Ok(definitions)
|
||||
}
|
||||
|
||||
/// Compute the set of impacted generated-column output field IDs.
|
||||
///
|
||||
/// `RowSetChanged` impacts every generated column. `UpdatedFields` computes a
|
||||
/// deterministic fixed point over generated output IDs: a definition is
|
||||
/// impacted when any field argument references a dirty ID, and each generated
|
||||
/// definition is added at most once so cycles terminate.
|
||||
fn compute_impacted_output_ids(
|
||||
definitions: &[(i32, GeneratedColumnDefinition)],
|
||||
impact: &GeneratedColumnMutationImpact,
|
||||
) -> BTreeSet<i32> {
|
||||
match impact {
|
||||
GeneratedColumnMutationImpact::RowSetChanged => {
|
||||
definitions.iter().map(|(id, _)| *id).collect()
|
||||
}
|
||||
GeneratedColumnMutationImpact::UpdatedFields(updated) => {
|
||||
let mut dirty = updated.clone();
|
||||
let mut impacted = BTreeSet::new();
|
||||
let mut progressed = true;
|
||||
while progressed {
|
||||
progressed = false;
|
||||
for (output_field_id, definition) in definitions {
|
||||
if impacted.contains(output_field_id) {
|
||||
continue;
|
||||
}
|
||||
let depends_on_dirty =
|
||||
definition
|
||||
.function_call()
|
||||
.arguments()
|
||||
.iter()
|
||||
.any(|(_, argument)| {
|
||||
argument
|
||||
.field_id()
|
||||
.is_some_and(|field_id| dirty.contains(&field_id))
|
||||
});
|
||||
if depends_on_dirty {
|
||||
impacted.insert(*output_field_id);
|
||||
dirty.insert(*output_field_id);
|
||||
progressed = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
impacted
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::{ArrayRef, Int32Array};
|
||||
use arrow_schema::{DataType, Field, FieldRef};
|
||||
|
||||
use super::*;
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY,
|
||||
};
|
||||
|
||||
fn int_field_function(id: &str) -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new(id).unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("upstream", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn int_field_bound_call(function: &Function, input_field_id: i32) -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
function,
|
||||
vec![(
|
||||
"upstream".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, DataType::Int32).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn definition(
|
||||
output_field_id: i32,
|
||||
call: FunctionCall,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
GeneratedColumnDefinition::try_new(
|
||||
output_field_id,
|
||||
call,
|
||||
dependency_epoch,
|
||||
materialized_epoch,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn generated_field(name: &str, def: &GeneratedColumnDefinition) -> FieldRef {
|
||||
let json = def.to_metadata_json().unwrap();
|
||||
Arc::new(
|
||||
Field::new(name, DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), json)].into()),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cyclic_dependency_fixed_point_impacts_each_definition_at_most_once() {
|
||||
// A <-> B cycle. Seeding either side must terminate and advance each
|
||||
// impacted definition exactly once. This proves planner termination; it
|
||||
// is not a public cyclic-dependency creation guarantee.
|
||||
let a_id = 60;
|
||||
let b_id = 70;
|
||||
let fn_a = int_field_function("fn.exact.b4a.cycle.a");
|
||||
let fn_b = int_field_function("fn.exact.b4a.cycle.b");
|
||||
let a = definition(a_id, int_field_bound_call(&fn_a, b_id), 1, 1);
|
||||
let b = definition(b_id, int_field_bound_call(&fn_b, a_id), 2, 2);
|
||||
let snap = GeneratedColumnBindingSnapshot::try_new(
|
||||
11,
|
||||
vec![generated_field("gen_a", &a), generated_field("gen_b", &b)],
|
||||
vec![a_id, b_id],
|
||||
)
|
||||
.unwrap();
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([a_id])),
|
||||
)
|
||||
.expect("cyclic fixed point must terminate");
|
||||
assert_eq!(snap, before);
|
||||
assert_eq!(plan.len(), 2);
|
||||
assert_eq!(plan[0].output_field_id(), a_id);
|
||||
assert_eq!(plan[1].output_field_id(), b_id);
|
||||
|
||||
let decoded_a =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[0].metadata_json(), a_id).unwrap();
|
||||
let decoded_b =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[1].metadata_json(), b_id).unwrap();
|
||||
assert_eq!(decoded_a.dependency_epoch(), 2);
|
||||
assert_eq!(decoded_a.materialized_epoch(), 1);
|
||||
assert_eq!(decoded_b.dependency_epoch(), 3);
|
||||
assert_eq!(decoded_b.materialized_epoch(), 2);
|
||||
assert_eq!(decoded_a.function_call(), a.function_call());
|
||||
assert_eq!(decoded_b.function_call(), b.function_call());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn row_set_change_with_cycle_still_invalidates_each_column_once() {
|
||||
let a_id = 61;
|
||||
let b_id = 71;
|
||||
let fn_a = int_field_function("fn.exact.b4a.cycle.row.a");
|
||||
let fn_b = int_field_function("fn.exact.b4a.cycle.row.b");
|
||||
let a = definition(a_id, int_field_bound_call(&fn_a, b_id), 5, 5);
|
||||
let b = definition(b_id, int_field_bound_call(&fn_b, a_id), 8, 8);
|
||||
let snap = GeneratedColumnBindingSnapshot::try_new(
|
||||
12,
|
||||
vec![generated_field("gen_b", &b), generated_field("gen_a", &a)],
|
||||
vec![b_id, a_id],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::RowSetChanged,
|
||||
)
|
||||
.expect("row-set change over a cycle must plan once per column");
|
||||
assert_eq!(plan.len(), 2);
|
||||
assert_eq!(plan[0].output_field_id(), b_id);
|
||||
assert_eq!(plan[1].output_field_id(), a_id);
|
||||
let decoded_b =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[0].metadata_json(), b_id).unwrap();
|
||||
let decoded_a =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[1].metadata_json(), a_id).unwrap();
|
||||
assert_eq!(decoded_b.dependency_epoch(), 9);
|
||||
assert_eq!(decoded_a.dependency_epoch(), 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn literal_only_is_ignored_by_updated_fields_even_with_empty_seed() {
|
||||
let literal_fn = Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.cycle.literal").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("constant", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let literal_id = 80;
|
||||
let literal = definition(
|
||||
literal_id,
|
||||
FunctionCall::try_new(
|
||||
&literal_fn,
|
||||
vec![(
|
||||
"constant".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap(),
|
||||
3,
|
||||
3,
|
||||
);
|
||||
let snap = GeneratedColumnBindingSnapshot::try_new(
|
||||
13,
|
||||
vec![generated_field("gen_literal", &literal)],
|
||||
vec![literal_id],
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::new()),
|
||||
)
|
||||
.expect("empty UpdatedFields must succeed");
|
||||
assert!(plan.is_empty());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,483 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Contract tests for the crate-private generated-column invalidation planner (B4a).
|
||||
//!
|
||||
//! These tests pin the pure planning surface implemented by
|
||||
//! [`super::plan_generated_column_invalidation`]. No runtime append/update/delete
|
||||
//! path is exercised.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::{ArrayRef, Int32Array};
|
||||
use arrow_schema::{DataType, Field, FieldRef};
|
||||
|
||||
use super::plan_generated_column_invalidation::{
|
||||
GeneratedColumnMutationImpact, PlannedGeneratedColumnMetadataUpdate,
|
||||
plan_generated_column_invalidation,
|
||||
};
|
||||
use super::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnBindingSnapshot,
|
||||
GeneratedColumnDefinition,
|
||||
};
|
||||
use crate::Error;
|
||||
|
||||
fn utf8_field_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.utf8").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("payload", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn literal_only_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.literal").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("constant", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn int_field_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.b4a.int").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("upstream", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn field_bound_call(input_field_id: i32) -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
&utf8_field_function(),
|
||||
vec![(
|
||||
"payload".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, DataType::Utf8).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn literal_only_call() -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
&literal_only_function(),
|
||||
vec![(
|
||||
"constant".to_string(),
|
||||
FunctionArgument::try_literal(Arc::new(Int32Array::from(vec![Some(7)])) as ArrayRef)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn int_field_bound_call(input_field_id: i32) -> FunctionCall {
|
||||
FunctionCall::try_new(
|
||||
&int_field_function(),
|
||||
vec![(
|
||||
"upstream".to_string(),
|
||||
FunctionArgument::try_field(input_field_id, DataType::Int32).unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn definition(
|
||||
output_field_id: i32,
|
||||
call: FunctionCall,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
GeneratedColumnDefinition::try_new(output_field_id, call, dependency_epoch, materialized_epoch)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn ordinary_field(name: &str, data_type: DataType) -> FieldRef {
|
||||
Arc::new(Field::new(name, data_type, true))
|
||||
}
|
||||
|
||||
fn generated_field(name: &str, def: &GeneratedColumnDefinition) -> FieldRef {
|
||||
let json = def.to_metadata_json().unwrap();
|
||||
Arc::new(
|
||||
Field::new(name, DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), json)].into()),
|
||||
)
|
||||
}
|
||||
|
||||
fn generated_field_with_raw_metadata(name: &str, raw: &str) -> FieldRef {
|
||||
Arc::new(
|
||||
Field::new(name, DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), raw.to_string())].into()),
|
||||
)
|
||||
}
|
||||
|
||||
fn snapshot(
|
||||
version: u64,
|
||||
fields: Vec<FieldRef>,
|
||||
field_ids: Vec<i32>,
|
||||
) -> GeneratedColumnBindingSnapshot {
|
||||
GeneratedColumnBindingSnapshot::try_new(version, fields, field_ids).unwrap()
|
||||
}
|
||||
|
||||
fn expected_invalidated(def: &GeneratedColumnDefinition) -> GeneratedColumnDefinition {
|
||||
let mut next = def.clone();
|
||||
next.invalidate().unwrap();
|
||||
next
|
||||
}
|
||||
|
||||
fn assert_planned_definition(
|
||||
update: &PlannedGeneratedColumnMetadataUpdate,
|
||||
expected: &GeneratedColumnDefinition,
|
||||
) {
|
||||
assert_eq!(update.output_field_id(), expected.output_field_id());
|
||||
let decoded = GeneratedColumnDefinition::from_metadata_json(
|
||||
update.metadata_json(),
|
||||
expected.output_field_id(),
|
||||
)
|
||||
.expect("planned metadata must decode");
|
||||
assert_eq!(&decoded, expected);
|
||||
assert_eq!(
|
||||
update.metadata_json(),
|
||||
expected.to_metadata_json().unwrap(),
|
||||
"planned metadata JSON must be canonical"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn no_generated_columns_returns_empty_plan() {
|
||||
let snap = snapshot(
|
||||
1,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
],
|
||||
vec![1, 2],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect("planner must succeed when no generated columns are present");
|
||||
assert!(plan.is_empty());
|
||||
assert_eq!(snap, before, "planner must not mutate the binding snapshot");
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([1])),
|
||||
)
|
||||
.expect("field update with no generated columns must succeed");
|
||||
assert!(plan.is_empty());
|
||||
assert_eq!(snap, before, "planner must not mutate the binding snapshot");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn row_set_change_invalidates_field_bound_and_literal_only_exactly_once() {
|
||||
let text_id = 10;
|
||||
let field_bound_id = 20;
|
||||
let literal_id = 30;
|
||||
let field_bound = definition(field_bound_id, field_bound_call(text_id), 3, 3);
|
||||
let literal_only = definition(literal_id, literal_only_call(), 4, 4);
|
||||
let snap = snapshot(
|
||||
2,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
generated_field("gen_field", &field_bound),
|
||||
generated_field("gen_literal", &literal_only),
|
||||
],
|
||||
vec![text_id, field_bound_id, literal_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect("row-set change must plan invalidation");
|
||||
assert_eq!(snap, before, "planner must not mutate the binding snapshot");
|
||||
assert_eq!(
|
||||
plan.len(),
|
||||
2,
|
||||
"each generated column invalidates exactly once"
|
||||
);
|
||||
assert_eq!(plan[0].output_field_id(), field_bound_id);
|
||||
assert_eq!(plan[1].output_field_id(), literal_id);
|
||||
assert_planned_definition(&plan[0], &expected_invalidated(&field_bound));
|
||||
assert_planned_definition(&plan[1], &expected_invalidated(&literal_only));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn already_incomplete_advances_dependency_epoch_and_preserves_materialized_epoch() {
|
||||
let text_id = 11;
|
||||
let gen_id = 21;
|
||||
let incomplete = definition(gen_id, field_bound_call(text_id), 9, 2);
|
||||
assert_eq!(incomplete.dependency_epoch(), 9);
|
||||
assert_eq!(incomplete.materialized_epoch(), 2);
|
||||
let snap = snapshot(
|
||||
3,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
generated_field("gen_incomplete", &incomplete),
|
||||
],
|
||||
vec![text_id, gen_id],
|
||||
);
|
||||
|
||||
let plan =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect("incomplete definition must still advance");
|
||||
assert_eq!(plan.len(), 1);
|
||||
let decoded =
|
||||
GeneratedColumnDefinition::from_metadata_json(plan[0].metadata_json(), gen_id).unwrap();
|
||||
assert_eq!(decoded.dependency_epoch(), 10);
|
||||
assert_eq!(decoded.materialized_epoch(), 2);
|
||||
assert_eq!(
|
||||
decoded.function_call(),
|
||||
incomplete.function_call(),
|
||||
"invalidation must preserve the embedded function call"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn direct_field_update_invalidates_only_dependent_generated_column() {
|
||||
let text_id = 12;
|
||||
let score_id = 13;
|
||||
let dependent_id = 22;
|
||||
let unrelated_gen_id = 23;
|
||||
let dependent = definition(dependent_id, field_bound_call(text_id), 5, 5);
|
||||
let unrelated_gen = definition(unrelated_gen_id, literal_only_call(), 6, 6);
|
||||
let snap = snapshot(
|
||||
4,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field("gen_dependent", &dependent),
|
||||
generated_field("gen_unrelated", &unrelated_gen),
|
||||
],
|
||||
vec![text_id, score_id, dependent_id, unrelated_gen_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([text_id])),
|
||||
)
|
||||
.expect("dependent update must plan a single invalidation");
|
||||
assert_eq!(snap, before);
|
||||
assert_eq!(plan.len(), 1);
|
||||
assert_planned_definition(&plan[0], &expected_invalidated(&dependent));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unrelated_field_update_returns_empty_plan() {
|
||||
let text_id = 14;
|
||||
let score_id = 15;
|
||||
let gen_id = 24;
|
||||
let dependent = definition(gen_id, field_bound_call(text_id), 2, 2);
|
||||
let snap = snapshot(
|
||||
5,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field("gen_text", &dependent),
|
||||
],
|
||||
vec![text_id, score_id, gen_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([score_id])),
|
||||
)
|
||||
.expect("unrelated update must not invent invalidation");
|
||||
assert!(plan.is_empty());
|
||||
assert_eq!(snap, before);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transitive_dependency_propagation_follows_snapshot_order() {
|
||||
// A (ordinary) -> B (generated) -> C (generated). Update A invalidates B and C.
|
||||
let a_id = 30;
|
||||
let b_id = 40;
|
||||
let c_id = 50;
|
||||
let b = definition(b_id, field_bound_call(a_id), 1, 1);
|
||||
let c = definition(c_id, int_field_bound_call(b_id), 1, 1);
|
||||
// Schema order places C before B so the plan must follow snapshot order, not
|
||||
// dependency discovery order.
|
||||
let snap = snapshot(
|
||||
6,
|
||||
vec![
|
||||
ordinary_field("a", DataType::Utf8),
|
||||
generated_field("gen_c", &c),
|
||||
generated_field("gen_b", &b),
|
||||
],
|
||||
vec![a_id, c_id, b_id],
|
||||
);
|
||||
let before = snap.clone();
|
||||
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([a_id])),
|
||||
)
|
||||
.expect("transitive dependents must invalidate");
|
||||
assert_eq!(snap, before);
|
||||
assert_eq!(plan.len(), 2);
|
||||
assert_eq!(plan[0].output_field_id(), c_id);
|
||||
assert_eq!(plan[1].output_field_id(), b_id);
|
||||
assert_planned_definition(&plan[0], &expected_invalidated(&c));
|
||||
assert_planned_definition(&plan[1], &expected_invalidated(&b));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_metadata_fails_closed_for_unrelated_update_without_echoing_payload() {
|
||||
const MARKER: &str = "SENSITIVE_B4A_METADATA_MARKER_7c91_e2aa";
|
||||
let text_id = 16;
|
||||
let score_id = 17;
|
||||
let bad_id = 25;
|
||||
let raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":{bad_id},"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(raw.contains(MARKER));
|
||||
let snap = snapshot(
|
||||
7,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field_with_raw_metadata("gen_bad", &raw),
|
||||
],
|
||||
vec![text_id, score_id, bad_id],
|
||||
);
|
||||
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([score_id])),
|
||||
)
|
||||
.expect_err("malformed metadata must fail closed even for an unrelated update");
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let text = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!text.contains(MARKER),
|
||||
"diagnostics must not echo raw metadata marker: {text}"
|
||||
);
|
||||
assert!(
|
||||
!text.contains(&raw),
|
||||
"diagnostics must not echo raw metadata payload: {text}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_input_field_id_fails_closed() {
|
||||
let missing_input_id = 99;
|
||||
let gen_id = 26;
|
||||
let orphan = definition(gen_id, field_bound_call(missing_input_id), 1, 1);
|
||||
let snap = snapshot(
|
||||
8,
|
||||
vec![
|
||||
ordinary_field("score", DataType::Int32),
|
||||
generated_field("gen_orphan", &orphan),
|
||||
],
|
||||
vec![18, gen_id],
|
||||
);
|
||||
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([18])),
|
||||
)
|
||||
.expect_err("missing stable input field id must fail closed");
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains("99") || message.contains("missing"),
|
||||
"diagnostic should identify the missing field id: {message}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn field_type_mismatch_fails_closed() {
|
||||
let text_id = 19;
|
||||
let gen_id = 27;
|
||||
// Definition claims Utf8 for field 19, but the snapshot entry is Int32.
|
||||
let mismatched = definition(gen_id, field_bound_call(text_id), 1, 1);
|
||||
let snap = snapshot(
|
||||
9,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Int32),
|
||||
generated_field("gen_mismatch", &mismatched),
|
||||
],
|
||||
vec![text_id, gen_id],
|
||||
);
|
||||
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([text_id])),
|
||||
)
|
||||
.expect_err("field type mismatch must fail closed");
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains("mismatch")
|
||||
|| (message.contains("Utf8") && message.contains("Int32"))
|
||||
|| message.contains(&text_id.to_string()),
|
||||
"diagnostic should identify the type mismatch: {message}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn epoch_overflow_fails_atomically_with_stable_sanitized_diagnostic() {
|
||||
let text_id = 31;
|
||||
let overflow_id = 41;
|
||||
let other_id = 42;
|
||||
let at_max = definition(overflow_id, field_bound_call(text_id), u64::MAX, u64::MAX);
|
||||
let other = definition(other_id, literal_only_call(), 1, 1);
|
||||
let snap = snapshot(
|
||||
10,
|
||||
vec![
|
||||
ordinary_field("text", DataType::Utf8),
|
||||
generated_field("gen_max", &at_max),
|
||||
generated_field("gen_other", &other),
|
||||
],
|
||||
vec![text_id, overflow_id, other_id],
|
||||
);
|
||||
|
||||
// Row-set change impacts every generated column, including the overflowed one.
|
||||
let err =
|
||||
plan_generated_column_invalidation(&snap, &GeneratedColumnMutationImpact::RowSetChanged)
|
||||
.expect_err("dependency_epoch overflow must fail closed");
|
||||
match err {
|
||||
Error::InvalidInput { message } => {
|
||||
assert_eq!(
|
||||
message, "dependency_epoch overflow",
|
||||
"overflow must use the existing sanitized InvalidInput diagnostic"
|
||||
);
|
||||
}
|
||||
other => panic!("expected InvalidInput overflow, got {other:?}"),
|
||||
}
|
||||
|
||||
// Direct update that impacts only the overflowed definition must also fail
|
||||
// atomically and must not return a partial plan for sibling columns.
|
||||
let err = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(BTreeSet::from([text_id])),
|
||||
)
|
||||
.expect_err("impacted overflow must fail with no partial plan");
|
||||
match err {
|
||||
Error::InvalidInput { message } => {
|
||||
assert_eq!(message, "dependency_epoch overflow");
|
||||
}
|
||||
other => panic!("expected InvalidInput overflow, got {other:?}"),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable RefreshGeneratedColumnJobSpec refresh-generated-column Job
|
||||
//! operation input (FF-010).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not look up catalogs or
|
||||
//! tables, execute Jobs, stage artifacts, call Lance, or mutate epochs.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{Function, GeneratedColumnDefinition, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for refreshing a generated column (format
|
||||
/// version 1).
|
||||
///
|
||||
/// Semantic field is exactly the nested [`GeneratedColumnDefinition`]. Wire
|
||||
/// keys are exactly `format_version` and `generated_column_definition`.
|
||||
///
|
||||
/// Construction via [`Self::try_new`] validates the nested call against a
|
||||
/// catalog [`Function`]. Structural deserialize does not; execution consumers
|
||||
/// must call [`Self::validate_against`], and later compare the full nested
|
||||
/// definition to current field metadata in the pinned snapshot.
|
||||
///
|
||||
/// Both complete and incomplete definitions are accepted. Status is not a
|
||||
/// constructor or wire restriction.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct RefreshGeneratedColumnJobSpec {
|
||||
generated_column_definition: GeneratedColumnDefinition,
|
||||
}
|
||||
|
||||
impl RefreshGeneratedColumnJobSpec {
|
||||
/// Create a refresh-generated-column Job operation input.
|
||||
///
|
||||
/// Requires [`crate::function::FunctionCall::validate_against`] to succeed
|
||||
/// for the nested call and `function` before returning (exact Function ID,
|
||||
/// parameter name/order, argument count, and Arrow type equality).
|
||||
pub fn try_new(
|
||||
function: &Function,
|
||||
generated_column_definition: GeneratedColumnDefinition,
|
||||
) -> Result<Self> {
|
||||
generated_column_definition
|
||||
.function_call()
|
||||
.validate_against(function)?;
|
||||
Ok(Self {
|
||||
generated_column_definition,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Nested generated-column definition to refresh.
|
||||
pub fn generated_column_definition(&self) -> &GeneratedColumnDefinition {
|
||||
&self.generated_column_definition
|
||||
}
|
||||
|
||||
/// Validate the nested call against a catalog [`Function`].
|
||||
///
|
||||
/// Structural decode does not perform this check. Execution consumers must
|
||||
/// call this before using the call.
|
||||
pub fn validate_against(&self, function: &Function) -> Result<()> {
|
||||
self.generated_column_definition
|
||||
.function_call()
|
||||
.validate_against(function)
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> RefreshGeneratedColumnJobSpecWire {
|
||||
RefreshGeneratedColumnJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
generated_column_definition: self.generated_column_definition.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: RefreshGeneratedColumnJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported RefreshGeneratedColumnJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
Ok(Self {
|
||||
generated_column_definition: wire.generated_column_definition,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for RefreshGeneratedColumnJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
let definition = &self.generated_column_definition;
|
||||
let call = definition.function_call();
|
||||
let field_ids: Vec<_> = call
|
||||
.arguments()
|
||||
.iter()
|
||||
.filter_map(|(_, argument)| argument.field_id())
|
||||
.collect();
|
||||
f.debug_struct("RefreshGeneratedColumnJobSpec")
|
||||
.field("output_field_id", &definition.output_field_id())
|
||||
.field("function_id", &call.function_id().as_str())
|
||||
.field("dependency_epoch", &definition.dependency_epoch())
|
||||
.field("materialized_epoch", &definition.materialized_epoch())
|
||||
.field("argument_count", &call.arguments().len())
|
||||
.field("field_ids", &field_ids)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested GeneratedColumnDefinition / FunctionCall may
|
||||
// carry typed literal payloads on the trusted refresh wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct RefreshGeneratedColumnJobSpecWire {
|
||||
format_version: u32,
|
||||
generated_column_definition: GeneratedColumnDefinition,
|
||||
}
|
||||
|
||||
impl Serialize for RefreshGeneratedColumnJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for RefreshGeneratedColumnJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = RefreshGeneratedColumnJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Immutable RegisterFunctionJobSpec registration Job operation input (B1d / FF-008).
|
||||
//!
|
||||
//! This type is Job operation input only. It does not execute registration,
|
||||
//! upsert into a catalog, mint identity, or manage Job lifecycle.
|
||||
|
||||
use std::fmt;
|
||||
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
|
||||
use super::{FunctionDefinition, FunctionId, invalid_input};
|
||||
use crate::Result;
|
||||
|
||||
const FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
/// Immutable Job operation input for registering a first-class Function
|
||||
/// (format version 1).
|
||||
///
|
||||
/// `expected_current_function_id` is a precondition only:
|
||||
/// - [`None`] means create-if-absent (no current Function is expected).
|
||||
/// - [`Some`] with an exact opaque [`FunctionId`] means conditional replace of
|
||||
/// that current Function.
|
||||
///
|
||||
/// This type does not perform catalog execution or upsert.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct RegisterFunctionJobSpec {
|
||||
name: String,
|
||||
definition: FunctionDefinition,
|
||||
expected_current_function_id: Option<FunctionId>,
|
||||
}
|
||||
|
||||
impl RegisterFunctionJobSpec {
|
||||
/// Create a registration Job operation input.
|
||||
///
|
||||
/// Rejects an empty `name`. Nested definition validation is enforced by
|
||||
/// [`FunctionDefinition`]. When `expected_current_function_id` is
|
||||
/// [`Some`], emptiness is enforced by [`FunctionId::try_new`].
|
||||
///
|
||||
/// - `expected_current_function_id = None`: create-if-absent.
|
||||
/// - `expected_current_function_id = Some(id)`: conditional replace of the
|
||||
/// Function with that exact opaque id.
|
||||
pub fn try_new(
|
||||
name: impl Into<String>,
|
||||
definition: FunctionDefinition,
|
||||
expected_current_function_id: Option<FunctionId>,
|
||||
) -> Result<Self> {
|
||||
let name = name.into();
|
||||
if name.is_empty() {
|
||||
return Err(invalid_input(
|
||||
"RegisterFunctionJobSpec name must be non-empty",
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
name,
|
||||
definition,
|
||||
expected_current_function_id,
|
||||
})
|
||||
}
|
||||
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Catalog Function name to register.
|
||||
pub fn name(&self) -> &str {
|
||||
&self.name
|
||||
}
|
||||
|
||||
/// Nested registration definition (exact FF-007 [`FunctionDefinition`]).
|
||||
pub fn definition(&self) -> &FunctionDefinition {
|
||||
&self.definition
|
||||
}
|
||||
|
||||
/// Precondition on the current Function id.
|
||||
///
|
||||
/// [`None`] is create-if-absent. [`Some`] is conditional replace of that
|
||||
/// exact opaque id.
|
||||
pub fn expected_current_function_id(&self) -> Option<&FunctionId> {
|
||||
self.expected_current_function_id.as_ref()
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> RegisterFunctionJobSpecWire {
|
||||
RegisterFunctionJobSpecWire {
|
||||
format_version: FORMAT_VERSION_V1,
|
||||
name: self.name.clone(),
|
||||
definition: self.definition.clone(),
|
||||
expected_current_function_id: self
|
||||
.expected_current_function_id
|
||||
.as_ref()
|
||||
.map(|id| id.as_str().to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: RegisterFunctionJobSpecWire) -> Result<Self> {
|
||||
if wire.format_version != FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported RegisterFunctionJobSpec format_version {}",
|
||||
wire.format_version
|
||||
)));
|
||||
}
|
||||
let expected_current_function_id = match wire.expected_current_function_id {
|
||||
None => None,
|
||||
Some(id) => Some(FunctionId::try_new(id)?),
|
||||
};
|
||||
Self::try_new(wire.name, wire.definition, expected_current_function_id)
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for RegisterFunctionJobSpec {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("RegisterFunctionJobSpec")
|
||||
.field("name", &self.name)
|
||||
.field("definition", &self.definition)
|
||||
.field(
|
||||
"expected_current_function_id",
|
||||
&self.expected_current_function_id,
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
// Do not derive Debug: nested FunctionDefinition carries Python source and
|
||||
// secret references on the trusted registration wire.
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
struct RegisterFunctionJobSpecWire {
|
||||
format_version: u32,
|
||||
name: String,
|
||||
definition: FunctionDefinition,
|
||||
expected_current_function_id: Option<String>,
|
||||
}
|
||||
|
||||
impl Serialize for RegisterFunctionJobSpec {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for RegisterFunctionJobSpec {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = RegisterFunctionJobSpecWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Schema admission for caller-authored generated-column definition ingress.
|
||||
//!
|
||||
//! General-purpose table-schema inputs (for example `Database::create_table`
|
||||
//! and Native `add_columns` schema-bearing transforms) must not invent or
|
||||
//! mutate Job-owned `lancedb::generated_column` top-level field metadata. Only
|
||||
//! generated-column create/change/refresh Job publication may create or change
|
||||
//! that reserved key.
|
||||
//!
|
||||
//! This helper checks raw key presence on top-level fields only. It does not
|
||||
//! recurse into nested children, inspect schema-level metadata, decode the
|
||||
//! payload, look up a Function, or validate epochs.
|
||||
|
||||
use arrow_schema::Schema;
|
||||
|
||||
use super::GENERATED_COLUMN_METADATA_KEY;
|
||||
use crate::{Error, Result};
|
||||
|
||||
/// Reject a caller-authored Arrow schema that carries reserved generated-column
|
||||
/// definition metadata on any top-level field.
|
||||
///
|
||||
/// Safe to call at the start of create-table and Native add-columns paths
|
||||
/// before source consumption, namespace mutation, or HTTP.
|
||||
pub fn reject_caller_authored_generated_column_schema(schema: &Schema) -> Result<()> {
|
||||
for field in schema.fields() {
|
||||
if field.metadata().contains_key(GENERATED_COLUMN_METADATA_KEY) {
|
||||
return Err(Error::NotSupported {
|
||||
message: "generated column definitions are owned by create/change/refresh Jobs \
|
||||
and cannot be supplied through general-purpose table schema input"
|
||||
.into(),
|
||||
});
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Conditionally admit an input schema for append vs overwrite.
|
||||
///
|
||||
/// Overwrite is schema replacement and must reject reserved top-level field
|
||||
/// metadata. Append is not schema replacement: caller field metadata is
|
||||
/// discarded by cast-to-table-schema, so reserved input keys remain accepted.
|
||||
pub fn reject_caller_authored_generated_column_schema_on_overwrite(
|
||||
schema: &Schema,
|
||||
is_overwrite: bool,
|
||||
) -> Result<()> {
|
||||
if is_overwrite {
|
||||
reject_caller_authored_generated_column_schema(schema)
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
+293
-13
@@ -6,10 +6,124 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde::de::Error as DeError;
|
||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||
use tokio::sync::watch;
|
||||
use tokio::task::{AbortHandle, JoinHandle};
|
||||
|
||||
use crate::error::{Error, JobFailure, Result};
|
||||
use crate::function::Function;
|
||||
|
||||
const JOB_RESULT_FORMAT_VERSION_V1: u32 = 1;
|
||||
|
||||
fn invalid_input(message: impl Into<String>) -> Error {
|
||||
Error::InvalidInput {
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Result value produced by a completed Job (format version 1).
|
||||
///
|
||||
/// This is a non-resource transport value. It is not a Job handle, does not
|
||||
/// observe lifecycle, and does not preserve unknown wire shapes.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[non_exhaustive]
|
||||
pub enum JobResult {
|
||||
/// The Job completed without a Function result.
|
||||
None,
|
||||
/// The Job completed with a [`Function`] value.
|
||||
Function(Function),
|
||||
}
|
||||
|
||||
impl JobResult {
|
||||
/// Wire format version (always 1 for this type).
|
||||
pub fn format_version(&self) -> u32 {
|
||||
JOB_RESULT_FORMAT_VERSION_V1
|
||||
}
|
||||
|
||||
/// Borrow the nested [`Function`] when this is [`JobResult::Function`].
|
||||
pub fn function(&self) -> Option<&Function> {
|
||||
match self {
|
||||
Self::None => None,
|
||||
Self::Function(function) => Some(function),
|
||||
}
|
||||
}
|
||||
|
||||
/// Consume this value and return the nested [`Function`] when present.
|
||||
pub fn into_function(self) -> Option<Function> {
|
||||
match self {
|
||||
Self::None => None,
|
||||
Self::Function(function) => Some(function),
|
||||
}
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> JobResultWire {
|
||||
match self {
|
||||
Self::None => JobResultWire::None {
|
||||
format_version: JOB_RESULT_FORMAT_VERSION_V1,
|
||||
},
|
||||
Self::Function(function) => JobResultWire::Function {
|
||||
format_version: JOB_RESULT_FORMAT_VERSION_V1,
|
||||
function: function.clone(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
fn from_wire(wire: JobResultWire) -> Result<Self> {
|
||||
match wire {
|
||||
JobResultWire::None { format_version } => {
|
||||
if format_version != JOB_RESULT_FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported JobResult format_version {format_version}"
|
||||
)));
|
||||
}
|
||||
Ok(Self::None)
|
||||
}
|
||||
JobResultWire::Function {
|
||||
format_version,
|
||||
function,
|
||||
} => {
|
||||
if format_version != JOB_RESULT_FORMAT_VERSION_V1 {
|
||||
return Err(invalid_input(format!(
|
||||
"unsupported JobResult format_version {format_version}"
|
||||
)));
|
||||
}
|
||||
Ok(Self::Function(function))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize)]
|
||||
#[serde(tag = "kind", deny_unknown_fields)]
|
||||
enum JobResultWire {
|
||||
#[serde(rename = "none")]
|
||||
None { format_version: u32 },
|
||||
#[serde(rename = "function")]
|
||||
Function {
|
||||
format_version: u32,
|
||||
function: Function,
|
||||
},
|
||||
}
|
||||
|
||||
impl Serialize for JobResult {
|
||||
fn serialize<S>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
self.to_wire().serialize(serializer)
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for JobResult {
|
||||
fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let wire = JobResultWire::deserialize(deserializer)?;
|
||||
Self::from_wire(wire).map_err(D::Error::custom)
|
||||
}
|
||||
}
|
||||
|
||||
/// Backend-specific tracking for an asynchronous operation.
|
||||
#[async_trait]
|
||||
@@ -19,7 +133,7 @@ pub(crate) trait JobHandle: Send + Sync {
|
||||
None
|
||||
}
|
||||
async fn status(&self) -> Result<String>;
|
||||
async fn wait(&self) -> Result<()>;
|
||||
async fn wait(&self) -> Result<JobResult>;
|
||||
async fn cancel(&self) -> Result<()>;
|
||||
}
|
||||
|
||||
@@ -52,7 +166,7 @@ impl Job {
|
||||
}
|
||||
|
||||
/// A job running as a task in this process.
|
||||
pub(crate) fn spawned(task: JoinHandle<Result<()>>) -> Self {
|
||||
pub(crate) fn spawned(task: JoinHandle<Result<JobResult>>) -> Self {
|
||||
Self::new(Box::new(SpawnedJob::new(task)))
|
||||
}
|
||||
|
||||
@@ -81,11 +195,14 @@ impl Job {
|
||||
|
||||
/// Waits until the operation reaches a terminal state.
|
||||
///
|
||||
/// On success, returns the job's [`JobResult`]. Operations that produce no
|
||||
/// resource result yield [`JobResult::None`].
|
||||
///
|
||||
/// Returns [`crate::Error::JobFailed`] if the operation failed and
|
||||
/// [`crate::Error::JobCancelled`] if it was cancelled.
|
||||
pub async fn wait(&self) -> Result<()> {
|
||||
pub async fn wait(&self) -> Result<JobResult> {
|
||||
match &self.handle {
|
||||
None => Ok(()),
|
||||
None => Ok(JobResult::None),
|
||||
Some(handle) => handle.wait().await,
|
||||
}
|
||||
}
|
||||
@@ -105,15 +222,15 @@ impl Job {
|
||||
/// the outcome; [`Error`] is not, so failures share one behind an [`Arc`].
|
||||
#[derive(Clone)]
|
||||
enum Outcome {
|
||||
Succeeded,
|
||||
Succeeded(JobResult),
|
||||
Failed(Arc<Error>),
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
impl Outcome {
|
||||
fn into_result(self) -> Result<()> {
|
||||
fn into_result(self) -> Result<JobResult> {
|
||||
match self {
|
||||
Self::Succeeded => Ok(()),
|
||||
Self::Succeeded(result) => Ok(result),
|
||||
Self::Failed(source) => Err(Error::JobFailed {
|
||||
job_id: None,
|
||||
failure: JobFailure::from_source(source),
|
||||
@@ -132,16 +249,16 @@ struct SpawnedJob {
|
||||
}
|
||||
|
||||
impl SpawnedJob {
|
||||
fn new(task: JoinHandle<Result<()>>) -> Self {
|
||||
fn new(task: JoinHandle<Result<JobResult>>) -> Self {
|
||||
let abort = task.abort_handle();
|
||||
let (tx, outcome) = watch::channel(None);
|
||||
tokio::spawn(async move {
|
||||
let outcome = match task.await {
|
||||
Ok(Ok(())) => Outcome::Succeeded,
|
||||
Ok(Ok(result)) => Outcome::Succeeded(result),
|
||||
Ok(Err(err)) => Outcome::Failed(Arc::new(err)),
|
||||
Err(err) if err.is_cancelled() => Outcome::Cancelled,
|
||||
Err(err) => Outcome::Failed(Arc::new(Error::Runtime {
|
||||
message: format!("index job task failed: {err}"),
|
||||
message: format!("job task failed: {err}"),
|
||||
})),
|
||||
};
|
||||
let _ = tx.send(Some(outcome));
|
||||
@@ -155,20 +272,20 @@ impl JobHandle for SpawnedJob {
|
||||
async fn status(&self) -> Result<String> {
|
||||
let label = match &*self.outcome.borrow() {
|
||||
None => "running",
|
||||
Some(Outcome::Succeeded) => "finished",
|
||||
Some(Outcome::Succeeded(_)) => "finished",
|
||||
Some(Outcome::Failed(_)) => "failed",
|
||||
Some(Outcome::Cancelled) => "cancelled",
|
||||
};
|
||||
Ok(label.to_string())
|
||||
}
|
||||
|
||||
async fn wait(&self) -> Result<()> {
|
||||
async fn wait(&self) -> Result<JobResult> {
|
||||
let mut outcome = self.outcome.clone();
|
||||
let settled = outcome
|
||||
.wait_for(|outcome| outcome.is_some())
|
||||
.await
|
||||
.map_err(|_| Error::Runtime {
|
||||
message: "index job outcome was dropped before it completed".to_string(),
|
||||
message: "job outcome was dropped before it completed".to_string(),
|
||||
})?
|
||||
.clone()
|
||||
.expect("wait_for returns once an outcome is set");
|
||||
@@ -180,3 +297,166 @@ impl JobHandle for SpawnedJob {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::future::Future;
|
||||
use std::pin::pin;
|
||||
use std::task::{Context, Poll, Waker};
|
||||
|
||||
use arrow_schema::DataType;
|
||||
use tokio::sync::oneshot;
|
||||
|
||||
use super::*;
|
||||
use crate::error::FunctionErrorCode;
|
||||
use crate::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
|
||||
fn sample_success_function() -> Function {
|
||||
let id = FunctionId::try_new("fn.exact.local-job-result").expect("valid FunctionId");
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("x", DataType::Int32)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.expect("valid FunctionSignature");
|
||||
Function::new(id, signature)
|
||||
}
|
||||
|
||||
fn assert_exact_function(actual: &Function, expected: &Function) {
|
||||
assert_eq!(actual.id(), expected.id());
|
||||
assert_eq!(actual.signature(), expected.signature());
|
||||
}
|
||||
|
||||
/// A completed-before-handle local job projects success as None.
|
||||
#[tokio::test]
|
||||
async fn local_job_result_new_done_wait_returns_none() {
|
||||
let job = Job::new_done();
|
||||
let result = job.wait().await.expect("new_done must succeed");
|
||||
assert_eq!(result, JobResult::None);
|
||||
}
|
||||
|
||||
/// A local spawned unit / no-resource success projects as None.
|
||||
#[tokio::test]
|
||||
async fn local_job_result_spawned_unit_success_projects_none() {
|
||||
let job = Job::spawned(tokio::spawn(async { Ok(JobResult::None) }));
|
||||
let result = job
|
||||
.wait()
|
||||
.await
|
||||
.expect("unit success must finish without error");
|
||||
assert_eq!(result, JobResult::None);
|
||||
}
|
||||
|
||||
/// Function success is cloneable and shared by concurrent + late waiters.
|
||||
///
|
||||
/// Wait futures are pinned and polled once to Pending while success is still
|
||||
/// gated, proving they observed the running state before publication.
|
||||
#[tokio::test]
|
||||
async fn local_job_result_spawned_function_shared_by_waiters() {
|
||||
let expected = sample_success_function();
|
||||
let (release_tx, release_rx) = oneshot::channel();
|
||||
|
||||
let job = Job::spawned(tokio::spawn({
|
||||
let function = expected.clone();
|
||||
async move {
|
||||
release_rx
|
||||
.await
|
||||
.expect("success task must be released by the test");
|
||||
Ok(JobResult::Function(function))
|
||||
}
|
||||
}));
|
||||
|
||||
let mut wait_a = pin!(job.wait());
|
||||
let mut wait_b = pin!(job.wait());
|
||||
let waker = Waker::noop();
|
||||
let mut cx = Context::from_waker(waker);
|
||||
|
||||
assert!(
|
||||
matches!(wait_a.as_mut().poll(&mut cx), Poll::Pending),
|
||||
"waiter A must poll Pending before success publication"
|
||||
);
|
||||
assert!(
|
||||
matches!(wait_b.as_mut().poll(&mut cx), Poll::Pending),
|
||||
"waiter B must poll Pending before success publication"
|
||||
);
|
||||
|
||||
release_tx
|
||||
.send(())
|
||||
.expect("success task must still be waiting on the gate");
|
||||
|
||||
let result_a = wait_a
|
||||
.await
|
||||
.expect("concurrent waiter A must observe success");
|
||||
let result_b = wait_b
|
||||
.await
|
||||
.expect("concurrent waiter B must observe success");
|
||||
let result_late = job
|
||||
.wait()
|
||||
.await
|
||||
.expect("late waiter must observe the same success");
|
||||
|
||||
for result in [&result_a, &result_b, &result_late] {
|
||||
match result {
|
||||
JobResult::Function(function) => assert_exact_function(function, &expected),
|
||||
JobResult::None => panic!("Function success must not project as JobResult::None"),
|
||||
}
|
||||
}
|
||||
assert_eq!(result_a, result_b);
|
||||
assert_eq!(result_a, result_late);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawned_job_function_failure_returns_job_failed_with_same_code() {
|
||||
let job = Job::spawned(tokio::spawn(async {
|
||||
Err(Error::Function {
|
||||
code: FunctionErrorCode::UdfExecutionFailure,
|
||||
// Message names a different category on purpose; code is structural.
|
||||
message: "looks like name_conflict to a string parser".to_string(),
|
||||
})
|
||||
}));
|
||||
|
||||
let err = job
|
||||
.wait()
|
||||
.await
|
||||
.expect_err("Function failure must fail the job");
|
||||
match err {
|
||||
Error::JobFailed { failure, .. } => match &failure.error_code {
|
||||
Some(code) => {
|
||||
assert_eq!(code, &FunctionErrorCode::UdfExecutionFailure);
|
||||
assert_ne!(code, &FunctionErrorCode::NameConflict);
|
||||
}
|
||||
None => panic!("local Function failure must project error_code onto JobFailure"),
|
||||
},
|
||||
other => panic!("expected Error::JobFailed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn spawned_job_preserves_unrecognized_function_error_code() {
|
||||
let raw = "enterprise_future_category_xyz";
|
||||
let job = Job::spawned(tokio::spawn({
|
||||
let raw = raw.to_string();
|
||||
async move {
|
||||
Err(Error::Function {
|
||||
code: FunctionErrorCode::Unrecognized(raw),
|
||||
message: "future server category".to_string(),
|
||||
})
|
||||
}
|
||||
}));
|
||||
|
||||
let err = job
|
||||
.wait()
|
||||
.await
|
||||
.expect_err("Function failure must fail the job");
|
||||
match err {
|
||||
Error::JobFailed { failure, .. } => match &failure.error_code {
|
||||
Some(FunctionErrorCode::Unrecognized(preserved)) => {
|
||||
assert_eq!(preserved, raw);
|
||||
}
|
||||
Some(other) => panic!("unrecognized code must not become known: {other:?}"),
|
||||
None => panic!("unrecognized Function code must be preserved on JobFailure"),
|
||||
},
|
||||
other => panic!("expected Error::JobFailed, got {other:?}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -181,6 +181,7 @@ pub mod dataloader;
|
||||
pub mod embeddings;
|
||||
pub mod error;
|
||||
pub mod expr;
|
||||
pub mod function;
|
||||
pub mod index;
|
||||
pub mod io;
|
||||
pub mod ipc;
|
||||
@@ -205,7 +206,7 @@ use serde::{Deserialize, Serialize};
|
||||
pub use blob::{BlobRangeRequest, blob, is_blob};
|
||||
pub use connection::{ConnectNamespaceBuilder, Connection};
|
||||
pub use error::{Error, JobFailure, Result};
|
||||
pub use job::Job;
|
||||
pub use job::{Job, JobResult};
|
||||
use lance_index::vector::ApproxMode as LanceApproxMode;
|
||||
use lance_linalg::distance::DistanceType as LanceDistanceType;
|
||||
/// Re-export of the [`metrics`](https://docs.rs/metrics) crate facade. Enable
|
||||
|
||||
+1345
-3
File diff suppressed because it is too large
Load Diff
@@ -8,10 +8,12 @@
|
||||
|
||||
pub(crate) mod client;
|
||||
pub(crate) mod db;
|
||||
pub(crate) mod function;
|
||||
pub(crate) mod job;
|
||||
pub mod oauth;
|
||||
mod retry;
|
||||
pub(crate) mod table;
|
||||
mod transport;
|
||||
pub(crate) mod util;
|
||||
|
||||
const ARROW_STREAM_CONTENT_TYPE: &str = "application/vnd.apache.arrow.stream";
|
||||
|
||||
@@ -15,6 +15,51 @@ use crate::remote::retry::{ResolvedRetryConfig, RetryCounter};
|
||||
|
||||
const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
|
||||
|
||||
/// Privacy mode for request logging and non-success response handling.
|
||||
///
|
||||
/// [`RequestPrivacy::Standard`] preserves the existing harmless JSON body
|
||||
/// visibility. [`RequestPrivacy::Sensitive`] never includes request bodies or
|
||||
/// headers in logs, and never folds response bodies into error chains.
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
enum RequestPrivacy {
|
||||
Standard,
|
||||
Sensitive,
|
||||
}
|
||||
|
||||
/// Format a request for debug logging according to [`RequestPrivacy`].
|
||||
fn format_request_log(request: &Request, request_id: &str, privacy: RequestPrivacy) -> String {
|
||||
match privacy {
|
||||
RequestPrivacy::Standard => {
|
||||
let content_type = request
|
||||
.headers()
|
||||
.get("content-type")
|
||||
.and_then(|v| v.to_str().ok());
|
||||
if content_type == Some("application/json") {
|
||||
let body = request
|
||||
.body()
|
||||
.and_then(|b| b.as_bytes())
|
||||
.map(|b| String::from_utf8_lossy(b).into_owned())
|
||||
.unwrap_or_default();
|
||||
format!(
|
||||
"Sending request_id={}: {:?} with body {}",
|
||||
request_id, request, body
|
||||
)
|
||||
} else {
|
||||
format!("Sending request_id={}: {:?}", request_id, request)
|
||||
}
|
||||
}
|
||||
RequestPrivacy::Sensitive => {
|
||||
// Safe context only: request id, method, and URL. Never body or headers.
|
||||
format!(
|
||||
"Sending request_id={}: {} {}",
|
||||
request_id,
|
||||
request.method(),
|
||||
request.url()
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Configuration for TLS/mTLS settings.
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct TlsConfig {
|
||||
@@ -746,6 +791,41 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
Ok((request_id, response))
|
||||
}
|
||||
|
||||
/// Send one attempt with a caller-owned request id.
|
||||
///
|
||||
/// Shared by explicit-`error_code` classifiers for Function catalog and
|
||||
/// Remote table query routes. Keeps the caller-owned request ID, uses
|
||||
/// sensitive logging, applies dynamic headers, sends one uninterpreted
|
||||
/// attempt, and leaves status/body classification to the caller.
|
||||
pub(crate) async fn send_attempt_with_request_id(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
request_id: &str,
|
||||
) -> Result<Response> {
|
||||
let (client, request) = req_builder.build_split();
|
||||
let mut request = request.map_err(|e| Error::Runtime {
|
||||
message: format!("Failed to build request: {}", e),
|
||||
})?;
|
||||
self.set_request_id(&mut request, request_id);
|
||||
request = self.apply_dynamic_headers(request).await?;
|
||||
if log::log_enabled!(log::Level::Debug) {
|
||||
debug!(
|
||||
"{}",
|
||||
format_request_log(&request, request_id, RequestPrivacy::Sensitive)
|
||||
);
|
||||
}
|
||||
let response = self
|
||||
.sender
|
||||
.send(&client, request)
|
||||
.await
|
||||
.err_to_http(request_id.to_string())?;
|
||||
debug!(
|
||||
"Received response for request_id={}: {:?}",
|
||||
request_id, response
|
||||
);
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
/// Send the request using retries configured in the RetryConfig.
|
||||
/// If retry_5xx is false, 5xx requests will not be retried regardless of the statuses configured
|
||||
/// in the RetryConfig.
|
||||
@@ -753,9 +833,37 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
pub async fn send_with_retry(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
mut make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
retry_5xx: bool,
|
||||
) -> Result<(String, Response)> {
|
||||
self.send_with_retry_inner(req_builder, make_body, retry_5xx, RequestPrivacy::Standard)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Like [`Self::send_with_retry`], but never logs request bodies/headers and
|
||||
/// never folds non-success response bodies into retry or HTTP error chains.
|
||||
///
|
||||
/// Privacy affects only logging and error-body exposure; retry budgets are
|
||||
/// identical to [`Self::send_with_retry`] for the same [`RetryConfig`].
|
||||
pub(crate) async fn send_sensitive_with_retry(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
retry_5xx: bool,
|
||||
) -> Result<(String, Response)> {
|
||||
self.send_with_retry_inner(req_builder, make_body, retry_5xx, RequestPrivacy::Sensitive)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn send_with_retry_inner(
|
||||
&self,
|
||||
req_builder: RequestBuilder,
|
||||
mut make_body: Option<Box<dyn FnMut() -> Result<Body> + Send + 'static>>,
|
||||
retry_5xx: bool,
|
||||
privacy: RequestPrivacy,
|
||||
) -> Result<(String, Response)> {
|
||||
// Privacy must never alter retry budgets: both Standard and Sensitive
|
||||
// share the same ResolvedRetryConfig / RetryCounter semantics.
|
||||
let retry_config = &self.retry_config;
|
||||
let non_5xx_statuses = retry_config
|
||||
.statuses
|
||||
@@ -772,6 +880,7 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
let mut r = r.map_err(|e| Error::Runtime {
|
||||
message: format!("Failed to build request: {}", e),
|
||||
})?;
|
||||
// One SDK-generated request id is reused across every retry attempt.
|
||||
let request_id = self.extract_request_id(&mut r);
|
||||
let mut retry_counter = RetryCounter::new(retry_config, request_id.clone());
|
||||
|
||||
@@ -790,12 +899,14 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
let mut request = request.map_err(|e| Error::Runtime {
|
||||
message: format!("Failed to build request: {}", e),
|
||||
})?;
|
||||
self.set_request_id(&mut request, &request_id.clone());
|
||||
self.set_request_id(&mut request, &request_id);
|
||||
|
||||
// Apply dynamic headers before each retry attempt
|
||||
request = self.apply_dynamic_headers(request).await?;
|
||||
|
||||
self.log_request(&request, &request_id);
|
||||
if log::log_enabled!(log::Level::Debug) {
|
||||
debug!("{}", format_request_log(&request, &request_id, privacy));
|
||||
}
|
||||
|
||||
let response = self.sender.send(&c, request).await.map(|r| (r.status(), r));
|
||||
|
||||
@@ -811,10 +922,16 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
if (retry_5xx && retry_config.statuses.contains(&status))
|
||||
|| non_5xx_statuses.contains(&status) =>
|
||||
{
|
||||
let source = self
|
||||
.check_response(&retry_counter.request_id, response)
|
||||
.await
|
||||
.unwrap_err();
|
||||
let source = match privacy {
|
||||
RequestPrivacy::Standard => self
|
||||
.check_response(&retry_counter.request_id, response)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
RequestPrivacy::Sensitive => self
|
||||
.check_sensitive_response(&retry_counter.request_id, response)
|
||||
.await
|
||||
.unwrap_err(),
|
||||
};
|
||||
retry_counter.increment_request_failures(source)?;
|
||||
}
|
||||
Err(err) if err.is_connect() => {
|
||||
@@ -839,22 +956,12 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn log_request(&self, request: &Request, request_id: &String) {
|
||||
pub(crate) fn log_request(&self, request: &Request, request_id: &str) {
|
||||
if log::log_enabled!(log::Level::Debug) {
|
||||
let content_type = request
|
||||
.headers()
|
||||
.get("content-type")
|
||||
.map(|v| v.to_str().unwrap());
|
||||
if content_type == Some("application/json") {
|
||||
let body = request.body().as_ref().unwrap().as_bytes().unwrap();
|
||||
let body = String::from_utf8_lossy(body);
|
||||
debug!(
|
||||
"Sending request_id={}: {:?} with body {}",
|
||||
request_id, request, body
|
||||
);
|
||||
} else {
|
||||
debug!("Sending request_id={}: {:?}", request_id, request);
|
||||
}
|
||||
debug!(
|
||||
"{}",
|
||||
format_request_log(request, request_id, RequestPrivacy::Standard)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -898,6 +1005,27 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Like [`Self::check_response`], but discards the response body on failure
|
||||
/// so marker-bearing payloads never enter [`Error::Http`] chains.
|
||||
pub(crate) async fn check_sensitive_response(
|
||||
&self,
|
||||
request_id: &str,
|
||||
response: Response,
|
||||
) -> Result<Response> {
|
||||
let status = response.status();
|
||||
if status.is_success() {
|
||||
Ok(response)
|
||||
} else {
|
||||
// Discard the body entirely; never fold it into Error::Http.
|
||||
let _ = response.bytes().await;
|
||||
Err(Error::Http {
|
||||
source: status.to_string().into(),
|
||||
request_id: request_id.into(),
|
||||
status_code: Some(status),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub trait RequestResultExt {
|
||||
@@ -1066,6 +1194,7 @@ pub mod test_utils {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serial_test::serial;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::time::Duration;
|
||||
|
||||
// Serializes the env-var-mutating tests below: cargo test runs tests in
|
||||
@@ -1664,4 +1793,253 @@ mod tests {
|
||||
}
|
||||
assert!(matches!(err, Error::InvalidInput { .. }), "got: {err:?}");
|
||||
}
|
||||
|
||||
// -------------------------------------------------------------------------
|
||||
// Sensitive-request privacy mode (generic transport; RED until helpers exist)
|
||||
// -------------------------------------------------------------------------
|
||||
|
||||
const PRIVACY_SOURCE_MARKER: &str = "SENSITIVE_PRIVACY_SOURCE_BODY_MARKER_client";
|
||||
const PRIVACY_SECRET_MARKER: &str = "secret://team/client-privacy-token";
|
||||
|
||||
fn privacy_json_request(url: &str, body: &str, request_id: &str) -> Request {
|
||||
reqwest::Client::new()
|
||||
.post(url)
|
||||
.header("content-type", "application/json")
|
||||
.header("x-request-id", request_id)
|
||||
.body(body.to_string())
|
||||
.build()
|
||||
.expect("build privacy fixture request")
|
||||
}
|
||||
|
||||
fn assert_markers_absent(text: &str) {
|
||||
assert!(
|
||||
!text.contains(PRIVACY_SOURCE_MARKER),
|
||||
"source marker must be absent: {text}"
|
||||
);
|
||||
assert!(
|
||||
!text.contains(PRIVACY_SECRET_MARKER),
|
||||
"secret marker must be absent: {text}"
|
||||
);
|
||||
}
|
||||
|
||||
fn error_chain_text(err: &Error) -> String {
|
||||
let mut text = format!("{err}\n{err:?}");
|
||||
let mut current: Option<&(dyn std::error::Error + 'static)> = Some(err);
|
||||
while let Some(e) = current {
|
||||
text.push('\n');
|
||||
text.push_str(&e.to_string());
|
||||
text.push('\n');
|
||||
text.push_str(&format!("{e:?}"));
|
||||
current = e.source();
|
||||
}
|
||||
text
|
||||
}
|
||||
|
||||
/// Standard JSON request logging keeps the current harmless body visibility.
|
||||
#[test]
|
||||
fn format_request_log_standard_retains_harmless_json_body() {
|
||||
let request_id = "req-privacy-standard";
|
||||
let body = r#"{"ok":true,"note":"harmless-visible-body"}"#;
|
||||
let request = privacy_json_request("http://localhost/v1/table/", body, request_id);
|
||||
|
||||
let log = format_request_log(&request, request_id, RequestPrivacy::Standard);
|
||||
|
||||
assert!(
|
||||
log.contains(request_id),
|
||||
"standard log must retain request id: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("POST"),
|
||||
"standard log must retain method: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("/v1/table/"),
|
||||
"standard log must retain URL path: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("harmless-visible-body"),
|
||||
"standard JSON logging must retain body visibility: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains(body) || log.contains(r#""note":"harmless-visible-body""#),
|
||||
"standard JSON logging must include the harmless JSON body: {log}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Sensitive JSON formatting redacts the entire body and keeps only safe context.
|
||||
#[test]
|
||||
fn format_request_log_sensitive_redacts_json_body_keeps_safe_context() {
|
||||
let request_id = "req-privacy-sensitive";
|
||||
let body =
|
||||
format!(r#"{{"source":"{PRIVACY_SOURCE_MARKER}","secret":"{PRIVACY_SECRET_MARKER}"}}"#);
|
||||
let request =
|
||||
privacy_json_request("http://localhost/v1/functions/register", &body, request_id);
|
||||
|
||||
let log = format_request_log(&request, request_id, RequestPrivacy::Sensitive);
|
||||
|
||||
assert!(
|
||||
log.contains(request_id),
|
||||
"sensitive log must retain request id: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("POST"),
|
||||
"sensitive log must retain method: {log}"
|
||||
);
|
||||
assert!(
|
||||
log.contains("/v1/functions/register"),
|
||||
"sensitive log must retain URL path: {log}"
|
||||
);
|
||||
assert_markers_absent(&log);
|
||||
assert!(
|
||||
!log.contains(&body),
|
||||
"sensitive JSON formatting must redact the entire body: {log}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Sensitive non-success responses omit the response body from Error::Http text.
|
||||
#[tokio::test]
|
||||
async fn check_sensitive_response_omits_non_success_response_body() {
|
||||
let client = test_utils::client_with_handler(|_| {
|
||||
http::Response::builder().status(200).body("").unwrap()
|
||||
});
|
||||
let response: Response = http::Response::builder()
|
||||
.status(400)
|
||||
.body(format!(
|
||||
"client error echoed {PRIVACY_SOURCE_MARKER} and {PRIVACY_SECRET_MARKER}"
|
||||
))
|
||||
.unwrap()
|
||||
.into();
|
||||
|
||||
let err = client
|
||||
.check_sensitive_response("req-privacy-check", response)
|
||||
.await
|
||||
.expect_err("non-success sensitive response must fail closed");
|
||||
|
||||
assert!(
|
||||
matches!(err, Error::Http { .. }),
|
||||
"expected Error::Http, got {err:?}"
|
||||
);
|
||||
assert_markers_absent(&error_chain_text(&err));
|
||||
}
|
||||
|
||||
/// Sensitive send+retry must not leak request/response markers into retry errors.
|
||||
#[tokio::test]
|
||||
async fn send_sensitive_with_retry_omits_markers_from_exhausted_retry_errors() {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let counted = call_count.clone();
|
||||
let client = test_utils::client_with_handler_and_config(
|
||||
move |request| {
|
||||
counted.fetch_add(1, Ordering::SeqCst);
|
||||
let body = request.body().and_then(|b| b.as_bytes()).unwrap_or(b"");
|
||||
let body = std::str::from_utf8(body).unwrap_or("");
|
||||
assert!(
|
||||
body.contains(PRIVACY_SOURCE_MARKER) && body.contains(PRIVACY_SECRET_MARKER),
|
||||
"trusted wire body must still carry sensitive fields"
|
||||
);
|
||||
http::Response::builder()
|
||||
.status(500)
|
||||
.body(format!(
|
||||
"server echoed {PRIVACY_SOURCE_MARKER} and {PRIVACY_SECRET_MARKER}"
|
||||
))
|
||||
.unwrap()
|
||||
},
|
||||
ClientConfig {
|
||||
retry_config: RetryConfig {
|
||||
// RetryCounter treats `retries` as max request failures, so
|
||||
// retries=2 yields exactly two transport attempts before Error::Retry.
|
||||
retries: Some(2),
|
||||
backoff_factor: Some(0.0),
|
||||
backoff_jitter: Some(0.0),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
|
||||
let payload = serde_json::json!({
|
||||
"source": PRIVACY_SOURCE_MARKER,
|
||||
"secret": PRIVACY_SECRET_MARKER,
|
||||
});
|
||||
let req = client.post("/v1/functions/register").json(&payload);
|
||||
let err = client
|
||||
.send_sensitive_with_retry(req, None, true)
|
||||
.await
|
||||
.expect_err("exhausted sensitive 5xx retries must fail");
|
||||
|
||||
assert!(
|
||||
matches!(err, Error::Retry { .. }),
|
||||
"expected Error::Retry, got {err:?}"
|
||||
);
|
||||
assert_markers_absent(&error_chain_text(&err));
|
||||
assert_eq!(
|
||||
call_count.load(Ordering::SeqCst),
|
||||
2,
|
||||
"RetryCounter max request failures=2 must make exactly two transport attempts"
|
||||
);
|
||||
}
|
||||
|
||||
/// Standard and Sensitive share the same RetryCounter attempt budget.
|
||||
#[tokio::test]
|
||||
async fn send_with_retry_standard_and_sensitive_share_attempt_budget() {
|
||||
async fn exhausted_attempts(sensitive: bool) -> usize {
|
||||
let call_count = Arc::new(AtomicUsize::new(0));
|
||||
let counted = call_count.clone();
|
||||
let client = test_utils::client_with_handler_and_config(
|
||||
move |_| {
|
||||
counted.fetch_add(1, Ordering::SeqCst);
|
||||
http::Response::builder()
|
||||
.status(500)
|
||||
.body(format!(
|
||||
"server echoed {PRIVACY_SOURCE_MARKER} and {PRIVACY_SECRET_MARKER}"
|
||||
))
|
||||
.unwrap()
|
||||
},
|
||||
ClientConfig {
|
||||
retry_config: RetryConfig {
|
||||
retries: Some(2),
|
||||
backoff_factor: Some(0.0),
|
||||
backoff_jitter: Some(0.0),
|
||||
..Default::default()
|
||||
},
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
|
||||
let payload = serde_json::json!({
|
||||
"source": PRIVACY_SOURCE_MARKER,
|
||||
"secret": PRIVACY_SECRET_MARKER,
|
||||
});
|
||||
let req = client.post("/v1/functions/register").json(&payload);
|
||||
let err = if sensitive {
|
||||
client
|
||||
.send_sensitive_with_retry(req, None, true)
|
||||
.await
|
||||
.expect_err("exhausted sensitive 5xx retries must fail")
|
||||
} else {
|
||||
client
|
||||
.send_with_retry(req, None, true)
|
||||
.await
|
||||
.expect_err("exhausted standard 5xx retries must fail")
|
||||
};
|
||||
assert!(
|
||||
matches!(err, Error::Retry { .. }),
|
||||
"expected Error::Retry, got {err:?}"
|
||||
);
|
||||
if sensitive {
|
||||
assert_markers_absent(&error_chain_text(&err));
|
||||
}
|
||||
call_count.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
let standard_attempts = exhausted_attempts(false).await;
|
||||
let sensitive_attempts = exhausted_attempts(true).await;
|
||||
assert_eq!(
|
||||
standard_attempts, sensitive_attempts,
|
||||
"Standard and Sensitive must share the same attempt budget for identical RetryConfig"
|
||||
);
|
||||
assert_eq!(
|
||||
standard_attempts, 2,
|
||||
"RetryCounter max request failures=2 must make exactly two transport attempts"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+3266
-53
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,270 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Remote first-class Function catalog wire helpers.
|
||||
//!
|
||||
//! POST `/v1/functions/lookup` resolves a database-scoped name or exact
|
||||
//! [`FunctionId`] to an immutable [`Function`] value. Name is lookup
|
||||
//! indirection only and never becomes part of the returned handle.
|
||||
//!
|
||||
//! POST `/v1/functions/remove` performs a direct synchronous catalog CAS that
|
||||
//! unbinds a name when the caller's observed [`Function`] id still matches.
|
||||
//! This is not a Job, not physical Function deletion, and not revocation.
|
||||
//!
|
||||
//! POST `/v1/functions/revoke` performs a direct synchronous administrator
|
||||
//! catalog set-bit for an exact [`Function`] id. This is not a Job, not name
|
||||
//! removal, not physical deletion, and not Function mutation.
|
||||
|
||||
use reqwest::{RequestBuilder, StatusCode};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::error::{Error, Result};
|
||||
use crate::function::{Function, FunctionId};
|
||||
|
||||
use super::client::{HttpSend, RestfulLanceDbClient};
|
||||
use super::transport::{
|
||||
BeforeBody, BodyAction, explicit_error_code, post_with_body_classification,
|
||||
};
|
||||
|
||||
const LOOKUP_PATH: &str = "/v1/functions/lookup";
|
||||
const REMOVE_PATH: &str = "/v1/functions/remove";
|
||||
const REVOKE_PATH: &str = "/v1/functions/revoke";
|
||||
|
||||
/// Fixed client diagnostic for [`Error::Function`]. Never carry server text,
|
||||
/// selector values, or response payload bytes.
|
||||
const LOOKUP_FUNCTION_ERROR_MESSAGE: &str = "function lookup failed";
|
||||
|
||||
/// Fixed client diagnostic for protocol / HTTP failures. Never include response
|
||||
/// payload bytes or selector values.
|
||||
const LOOKUP_HTTP_ERROR_MESSAGE: &str = "function lookup request failed";
|
||||
|
||||
/// Fixed client diagnostic for malformed success payloads.
|
||||
const LOOKUP_INVALID_SUCCESS_MESSAGE: &str = "function lookup response missing or invalid function";
|
||||
|
||||
/// Fixed client diagnostic for remove [`Error::Function`]. Never carry server
|
||||
/// text, catalog name, Function id, or response payload bytes.
|
||||
const REMOVE_FUNCTION_ERROR_MESSAGE: &str = "function name removal failed";
|
||||
|
||||
/// Fixed client diagnostic for remove protocol / HTTP failures.
|
||||
const REMOVE_HTTP_ERROR_MESSAGE: &str = "function name removal request failed";
|
||||
|
||||
/// Fixed client diagnostic for revoke [`Error::Function`]. Never carry server
|
||||
/// text, Function id, or response payload bytes.
|
||||
const REVOKE_FUNCTION_ERROR_MESSAGE: &str = "function revocation failed";
|
||||
|
||||
/// Fixed client diagnostic for revoke protocol / HTTP failures.
|
||||
const REVOKE_HTTP_ERROR_MESSAGE: &str = "function revocation request failed";
|
||||
|
||||
/// One exact lookup selector. Exactly one variant is serialized on the wire.
|
||||
pub enum FunctionLookupSelector {
|
||||
Name(String),
|
||||
FunctionId(String),
|
||||
}
|
||||
|
||||
impl FunctionLookupSelector {
|
||||
pub fn by_name(name: impl Into<String>) -> Result<Self> {
|
||||
let name = name.into();
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function lookup name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
Ok(Self::Name(name))
|
||||
}
|
||||
|
||||
pub fn by_function_id(function_id: &FunctionId) -> Self {
|
||||
Self::FunctionId(function_id.as_str().to_string())
|
||||
}
|
||||
|
||||
fn to_wire(&self) -> Value {
|
||||
match self {
|
||||
Self::Name(name) => serde_json::json!({ "name": name }),
|
||||
Self::FunctionId(function_id) => {
|
||||
serde_json::json!({ "function_id": function_id })
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct LookupSuccessResponse {
|
||||
function: Function,
|
||||
}
|
||||
|
||||
/// Resolve a Function via POST `/v1/functions/lookup`.
|
||||
///
|
||||
/// Transport classification matches [`RestfulLanceDbClient::send_with_retry`]:
|
||||
/// connect → connect_retries; timeout/body/decode (including response-byte
|
||||
/// reads) → read_retries; configured retryable statuses without an explicit
|
||||
/// `error_code` → request retries; all other transport/client errors return
|
||||
/// immediately. An explicit nonempty `error_code` is terminal and wins over
|
||||
/// HTTP status. Request/response payload bytes never enter error chains.
|
||||
pub async fn lookup_function<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
selector: FunctionLookupSelector,
|
||||
) -> Result<Function> {
|
||||
let req_builder = client.post(LOOKUP_PATH).json(&selector.to_wire());
|
||||
post_with_body_classification(
|
||||
client,
|
||||
req_builder,
|
||||
LOOKUP_HTTP_ERROR_MESSAGE,
|
||||
|_status, _request_id| BeforeBody::ReadBody,
|
||||
|status, bytes, request_id| {
|
||||
if status.is_success() {
|
||||
return BodyAction::Done(decode_lookup_success(bytes, request_id));
|
||||
}
|
||||
if let Some(code) = explicit_error_code(bytes) {
|
||||
return BodyAction::Done(Err(Error::Function {
|
||||
code,
|
||||
message: LOOKUP_FUNCTION_ERROR_MESSAGE.to_string(),
|
||||
}));
|
||||
}
|
||||
if client.retry_config.statuses.contains(&status) {
|
||||
return BodyAction::RetryRequest;
|
||||
}
|
||||
BodyAction::Done(Err(Error::Http {
|
||||
source: LOOKUP_HTTP_ERROR_MESSAGE.into(),
|
||||
request_id,
|
||||
status_code: Some(status),
|
||||
}))
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Conditionally remove a database-scoped Function name via POST
|
||||
/// `/v1/functions/remove`.
|
||||
///
|
||||
/// Direct synchronous catalog CAS: the wire body is exactly
|
||||
/// `{"name","expected_current_function_id"}` using only `current.id`. Only
|
||||
/// HTTP 204 means the CAS completed; other 2xx are payload-free protocol
|
||||
/// [`Error::Http`]. Empty names are [`Error::InvalidInput`] before transport.
|
||||
///
|
||||
/// Retry budgets match lookup: stable internal request id and exact cloned
|
||||
/// body across attempts; response-byte failures consume read budget; configured
|
||||
/// retryable status without explicit `error_code` consumes request budget;
|
||||
/// header/client errors are immediate. Sophon deduplicates the internal request
|
||||
/// id; it is not a user-facing idempotency key.
|
||||
pub async fn remove_function_name<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
name: &str,
|
||||
current: &Function,
|
||||
) -> Result<()> {
|
||||
if name.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "function name removal name must be non-empty".into(),
|
||||
});
|
||||
}
|
||||
|
||||
// Authority is the observed immutable Function id only; never send
|
||||
// signature, raw Function objects, Job fields, or user idempotency keys.
|
||||
let body = serde_json::json!({
|
||||
"name": name,
|
||||
"expected_current_function_id": current.id().as_str(),
|
||||
});
|
||||
let req_builder = client.post(REMOVE_PATH).json(&body);
|
||||
|
||||
catalog_mutation_with_retry(
|
||||
client,
|
||||
req_builder,
|
||||
REMOVE_HTTP_ERROR_MESSAGE,
|
||||
REMOVE_FUNCTION_ERROR_MESSAGE,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Revoke an exact Function via POST `/v1/functions/revoke`.
|
||||
///
|
||||
/// Direct synchronous administrator catalog set-bit: the wire body is exactly
|
||||
/// `{"function_id"}` from `function.id`. Only HTTP 204 means the set-bit
|
||||
/// completed; other 2xx are payload-free protocol [`Error::Http`] and are not
|
||||
/// retried or body-read. There is no empty-input validation because
|
||||
/// [`Function`] is already a validated exact handle.
|
||||
///
|
||||
/// Retry and explicit-code classification match remove. Sophon owns durable
|
||||
/// idempotent set-bit semantics; repeated logical calls that each receive 204
|
||||
/// succeed with no client already-revoked branch.
|
||||
pub async fn revoke_function<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
function: &Function,
|
||||
) -> Result<()> {
|
||||
let body = serde_json::json!({
|
||||
"function_id": function.id().as_str(),
|
||||
});
|
||||
let req_builder = client.post(REVOKE_PATH).json(&body);
|
||||
|
||||
catalog_mutation_with_retry(
|
||||
client,
|
||||
req_builder,
|
||||
REVOKE_HTTP_ERROR_MESSAGE,
|
||||
REVOKE_FUNCTION_ERROR_MESSAGE,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Shared remove/revoke catalog-mutation response classification.
|
||||
///
|
||||
/// Exact HTTP 204 succeeds without reading the body. Other 2xx are immediate
|
||||
/// payload-free [`Error::Http`]. Non-success bodies use explicit nonempty
|
||||
/// `error_code` as terminal [`Error::Function`], else configured retryable
|
||||
/// status retry, else payload-free [`Error::Http`]. Each caller supplies its
|
||||
/// own fixed sanitized messages.
|
||||
async fn catalog_mutation_with_retry<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
req_builder: RequestBuilder,
|
||||
http_error_message: &'static str,
|
||||
function_error_message: &'static str,
|
||||
) -> Result<()> {
|
||||
post_with_body_classification(
|
||||
client,
|
||||
req_builder,
|
||||
http_error_message,
|
||||
|status, request_id| {
|
||||
// Exact HTTP 204 completes the mutation; do not read or interpret any body.
|
||||
if status == StatusCode::NO_CONTENT {
|
||||
BeforeBody::Done(Ok(()))
|
||||
} else if status.is_success() {
|
||||
// Other 2xx are payload-free protocol failures from status alone.
|
||||
BeforeBody::Done(Err(Error::Http {
|
||||
source: http_error_message.into(),
|
||||
request_id: request_id.to_string(),
|
||||
status_code: Some(status),
|
||||
}))
|
||||
} else {
|
||||
BeforeBody::ReadBody
|
||||
}
|
||||
},
|
||||
|status, bytes, request_id| {
|
||||
// Explicit nonempty error_code wins over HTTP status and precludes retry.
|
||||
if let Some(code) = explicit_error_code(bytes) {
|
||||
return BodyAction::Done(Err(Error::Function {
|
||||
code,
|
||||
message: function_error_message.to_string(),
|
||||
}));
|
||||
}
|
||||
|
||||
if client.retry_config.statuses.contains(&status) {
|
||||
return BodyAction::RetryRequest;
|
||||
}
|
||||
|
||||
BodyAction::Done(Err(Error::Http {
|
||||
source: http_error_message.into(),
|
||||
request_id,
|
||||
status_code: Some(status),
|
||||
}))
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn decode_lookup_success(bytes: &[u8], request_id: String) -> Result<Function> {
|
||||
match serde_json::from_slice::<LookupSuccessResponse>(bytes) {
|
||||
Ok(body) => Ok(body.function),
|
||||
Err(_) => Err(Error::Http {
|
||||
source: LOOKUP_INVALID_SUCCESS_MESSAGE.into(),
|
||||
request_id,
|
||||
status_code: None,
|
||||
}),
|
||||
}
|
||||
}
|
||||
+962
-19
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,247 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Shared Remote HTTP transport primitives for explicit `error_code` classification.
|
||||
//!
|
||||
//! These helpers implement the same retry-bucket policy as
|
||||
//! [`RestfulLanceDbClient::send_with_retry`]: connect → connect budget; timeout /
|
||||
//! body / decode (including non-success body reads) → read budget; configured
|
||||
//! retryable status without an explicit nonempty `error_code` → request budget;
|
||||
//! other client/header errors are immediate. An explicit nonempty top-level
|
||||
//! `error_code` is terminal [`Error::Function`] and wins over HTTP status.
|
||||
//! Request/response payload bytes never enter error chains on the classified
|
||||
//! path.
|
||||
|
||||
use reqwest::{RequestBuilder, Response, StatusCode};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::error::{Error, FunctionErrorCode, Result};
|
||||
|
||||
use super::client::{HttpSend, RestfulLanceDbClient};
|
||||
use super::retry::RetryCounter;
|
||||
|
||||
/// Prepare a [`RetryCounter`] with one SDK-generated request id taken from the
|
||||
/// request builder. The same id is reused across every attempt.
|
||||
fn prepare_transport_retry<'a, S: HttpSend>(
|
||||
client: &'a RestfulLanceDbClient<S>,
|
||||
req_builder: &RequestBuilder,
|
||||
) -> Result<RetryCounter<'a>> {
|
||||
let tmp_req = req_builder.try_clone().ok_or_else(|| Error::Runtime {
|
||||
message: "Attempted to retry a request that cannot be cloned".to_string(),
|
||||
})?;
|
||||
let (_, built) = tmp_req.build_split();
|
||||
let mut built = built.map_err(|e| Error::Runtime {
|
||||
message: format!("Failed to build request: {}", e),
|
||||
})?;
|
||||
let request_id = client.extract_request_id(&mut built);
|
||||
Ok(RetryCounter::new(&client.retry_config, request_id))
|
||||
}
|
||||
|
||||
/// Classify a send-attempt error using the same buckets as `send_with_retry`.
|
||||
///
|
||||
/// Returns `Ok(())` when the caller should sleep and retry. Returns `Err` for
|
||||
/// nonretryable failures or when a retry budget is exhausted (no extra attempt).
|
||||
fn classify_transport_send_error(retry_counter: &mut RetryCounter<'_>, err: Error) -> Result<()> {
|
||||
match err {
|
||||
Error::Http {
|
||||
source,
|
||||
request_id,
|
||||
status_code,
|
||||
} => match source.downcast::<reqwest::Error>() {
|
||||
Ok(reqwest_err) if reqwest_err.is_connect() => {
|
||||
retry_counter.increment_connect_failures(*reqwest_err)
|
||||
}
|
||||
Ok(reqwest_err)
|
||||
if reqwest_err.is_timeout() || reqwest_err.is_body() || reqwest_err.is_decode() =>
|
||||
{
|
||||
retry_counter.increment_read_failures(*reqwest_err)
|
||||
}
|
||||
Ok(reqwest_err) => Err(Error::Http {
|
||||
source: Box::new(*reqwest_err),
|
||||
request_id,
|
||||
status_code,
|
||||
}),
|
||||
Err(source) => Err(Error::Http {
|
||||
source,
|
||||
request_id,
|
||||
status_code,
|
||||
}),
|
||||
},
|
||||
// Header-provider / client failures are not transport retries.
|
||||
other => Err(other),
|
||||
}
|
||||
}
|
||||
|
||||
/// Decode a stable category only from an explicit nonempty string `error_code`.
|
||||
/// Missing, empty, wrong-type, nested-only, or non-JSON bodies yield [`None`].
|
||||
pub(super) fn explicit_error_code(bytes: &[u8]) -> Option<FunctionErrorCode> {
|
||||
let value: Value = serde_json::from_slice(bytes).ok()?;
|
||||
let code = value.get("error_code")?;
|
||||
match code {
|
||||
Value::String(raw) if !raw.is_empty() => {
|
||||
serde_json::from_value(Value::String(raw.clone())).ok()
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Send a clonable request with explicit-`error_code` classification.
|
||||
///
|
||||
/// Uses the sensitive attempt path (no request body/header logging). Successful
|
||||
/// 2xx responses are returned **unconsumed** so callers can run their own
|
||||
/// success decoders. Non-success bodies are inspected for an explicit nonempty
|
||||
/// `error_code` before status-based retry. Payload bytes never enter
|
||||
/// [`Error::Function`], [`Error::Http`], or exhausted [`Error::Retry`] chains.
|
||||
pub(super) async fn send_with_explicit_error_code<S: HttpSend>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
req_builder: RequestBuilder,
|
||||
http_error_message: &'static str,
|
||||
function_error_message: &'static str,
|
||||
) -> Result<(String, Response)> {
|
||||
let mut retry_counter = prepare_transport_retry(client, &req_builder)?;
|
||||
|
||||
loop {
|
||||
let attempt = req_builder.try_clone().ok_or_else(|| Error::Runtime {
|
||||
message: "Attempted to retry a request that cannot be cloned".to_string(),
|
||||
})?;
|
||||
|
||||
let rsp = match client
|
||||
.send_attempt_with_request_id(attempt, &retry_counter.request_id)
|
||||
.await
|
||||
{
|
||||
Ok(rsp) => rsp,
|
||||
Err(err) => {
|
||||
classify_transport_send_error(&mut retry_counter, err)?;
|
||||
tokio::time::sleep(retry_counter.next_sleep_time()).await;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let status = rsp.status();
|
||||
if status.is_success() {
|
||||
// Leave the body unconsumed for Arrow / plan success decoders.
|
||||
return Ok((retry_counter.request_id.clone(), rsp));
|
||||
}
|
||||
|
||||
// Inspect the body before deciding whether the status is retryable.
|
||||
let bytes = match rsp.bytes().await {
|
||||
Ok(bytes) => bytes,
|
||||
Err(err) => {
|
||||
// Response body/decode failures share the read budget with
|
||||
// send-time timeout/body/decode errors.
|
||||
retry_counter.increment_read_failures(err)?;
|
||||
tokio::time::sleep(retry_counter.next_sleep_time()).await;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if let Some(code) = explicit_error_code(&bytes) {
|
||||
return Err(Error::Function {
|
||||
code,
|
||||
message: function_error_message.to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
if client.retry_config.statuses.contains(&status) {
|
||||
let source = Error::Http {
|
||||
source: http_error_message.into(),
|
||||
request_id: retry_counter.request_id.clone(),
|
||||
status_code: Some(status),
|
||||
};
|
||||
retry_counter.increment_request_failures(source)?;
|
||||
tokio::time::sleep(retry_counter.next_sleep_time()).await;
|
||||
continue;
|
||||
}
|
||||
|
||||
return Err(Error::Http {
|
||||
source: http_error_message.into(),
|
||||
request_id: retry_counter.request_id.clone(),
|
||||
status_code: Some(status),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
/// Decision before reading response bytes (catalog mutations / lookup).
|
||||
pub(super) enum BeforeBody<T> {
|
||||
/// Finish without reading or interpreting any body (HTTP 204 mutations).
|
||||
Done(Result<T>),
|
||||
/// Read bytes and continue classification.
|
||||
ReadBody,
|
||||
}
|
||||
|
||||
/// How a classified body is treated after a successful read.
|
||||
pub(super) enum BodyAction<T> {
|
||||
/// Terminal success or failure for this attempt.
|
||||
Done(Result<T>),
|
||||
/// Configured retryable status without an explicit `error_code`: consume
|
||||
/// the request budget and retry with the same request id and body.
|
||||
RetryRequest,
|
||||
}
|
||||
|
||||
/// Shared POST retry loop used by Function-catalog helpers.
|
||||
///
|
||||
/// Sensitive attempt sending logs no body/header selectors. One SDK-generated
|
||||
/// request id and the exact cloned JSON body are reused across attempts.
|
||||
/// Callers supply before/after body classifiers; this loop owns budgets only.
|
||||
pub(super) async fn post_with_body_classification<S, Before, After, T>(
|
||||
client: &RestfulLanceDbClient<S>,
|
||||
req_builder: RequestBuilder,
|
||||
http_error_message: &'static str,
|
||||
mut before_body: Before,
|
||||
mut after_body: After,
|
||||
) -> Result<T>
|
||||
where
|
||||
S: HttpSend,
|
||||
Before: FnMut(StatusCode, &str) -> BeforeBody<T>,
|
||||
After: FnMut(StatusCode, &[u8], String) -> BodyAction<T>,
|
||||
{
|
||||
let mut retry_counter = prepare_transport_retry(client, &req_builder)?;
|
||||
|
||||
loop {
|
||||
let attempt = req_builder.try_clone().ok_or_else(|| Error::Runtime {
|
||||
message: "Attempted to retry a request that cannot be cloned".to_string(),
|
||||
})?;
|
||||
|
||||
let rsp = match client
|
||||
.send_attempt_with_request_id(attempt, &retry_counter.request_id)
|
||||
.await
|
||||
{
|
||||
Ok(rsp) => rsp,
|
||||
Err(err) => {
|
||||
classify_transport_send_error(&mut retry_counter, err)?;
|
||||
tokio::time::sleep(retry_counter.next_sleep_time()).await;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let status = rsp.status();
|
||||
match before_body(status, &retry_counter.request_id) {
|
||||
BeforeBody::Done(result) => return result,
|
||||
BeforeBody::ReadBody => {}
|
||||
}
|
||||
|
||||
let bytes = match rsp.bytes().await {
|
||||
Ok(bytes) => bytes,
|
||||
Err(err) => {
|
||||
// Response body/decode failures share the read budget with
|
||||
// send-time timeout/body/decode errors.
|
||||
retry_counter.increment_read_failures(err)?;
|
||||
tokio::time::sleep(retry_counter.next_sleep_time()).await;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
match after_body(status, &bytes, retry_counter.request_id.clone()) {
|
||||
BodyAction::Done(result) => return result,
|
||||
BodyAction::RetryRequest => {
|
||||
let source = Error::Http {
|
||||
source: http_error_message.into(),
|
||||
request_id: retry_counter.request_id.clone(),
|
||||
status_code: Some(status),
|
||||
};
|
||||
retry_counter.increment_request_failures(source)?;
|
||||
tokio::time::sleep(retry_counter.next_sleep_time()).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+1243
-8
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,441 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! RED runtime contract tests for Native add-columns schema admission (B4h).
|
||||
//!
|
||||
//! Caller-authored Arrow field metadata under
|
||||
//! [`crate::function::GENERATED_COLUMN_METADATA_KEY`] must not enter table
|
||||
//! schema state through general-purpose Native `add_columns`. Generated
|
||||
//! definitions are Job-owned. Schema-bearing transforms (`BatchUDF`, `Stream`,
|
||||
//! `Reader`, `AllNulls`) currently accept and persist reserved top-level field
|
||||
//! metadata; these tests pin the missing pre-consumption admission guard.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator, RecordBatchReader, StringArray};
|
||||
use arrow_schema::{ArrowError, DataType, Field, Schema, SchemaRef};
|
||||
use datafusion_physical_plan::stream::RecordBatchStreamAdapter;
|
||||
use futures::{TryStreamExt, stream};
|
||||
use lance::dataset::{BatchUDF, NewColumnTransform};
|
||||
use tempfile::TempDir;
|
||||
|
||||
use crate::connection::ConnectBuilder;
|
||||
use crate::error::Error;
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
|
||||
};
|
||||
use crate::query::{ExecutableQuery, QueryBase, Select};
|
||||
use crate::table::Table;
|
||||
|
||||
const ID: &str = "id";
|
||||
const ORDINARY: &str = "ordinary";
|
||||
const GEN_OUT: &str = "gen_out";
|
||||
const ORDINARY_META_KEY: &str = "unit";
|
||||
const ORDINARY_META_VALUE: &str = "label";
|
||||
const FN_ID: &str = "fn.exact.b4h.add_columns.literal";
|
||||
const MALFORMED_MARKER: &str = "SENSITIVE_B4H_ADD_COLUMNS_METADATA_MARKER_4f8a_c3e2";
|
||||
|
||||
struct Fixture {
|
||||
_tmp: TempDir,
|
||||
table: Table,
|
||||
}
|
||||
|
||||
/// Counts [`RecordBatchReader::next`] calls. [`RecordBatchReader::schema`] is free.
|
||||
struct ObservableReader {
|
||||
inner: Box<dyn RecordBatchReader + Send>,
|
||||
next_calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
impl ObservableReader {
|
||||
fn wrap(
|
||||
inner: Box<dyn RecordBatchReader + Send>,
|
||||
next_calls: Arc<AtomicUsize>,
|
||||
) -> Box<dyn RecordBatchReader + Send> {
|
||||
Box::new(Self { inner, next_calls })
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for ObservableReader {
|
||||
type Item = Result<RecordBatch, ArrowError>;
|
||||
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.next_calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.inner.next()
|
||||
}
|
||||
}
|
||||
|
||||
impl RecordBatchReader for ObservableReader {
|
||||
fn schema(&self) -> SchemaRef {
|
||||
self.inner.schema()
|
||||
}
|
||||
}
|
||||
|
||||
fn literal_definition(output_field_id: i32) -> GeneratedColumnDefinition {
|
||||
let function = Function::new(
|
||||
FunctionId::try_new(FN_ID).unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("literal-only")])) as arrow_array::ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
GeneratedColumnDefinition::try_new(output_field_id, call, 1, 1).unwrap()
|
||||
}
|
||||
|
||||
fn valid_reserved_payload() -> String {
|
||||
literal_definition(1).to_metadata_json().unwrap()
|
||||
}
|
||||
|
||||
fn malformed_reserved_payload() -> String {
|
||||
format!(
|
||||
r#"{{"format_version":1,"output_field_id":1,"function_call":"{MALFORMED_MARKER}","dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
)
|
||||
}
|
||||
|
||||
fn seed_batch() -> RecordBatch {
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(ID, DataType::Int32, false),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
]));
|
||||
RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1, 2])),
|
||||
Arc::new(StringArray::from(vec![Some("a"), Some("b")])),
|
||||
],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn field_with_metadata(metadata: HashMap<String, String>) -> Field {
|
||||
Field::new(GEN_OUT, DataType::Int32, true).with_metadata(metadata)
|
||||
}
|
||||
|
||||
fn reserved_field(payload: &str) -> Field {
|
||||
field_with_metadata(
|
||||
[(
|
||||
GENERATED_COLUMN_METADATA_KEY.to_string(),
|
||||
payload.to_string(),
|
||||
)]
|
||||
.into(),
|
||||
)
|
||||
}
|
||||
|
||||
fn ordinary_metadata_field() -> Field {
|
||||
field_with_metadata(
|
||||
[(
|
||||
ORDINARY_META_KEY.to_string(),
|
||||
ORDINARY_META_VALUE.to_string(),
|
||||
)]
|
||||
.into(),
|
||||
)
|
||||
}
|
||||
|
||||
fn values_batch(schema: SchemaRef, values: Vec<i32>) -> RecordBatch {
|
||||
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(values))]).unwrap()
|
||||
}
|
||||
|
||||
fn boxed_reader(batch: RecordBatch) -> Box<dyn RecordBatchReader + Send> {
|
||||
let schema = batch.schema();
|
||||
Box::new(RecordBatchIterator::new(
|
||||
vec![Ok(batch)].into_iter(),
|
||||
schema,
|
||||
))
|
||||
}
|
||||
|
||||
fn observable_stream(
|
||||
batch: RecordBatch,
|
||||
yield_calls: Arc<AtomicUsize>,
|
||||
) -> datafusion_physical_plan::SendableRecordBatchStream {
|
||||
let schema = batch.schema();
|
||||
let counter = yield_calls.clone();
|
||||
Box::pin(RecordBatchStreamAdapter::new(
|
||||
schema,
|
||||
stream::once(async move {
|
||||
counter.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(batch)
|
||||
}),
|
||||
))
|
||||
}
|
||||
|
||||
fn assert_not_supported_redacted(err: &Error, label: &str, payload: &str) {
|
||||
match err {
|
||||
Error::NotSupported { message } => {
|
||||
let rendered = format!("{err}\n{err:?}\n{message}");
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: leaked metadata wire key: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(payload),
|
||||
"{label}: leaked raw payload: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(FN_ID),
|
||||
"{label}: leaked Function ID: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(GEN_OUT),
|
||||
"{label}: leaked output field name: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(MALFORMED_MARKER),
|
||||
"{label}: leaked malformed marker: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
message.to_lowercase().contains("generated")
|
||||
|| message.to_lowercase().contains("job"),
|
||||
"{label}: message must describe Job-owned generated-column boundary: {message}"
|
||||
);
|
||||
}
|
||||
other => panic!("{label}: expected Error::NotSupported, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
async fn create_table(name: &str) -> Fixture {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().to_str().unwrap().to_string();
|
||||
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
|
||||
let table = conn
|
||||
.create_table(name, seed_batch())
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
Fixture { _tmp: tmp, table }
|
||||
}
|
||||
|
||||
async fn snapshot_rows(table: &Table) -> Vec<(i32, String)> {
|
||||
let batches: Vec<RecordBatch> = table
|
||||
.query()
|
||||
.select(Select::columns(&[ID, ORDINARY]))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.try_collect()
|
||||
.await
|
||||
.unwrap();
|
||||
let mut rows = Vec::new();
|
||||
for batch in batches {
|
||||
let ids = batch
|
||||
.column(0)
|
||||
.as_any()
|
||||
.downcast_ref::<Int32Array>()
|
||||
.unwrap();
|
||||
let ordinary = batch
|
||||
.column(1)
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.unwrap();
|
||||
for i in 0..batch.num_rows() {
|
||||
rows.push((ids.value(i), ordinary.value(i).to_string()));
|
||||
}
|
||||
}
|
||||
rows.sort_by_key(|(id, _)| *id);
|
||||
rows
|
||||
}
|
||||
|
||||
async fn assert_table_unchanged(
|
||||
table: &Table,
|
||||
version_before: u64,
|
||||
schema_before: &Schema,
|
||||
rows_before: &[(i32, String)],
|
||||
) {
|
||||
assert_eq!(table.version().await.unwrap(), version_before);
|
||||
let schema_after = table.schema().await.unwrap();
|
||||
assert_eq!(schema_after.as_ref(), schema_before);
|
||||
assert!(
|
||||
schema_after.field_with_name(GEN_OUT).is_err(),
|
||||
"rejected add_columns must leave column `{GEN_OUT}` absent"
|
||||
);
|
||||
assert_eq!(snapshot_rows(table).await, rows_before);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn batch_udf_rejects_valid_reserved_before_mapper() {
|
||||
let fixture = create_table("b4h_batch_udf").await;
|
||||
let table = &fixture.table;
|
||||
let version_before = table.version().await.unwrap();
|
||||
let schema_before = table.schema().await.unwrap();
|
||||
let rows_before = snapshot_rows(table).await;
|
||||
|
||||
let payload = valid_reserved_payload();
|
||||
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
|
||||
let mapper_schema = output_schema.clone();
|
||||
let mapper_calls = Arc::new(AtomicUsize::new(0));
|
||||
let calls = mapper_calls.clone();
|
||||
let udf = BatchUDF {
|
||||
mapper: Box::new(move |batch: &RecordBatch| {
|
||||
calls.fetch_add(1, Ordering::SeqCst);
|
||||
let values = Int32Array::from(vec![Some(10); batch.num_rows()]);
|
||||
Ok(RecordBatch::try_new(
|
||||
mapper_schema.clone(),
|
||||
vec![Arc::new(values)],
|
||||
)?)
|
||||
}),
|
||||
output_schema,
|
||||
result_checkpoint: None,
|
||||
};
|
||||
|
||||
// Public Table::add_columns builder path.
|
||||
let err = table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::BatchUDF(udf))
|
||||
.execute()
|
||||
.await
|
||||
.expect_err("BatchUDF must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(&err, "BatchUDF reserved admission", &payload);
|
||||
assert_eq!(
|
||||
mapper_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"rejection must occur before invoking the BatchUDF mapper"
|
||||
);
|
||||
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stream_rejects_malformed_reserved_before_yield() {
|
||||
let fixture = create_table("b4h_stream").await;
|
||||
let table = &fixture.table;
|
||||
let version_before = table.version().await.unwrap();
|
||||
let schema_before = table.schema().await.unwrap();
|
||||
let rows_before = snapshot_rows(table).await;
|
||||
|
||||
let payload = malformed_reserved_payload();
|
||||
assert!(payload.contains(MALFORMED_MARKER));
|
||||
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
|
||||
let yield_calls = Arc::new(AtomicUsize::new(0));
|
||||
let stream = observable_stream(
|
||||
values_batch(output_schema, vec![10, 20]),
|
||||
yield_calls.clone(),
|
||||
);
|
||||
|
||||
// Direct experimental BaseTable::add_columns path.
|
||||
let err = table
|
||||
.base_table()
|
||||
.add_columns(NewColumnTransform::Stream(stream), None)
|
||||
.await
|
||||
.expect_err("Stream must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(&err, "Stream reserved admission", &payload);
|
||||
assert_eq!(
|
||||
yield_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"rejection must occur before polling/yielding the user Stream"
|
||||
);
|
||||
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reader_rejects_valid_reserved_before_next() {
|
||||
let fixture = create_table("b4h_reader").await;
|
||||
let table = &fixture.table;
|
||||
let version_before = table.version().await.unwrap();
|
||||
let schema_before = table.schema().await.unwrap();
|
||||
let rows_before = snapshot_rows(table).await;
|
||||
|
||||
let payload = valid_reserved_payload();
|
||||
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
|
||||
let next_calls = Arc::new(AtomicUsize::new(0));
|
||||
let reader = ObservableReader::wrap(
|
||||
boxed_reader(values_batch(output_schema, vec![10, 20])),
|
||||
next_calls.clone(),
|
||||
);
|
||||
|
||||
let err = table
|
||||
.base_table()
|
||||
.add_columns(NewColumnTransform::Reader(reader), None)
|
||||
.await
|
||||
.expect_err("Reader must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(&err, "Reader reserved admission", &payload);
|
||||
assert_eq!(
|
||||
next_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"rejection must occur before RecordBatchReader::next"
|
||||
);
|
||||
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn all_nulls_rejects_malformed_reserved_before_commit() {
|
||||
let fixture = create_table("b4h_all_nulls").await;
|
||||
let table = &fixture.table;
|
||||
let version_before = table.version().await.unwrap();
|
||||
let schema_before = table.schema().await.unwrap();
|
||||
let rows_before = snapshot_rows(table).await;
|
||||
|
||||
let payload = malformed_reserved_payload();
|
||||
let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)]));
|
||||
|
||||
let err = table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::AllNulls(output_schema))
|
||||
.execute()
|
||||
.await
|
||||
.expect_err("AllNulls must reject reserved generated-column metadata");
|
||||
assert_not_supported_redacted(&err, "AllNulls reserved admission", &payload);
|
||||
assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sql_expressions_add_columns_still_succeeds() {
|
||||
let fixture = create_table("b4h_sql_control").await;
|
||||
let table = &fixture.table;
|
||||
|
||||
table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::SqlExpressions(vec![(
|
||||
"doubled".into(),
|
||||
"id * 2".into(),
|
||||
)]))
|
||||
.execute()
|
||||
.await
|
||||
.expect("ordinary SqlExpressions add_columns must remain supported");
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
assert!(schema.field_with_name("doubled").is_ok());
|
||||
assert!(schema.field_with_name(GEN_OUT).is_err());
|
||||
assert!(
|
||||
!schema
|
||||
.field_with_name("doubled")
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn schema_bearing_ordinary_metadata_is_preserved() {
|
||||
let fixture = create_table("b4h_ordinary_meta").await;
|
||||
let table = &fixture.table;
|
||||
let output_schema = Arc::new(Schema::new(vec![ordinary_metadata_field()]));
|
||||
|
||||
// AllNulls is schema-bearing and metadata-only; proves ordinary metadata
|
||||
// remains accepted so a later guard cannot reject every field metadata map.
|
||||
table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::AllNulls(output_schema))
|
||||
.execute()
|
||||
.await
|
||||
.expect("ordinary non-reserved field metadata must remain accepted");
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
let md = schema.field_with_name(GEN_OUT).unwrap().metadata();
|
||||
assert_eq!(
|
||||
md.get(ORDINARY_META_KEY).map(String::as_str),
|
||||
Some(ORDINARY_META_VALUE)
|
||||
);
|
||||
assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY));
|
||||
}
|
||||
@@ -10,6 +10,7 @@ use serde::{Deserialize, Serialize};
|
||||
use crate::data::scannable::Scannable;
|
||||
use crate::data::scannable::scannable_with_embeddings;
|
||||
use crate::embeddings::EmbeddingRegistry;
|
||||
use crate::function::schema_admission::reject_caller_authored_generated_column_schema_on_overwrite;
|
||||
use crate::table::datafusion::cast::cast_to_table_schema;
|
||||
use crate::table::datafusion::reject_nan::reject_nan_vectors;
|
||||
use crate::table::datafusion::scannable_exec::ScannableExec;
|
||||
@@ -151,6 +152,27 @@ impl AddDataBuilder {
|
||||
self.parent.clone().add(self).await
|
||||
}
|
||||
|
||||
/// Effective overwrite for schema-replacement admission and planning.
|
||||
///
|
||||
/// True when either `WriteOptions.lance_write_params.mode` is
|
||||
/// [`WriteMode::Overwrite`] or [`AddDataMode::Overwrite`] is set.
|
||||
pub(crate) fn is_effective_overwrite(&self) -> bool {
|
||||
self.write_options
|
||||
.lance_write_params
|
||||
.as_ref()
|
||||
.is_some_and(|p| matches!(p.mode, WriteMode::Overwrite))
|
||||
|| matches!(self.mode, AddDataMode::Overwrite)
|
||||
}
|
||||
|
||||
/// Borrowed preflight: reject reserved generated-column metadata on
|
||||
/// effective overwrite before source scan or table/network work.
|
||||
pub(crate) fn admit_input_schema(&self) -> Result<()> {
|
||||
reject_caller_authored_generated_column_schema_on_overwrite(
|
||||
self.data.schema().as_ref(),
|
||||
self.is_effective_overwrite(),
|
||||
)
|
||||
}
|
||||
|
||||
/// Build a DataFusion execution plan that applies embeddings, casts data to
|
||||
/// the table schema, and optionally rejects NaN vectors.
|
||||
///
|
||||
@@ -161,12 +183,7 @@ impl AddDataBuilder {
|
||||
table_schema: &Schema,
|
||||
table_def: &TableDefinition,
|
||||
) -> Result<PreprocessingOutput> {
|
||||
let overwrite = self
|
||||
.write_options
|
||||
.lance_write_params
|
||||
.as_ref()
|
||||
.is_some_and(|p| matches!(p.mode, WriteMode::Overwrite))
|
||||
|| matches!(self.mode, AddDataMode::Overwrite);
|
||||
let overwrite = self.is_effective_overwrite();
|
||||
|
||||
if !overwrite {
|
||||
validate_schema(&self.data.schema(), table_schema)?;
|
||||
|
||||
@@ -0,0 +1,797 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! RED runtime contract tests for Native append invalidation (B4b).
|
||||
//!
|
||||
//! These tests pin Native Table API and DataFusion SQL insert behavior for
|
||||
//! generated-column dependency-epoch invalidation. They use real local Native
|
||||
//! tables and existing public/internal APIs; Lance commits and query guards are
|
||||
//! not mocked.
|
||||
|
||||
use std::collections::HashSet;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use arrow_array::{Array, Int32Array, RecordBatch, RecordBatchIterator, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use datafusion::prelude::SessionContext;
|
||||
use futures::TryStreamExt;
|
||||
use lance::dataset::{WriteMode, WriteParams};
|
||||
use tempfile::TempDir;
|
||||
|
||||
use crate::connection::ConnectBuilder;
|
||||
use crate::error::{Error, FunctionErrorCode};
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
|
||||
GeneratedColumnStatus,
|
||||
};
|
||||
use crate::query::{ExecutableQuery, QueryBase, Select};
|
||||
use crate::table::datafusion::BaseTableAdapter;
|
||||
use crate::table::{AddDataMode, Table, WriteOptions};
|
||||
|
||||
const GEN_OUT: &str = "gen_out";
|
||||
const ORDINARY: &str = "ordinary";
|
||||
const INITIAL_DEPENDENCY_EPOCH: u64 = 3;
|
||||
const INITIAL_MATERIALIZED_EPOCH: u64 = 3;
|
||||
const MALFORMED_MARKER: &str = "SENSITIVE_B4B_APPEND_METADATA_MARKER_9f2c_a81d";
|
||||
|
||||
struct Fixture {
|
||||
_tmp: TempDir,
|
||||
table: Table,
|
||||
table_name: String,
|
||||
uri: String,
|
||||
}
|
||||
|
||||
fn literal_only_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new("fn.exact.b4b.append.literal").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn literal_only_definition(
|
||||
output_field_id: i32,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
let function = literal_only_function();
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("literal-only")])) as arrow_array::ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
assert!(
|
||||
call.arguments()
|
||||
.iter()
|
||||
.all(|(_, argument)| argument.field_id().is_none()),
|
||||
"fixture must be literal-only so row-set coverage, not field dependency, drives invalidation"
|
||||
);
|
||||
GeneratedColumnDefinition::try_new(output_field_id, call, dependency_epoch, materialized_epoch)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
async fn create_table_with_complete_literal_generated(name: &str) -> Fixture {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().to_str().unwrap().to_string();
|
||||
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
Arc::new(StringArray::from(vec![Some("seed")])),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
let table = conn.create_table(name, batch).execute().await.unwrap();
|
||||
|
||||
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
|
||||
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
|
||||
let definition = literal_only_definition(
|
||||
field_id,
|
||||
INITIAL_DEPENDENCY_EPOCH,
|
||||
INITIAL_MATERIALIZED_EPOCH,
|
||||
);
|
||||
let json = definition.to_metadata_json().unwrap();
|
||||
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
|
||||
table
|
||||
.as_native()
|
||||
.expect("generated-column fixture planting requires a Native table"),
|
||||
GEN_OUT,
|
||||
json,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
table.generated_column_status(GEN_OUT).await.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
let planted = read_generated_definition(&table).await;
|
||||
assert!(
|
||||
planted
|
||||
.function_call()
|
||||
.arguments()
|
||||
.iter()
|
||||
.all(|(_, argument)| argument.field_id().is_none()),
|
||||
"planted metadata must remain literal-only"
|
||||
);
|
||||
|
||||
Fixture {
|
||||
_tmp: tmp,
|
||||
table,
|
||||
table_name: name.to_string(),
|
||||
uri,
|
||||
}
|
||||
}
|
||||
|
||||
async fn create_ordinary_table(name: &str) -> Fixture {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().to_str().unwrap().to_string();
|
||||
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
Arc::new(StringArray::from(vec![Some("seed")])),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
let table = conn.create_table(name, batch).execute().await.unwrap();
|
||||
Fixture {
|
||||
_tmp: tmp,
|
||||
table,
|
||||
table_name: name.to_string(),
|
||||
uri,
|
||||
}
|
||||
}
|
||||
|
||||
async fn read_generated_definition(table: &Table) -> GeneratedColumnDefinition {
|
||||
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
|
||||
snapshot
|
||||
.field(GEN_OUT)
|
||||
.expect(GEN_OUT)
|
||||
.generated_column_definition()
|
||||
.expect("generated metadata must decode")
|
||||
.expect("generated metadata must be present")
|
||||
}
|
||||
|
||||
fn ordinary_rows_batch(values: &[&str]) -> RecordBatch {
|
||||
let schema = Arc::new(Schema::new(vec![Field::new(
|
||||
ORDINARY,
|
||||
DataType::Utf8,
|
||||
true,
|
||||
)]));
|
||||
RecordBatch::try_new(
|
||||
schema,
|
||||
vec![Arc::new(StringArray::from(
|
||||
values.iter().map(|value| Some(*value)).collect::<Vec<_>>(),
|
||||
))],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn full_rows_batch(gen_values: &[Option<i32>], ordinary_values: &[&str]) -> RecordBatch {
|
||||
assert_eq!(gen_values.len(), ordinary_values.len());
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
]));
|
||||
RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(gen_values.to_vec())),
|
||||
Arc::new(StringArray::from(
|
||||
ordinary_values
|
||||
.iter()
|
||||
.map(|value| Some(*value))
|
||||
.collect::<Vec<_>>(),
|
||||
)),
|
||||
],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
async fn ordinary_values(table: &Table) -> HashSet<String> {
|
||||
let batches = table
|
||||
.query()
|
||||
.select(Select::columns(&[ORDINARY]))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
let mut values = HashSet::new();
|
||||
for batch in batches {
|
||||
let column = batch
|
||||
.column(0)
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.unwrap();
|
||||
for index in 0..column.len() {
|
||||
if !column.is_null(index) {
|
||||
values.insert(column.value(index).to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
values
|
||||
}
|
||||
|
||||
fn assert_generated_column_incomplete(err: &Error, label: &str) {
|
||||
match err {
|
||||
Error::Function {
|
||||
code: FunctionErrorCode::GeneratedColumnIncomplete,
|
||||
..
|
||||
} => {}
|
||||
other => panic!("{label}: expected generated_column_incomplete, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_not_supported(err: &Error, label: &str) {
|
||||
assert!(
|
||||
matches!(err, Error::NotSupported { .. }),
|
||||
"{label}: expected NotSupported, got {err:?}"
|
||||
);
|
||||
}
|
||||
|
||||
fn assert_invalid_input_redacted(err: &Error, label: &str) {
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"{label}: expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let rendered = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!rendered.contains(MALFORMED_MARKER),
|
||||
"{label}: diagnostic echoed unique metadata marker: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: diagnostic echoed metadata wire key: {rendered}"
|
||||
);
|
||||
}
|
||||
|
||||
fn assert_conflict_error(err: &Error, label: &str) {
|
||||
match err {
|
||||
Error::Lance { source } => {
|
||||
assert!(
|
||||
matches!(
|
||||
source,
|
||||
lance::Error::IncompatibleTransaction { .. }
|
||||
| lance::Error::RetryableCommitConflict { .. }
|
||||
| lance::Error::CommitConflict { .. }
|
||||
),
|
||||
"{label}: expected Lance commit conflict category, got {source:?}"
|
||||
);
|
||||
}
|
||||
Error::Function {
|
||||
code: FunctionErrorCode::StaleOrConflictingInput,
|
||||
..
|
||||
} => {}
|
||||
other => panic!("{label}: expected conflict error category, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
fn from_datafusion_error(err: datafusion_common::DataFusionError) -> Error {
|
||||
Error::from(err)
|
||||
}
|
||||
|
||||
async fn sql_ctx_for(table: &Table, name: &str) -> SessionContext {
|
||||
let ctx = SessionContext::new();
|
||||
let provider = BaseTableAdapter::try_new(table.base_table().clone())
|
||||
.await
|
||||
.unwrap();
|
||||
ctx.register_table(name, Arc::new(provider)).unwrap();
|
||||
ctx
|
||||
}
|
||||
|
||||
async fn run_sql(ctx: &SessionContext, sql: &str) -> Result<(), Error> {
|
||||
match ctx.sql(sql).await {
|
||||
Err(err) => Err(from_datafusion_error(err)),
|
||||
Ok(df) => match df.collect().await {
|
||||
Ok(_) => Ok(()),
|
||||
Err(err) => Err(from_datafusion_error(err)),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn nonempty_table_api_append_invalidates_literal_only_generated_column() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_table_append").await;
|
||||
let before = read_generated_definition(&fixture.table).await;
|
||||
|
||||
fixture
|
||||
.table
|
||||
.add(ordinary_rows_batch(&["appended"]))
|
||||
.execute()
|
||||
.await
|
||||
.expect("non-empty Table API append must commit");
|
||||
|
||||
let values = ordinary_values(&fixture.table).await;
|
||||
assert!(values.contains("seed"));
|
||||
assert!(values.contains("appended"));
|
||||
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Incomplete
|
||||
);
|
||||
let after = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(after.dependency_epoch(), before.dependency_epoch() + 1);
|
||||
assert_eq!(after.materialized_epoch(), before.materialized_epoch());
|
||||
assert_eq!(after.output_field_id(), before.output_field_id());
|
||||
assert_eq!(after.function_call(), before.function_call());
|
||||
|
||||
let Err(err) = fixture
|
||||
.table
|
||||
.query()
|
||||
.select(Select::columns(&[GEN_OUT]))
|
||||
.execute()
|
||||
.await
|
||||
else {
|
||||
panic!("incomplete generated column query must fail");
|
||||
};
|
||||
assert_generated_column_incomplete(&err, "table api append query");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn nonempty_table_api_append_atomic_version_visibility() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_atomic_visibility").await;
|
||||
let previous_version = fixture.table.version().await.unwrap();
|
||||
let previous_rows = ordinary_values(&fixture.table).await;
|
||||
let previous_definition = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(
|
||||
previous_definition.status(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
|
||||
fixture
|
||||
.table
|
||||
.add(ordinary_rows_batch(&["atomic-new"]))
|
||||
.execute()
|
||||
.await
|
||||
.expect("non-empty append must commit");
|
||||
let new_version = fixture.table.version().await.unwrap();
|
||||
assert_ne!(new_version, previous_version);
|
||||
|
||||
// Exact new version: new rows + incomplete metadata together.
|
||||
let new_rows = ordinary_values(&fixture.table).await;
|
||||
assert!(new_rows.contains("atomic-new"));
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Incomplete
|
||||
);
|
||||
let new_definition = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(
|
||||
new_definition.dependency_epoch(),
|
||||
previous_definition.dependency_epoch() + 1
|
||||
);
|
||||
|
||||
// Immediately previous version: neither new rows nor incomplete metadata.
|
||||
fixture.table.checkout(previous_version).await.unwrap();
|
||||
assert_eq!(ordinary_values(&fixture.table).await, previous_rows);
|
||||
assert!(!ordinary_values(&fixture.table).await.contains("atomic-new"));
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
let checked_out = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(checked_out, previous_definition);
|
||||
fixture
|
||||
.table
|
||||
.query()
|
||||
.select(Select::columns(&[GEN_OUT]))
|
||||
.execute()
|
||||
.await
|
||||
.expect("previous complete version must remain readable");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn empty_table_api_append_leaves_complete_generated_column() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_empty_table_append").await;
|
||||
let before = read_generated_definition(&fixture.table).await;
|
||||
let rows_before = ordinary_values(&fixture.table).await;
|
||||
|
||||
fixture
|
||||
.table
|
||||
.add(RecordBatch::new_empty(Arc::new(Schema::new(vec![
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
]))))
|
||||
.execute()
|
||||
.await
|
||||
.expect("empty Table API append is a supported path");
|
||||
|
||||
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
let after = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(after, before);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn multipartition_table_api_append_advances_dependency_epoch_once() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_multipartition").await;
|
||||
let before = read_generated_definition(&fixture.table).await;
|
||||
|
||||
fixture
|
||||
.table
|
||||
.add(ordinary_rows_batch(&["p0", "p1", "p2", "p3"]))
|
||||
.write_parallelism(2)
|
||||
.execute()
|
||||
.await
|
||||
.expect("multi-partition append must commit");
|
||||
|
||||
let values = ordinary_values(&fixture.table).await;
|
||||
assert!(values.contains("p0"));
|
||||
assert!(values.contains("p3"));
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Incomplete
|
||||
);
|
||||
let after = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(
|
||||
after.dependency_epoch(),
|
||||
before.dependency_epoch() + 1,
|
||||
"multi-partition append must attach one whole-transaction patch"
|
||||
);
|
||||
assert_eq!(after.materialized_epoch(), before.materialized_epoch());
|
||||
assert_eq!(after.function_call(), before.function_call());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn table_api_overwrite_rejects_before_mutation_when_generated_column_present() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_table_overwrite").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_values(&fixture.table).await;
|
||||
let definition_before = read_generated_definition(&fixture.table).await;
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.add(full_rows_batch(&[Some(9)], &["overwrite"]))
|
||||
.mode(AddDataMode::Overwrite)
|
||||
.execute()
|
||||
.await
|
||||
.expect_err("overwrite must reject when any generated column is present");
|
||||
assert_not_supported(&err, "table api overwrite");
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
read_generated_definition(&fixture.table).await,
|
||||
definition_before
|
||||
);
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn table_api_effective_overwrite_from_add_data_mode_rejects_when_lance_params_append() {
|
||||
let fixture =
|
||||
create_table_with_complete_literal_generated("b4b_table_effective_overwrite").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_values(&fixture.table).await;
|
||||
let definition_before = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(definition_before.status(), GeneratedColumnStatus::Complete);
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.add(full_rows_batch(&[Some(9)], &["effective-overwrite"]))
|
||||
.mode(AddDataMode::Overwrite)
|
||||
.write_options(WriteOptions {
|
||||
lance_write_params: Some(WriteParams {
|
||||
mode: WriteMode::Append,
|
||||
..Default::default()
|
||||
}),
|
||||
})
|
||||
.execute()
|
||||
.await
|
||||
.expect_err(
|
||||
"AddDataMode::Overwrite must reject generated-table writes even when \
|
||||
explicit lance WriteParams.mode is Append",
|
||||
);
|
||||
assert_not_supported(&err, "table api effective overwrite");
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
read_generated_definition(&fixture.table).await,
|
||||
definition_before
|
||||
);
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ordinary_table_api_overwrite_still_supported() {
|
||||
let fixture = create_ordinary_table("b4b_ordinary_overwrite_control").await;
|
||||
|
||||
fixture
|
||||
.table
|
||||
.add(full_rows_batch(&[Some(42)], &["replaced"]))
|
||||
.mode(AddDataMode::Overwrite)
|
||||
.execute()
|
||||
.await
|
||||
.expect("ordinary tables must keep overwrite support");
|
||||
|
||||
let values = ordinary_values(&fixture.table).await;
|
||||
assert_eq!(values, HashSet::from(["replaced".to_string()]));
|
||||
assert_eq!(fixture.table.count_rows(None).await.unwrap(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn nonempty_sql_insert_invalidates_generated_column() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_sql_insert").await;
|
||||
let before = read_generated_definition(&fixture.table).await;
|
||||
let ctx = sql_ctx_for(&fixture.table, &fixture.table_name).await;
|
||||
|
||||
run_sql(
|
||||
&ctx,
|
||||
&format!(
|
||||
"INSERT INTO {} VALUES (CAST(NULL AS INT), 'sql-appended')",
|
||||
fixture.table_name
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("non-empty SQL INSERT must commit");
|
||||
|
||||
fixture.table.checkout_latest().await.unwrap();
|
||||
let values = ordinary_values(&fixture.table).await;
|
||||
assert!(values.contains("sql-appended"));
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Incomplete
|
||||
);
|
||||
let after = read_generated_definition(&fixture.table).await;
|
||||
assert_eq!(after.dependency_epoch(), before.dependency_epoch() + 1);
|
||||
assert_eq!(after.materialized_epoch(), before.materialized_epoch());
|
||||
assert_eq!(after.function_call(), before.function_call());
|
||||
|
||||
let Err(err) = fixture
|
||||
.table
|
||||
.query()
|
||||
.select(Select::columns(&[GEN_OUT]))
|
||||
.execute()
|
||||
.await
|
||||
else {
|
||||
panic!("SQL INSERT invalidation must trip generated query guard");
|
||||
};
|
||||
assert_generated_column_incomplete(&err, "sql insert query");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn empty_sql_insert_leaves_complete_generated_column() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_empty_sql_insert").await;
|
||||
let before = read_generated_definition(&fixture.table).await;
|
||||
let rows_before = ordinary_values(&fixture.table).await;
|
||||
|
||||
let conn = ConnectBuilder::new(&fixture.uri).execute().await.unwrap();
|
||||
let source_schema = Arc::new(Schema::new(vec![
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
]));
|
||||
let empty_reader: Box<dyn arrow_array::RecordBatchReader + Send> =
|
||||
Box::new(RecordBatchIterator::new(
|
||||
std::iter::empty::<Result<RecordBatch, arrow_schema::ArrowError>>(),
|
||||
source_schema,
|
||||
));
|
||||
let source = conn
|
||||
.create_table("empty_source", empty_reader)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let ctx = sql_ctx_for(&fixture.table, &fixture.table_name).await;
|
||||
let source_provider = BaseTableAdapter::try_new(source.base_table().clone())
|
||||
.await
|
||||
.unwrap();
|
||||
ctx.register_table("empty_source", Arc::new(source_provider))
|
||||
.unwrap();
|
||||
|
||||
run_sql(
|
||||
&ctx,
|
||||
&format!(
|
||||
"INSERT INTO {} SELECT * FROM empty_source",
|
||||
fixture.table_name
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect("empty SQL INSERT is a supported path");
|
||||
|
||||
fixture.table.checkout_latest().await.unwrap();
|
||||
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
assert_eq!(read_generated_definition(&fixture.table).await, before);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sql_insert_overwrite_rejects_before_mutation_when_generated_column_present() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_sql_overwrite").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_values(&fixture.table).await;
|
||||
let definition_before = read_generated_definition(&fixture.table).await;
|
||||
let ctx = sql_ctx_for(&fixture.table, &fixture.table_name).await;
|
||||
|
||||
let err = run_sql(
|
||||
&ctx,
|
||||
&format!(
|
||||
"INSERT OVERWRITE INTO {} VALUES (10, 'sql-overwrite')",
|
||||
fixture.table_name
|
||||
),
|
||||
)
|
||||
.await
|
||||
.expect_err("SQL INSERT OVERWRITE must reject when any generated column is present");
|
||||
assert_not_supported(&err, "sql insert overwrite");
|
||||
|
||||
fixture.table.checkout_latest().await.unwrap();
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
read_generated_definition(&fixture.table).await,
|
||||
definition_before
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn malformed_generated_metadata_rejects_append_before_mutation_and_redacts_marker() {
|
||||
let fixture = create_ordinary_table("b4b_malformed_preflight").await;
|
||||
let snapshot = fixture
|
||||
.table
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.unwrap();
|
||||
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
|
||||
let raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{MALFORMED_MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(raw.contains(MALFORMED_MARKER));
|
||||
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
|
||||
fixture
|
||||
.table
|
||||
.as_native()
|
||||
.expect("generated-column fixture planting requires a Native table"),
|
||||
GEN_OUT,
|
||||
raw.clone(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_values(&fixture.table).await;
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.add(ordinary_rows_batch(&["must-not-land"]))
|
||||
.execute()
|
||||
.await
|
||||
.expect_err("malformed generated metadata must fail closed before append visibility");
|
||||
assert_invalid_input_redacted(&err, "malformed append preflight");
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
|
||||
assert!(
|
||||
!ordinary_values(&fixture.table)
|
||||
.await
|
||||
.contains("must-not-land")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn concurrent_same_field_append_one_winner_one_conflict() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4b_concurrent_append").await;
|
||||
let conn = ConnectBuilder::new(&fixture.uri)
|
||||
.read_consistency_interval(Duration::from_secs(3600))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
let table_a = conn
|
||||
.open_table(&fixture.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
let table_b = conn
|
||||
.open_table(&fixture.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
let basis_version = table_a.version().await.unwrap();
|
||||
assert_eq!(table_b.version().await.unwrap(), basis_version);
|
||||
|
||||
let (result_a, result_b) = tokio::join!(
|
||||
table_a.add(ordinary_rows_batch(&["winner-a"])).execute(),
|
||||
table_b.add(ordinary_rows_batch(&["winner-b"])).execute(),
|
||||
);
|
||||
|
||||
let outcomes = [result_a, result_b];
|
||||
let wins = outcomes.iter().filter(|result| result.is_ok()).count();
|
||||
let losses = outcomes.iter().filter(|result| result.is_err()).count();
|
||||
assert_eq!(wins, 1, "exactly one same-basis append may publish");
|
||||
assert_eq!(losses, 1, "exactly one same-basis append must conflict");
|
||||
for result in &outcomes {
|
||||
if let Err(err) = result {
|
||||
assert_conflict_error(err, "concurrent same-field append loser");
|
||||
}
|
||||
}
|
||||
|
||||
let fresh = conn
|
||||
.open_table(&fixture.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
let values = ordinary_values(&fresh).await;
|
||||
assert!(values.contains("seed"));
|
||||
let has_a = values.contains("winner-a");
|
||||
let has_b = values.contains("winner-b");
|
||||
assert!(
|
||||
has_a ^ has_b,
|
||||
"only winner rows may be visible, got {values:?}"
|
||||
);
|
||||
assert_eq!(
|
||||
fresh.generated_column_status(GEN_OUT).await.unwrap(),
|
||||
GeneratedColumnStatus::Incomplete
|
||||
);
|
||||
let definition = read_generated_definition(&fresh).await;
|
||||
assert_eq!(definition.dependency_epoch(), INITIAL_DEPENDENCY_EPOCH + 1);
|
||||
assert_eq!(definition.materialized_epoch(), INITIAL_MATERIALIZED_EPOCH);
|
||||
}
|
||||
@@ -423,6 +423,7 @@ mod tests {
|
||||
use futures::TryStreamExt;
|
||||
use tempfile::tempdir;
|
||||
|
||||
use crate::JobResult;
|
||||
use crate::connect;
|
||||
use crate::connection::ConnectBuilder;
|
||||
use crate::index::Index;
|
||||
@@ -538,7 +539,8 @@ mod tests {
|
||||
assert_eq!(job.id(), None);
|
||||
// The build runs as a task, so the index need not exist yet; it must
|
||||
// once the job resolves.
|
||||
job.wait().await.unwrap();
|
||||
let result = job.wait().await.unwrap();
|
||||
assert_eq!(result, JobResult::None);
|
||||
assert_eq!(table.list_indices().await.unwrap().len(), 1);
|
||||
// Cancelling a finished job is a no-op.
|
||||
job.cancel().await.unwrap();
|
||||
@@ -570,10 +572,12 @@ mod tests {
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
for waiter in waiters {
|
||||
waiter.await.unwrap().unwrap();
|
||||
let result = waiter.await.unwrap().unwrap();
|
||||
assert_eq!(result, JobResult::None);
|
||||
}
|
||||
// A wait after the job settled still reports the same outcome.
|
||||
job.wait().await.unwrap();
|
||||
let late = job.wait().await.unwrap();
|
||||
assert_eq!(late, JobResult::None);
|
||||
assert_eq!(table.list_indices().await.unwrap().len(), 1);
|
||||
}
|
||||
|
||||
@@ -716,7 +720,7 @@ mod tests {
|
||||
match job.wait().await {
|
||||
Err(crate::Error::JobCancelled { .. }) => {}
|
||||
// The build may finish before the abort lands.
|
||||
Ok(()) => {}
|
||||
Ok(JobResult::None) => {}
|
||||
other => panic!("unexpected job outcome: {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -19,7 +19,7 @@ use datafusion_physical_plan::{
|
||||
};
|
||||
use futures::TryStreamExt;
|
||||
use lance::Dataset;
|
||||
use lance::dataset::transaction::{Operation, Transaction};
|
||||
use lance::dataset::transaction::{Operation, SchemaMetadataUpdates, Transaction};
|
||||
use lance::dataset::{CommitBuilder, InsertBuilder, WriteParams, WriteProgressFn};
|
||||
use lance::io::exec::utils::InstrumentedRecordBatchStreamAdapter;
|
||||
use lance_table::format::Fragment;
|
||||
@@ -74,7 +74,9 @@ fn merge_transactions(mut transactions: Vec<Transaction>) -> Option<Transaction>
|
||||
///
|
||||
/// This plan executes inserts by:
|
||||
/// 1. Each partition writes data independently using InsertBuilder::execute_uncommitted_stream
|
||||
/// 2. The last partition to complete commits all transactions atomically
|
||||
/// 2. The last partition to complete merges transactions, optionally attaches one
|
||||
/// precomputed generated-column metadata patch when the merged write has rows,
|
||||
/// then commits once
|
||||
/// 3. Returns the count of inserted rows per partition
|
||||
#[derive(Debug)]
|
||||
pub struct InsertExec {
|
||||
@@ -83,6 +85,10 @@ pub struct InsertExec {
|
||||
input: Arc<dyn ExecutionPlan>,
|
||||
write_params: WriteParams,
|
||||
tracker: Option<Arc<WriteProgressTracker>>,
|
||||
/// Optional whole-transaction field-metadata patch for generated-column
|
||||
/// invalidation. Attached once after partition merge, and only when the
|
||||
/// merged operation contains at least one written row.
|
||||
schema_metadata_updates: Option<SchemaMetadataUpdates>,
|
||||
properties: Arc<PlanProperties>,
|
||||
partial_transactions: Arc<Mutex<Vec<Transaction>>>,
|
||||
metrics: ExecutionPlanMetricsSet,
|
||||
@@ -95,7 +101,7 @@ impl InsertExec {
|
||||
input: Arc<dyn ExecutionPlan>,
|
||||
write_params: WriteParams,
|
||||
) -> Self {
|
||||
Self::new_with_tracker(ds_wrapper, dataset, input, write_params, None)
|
||||
Self::new_with_tracker(ds_wrapper, dataset, input, write_params, None, None)
|
||||
}
|
||||
|
||||
pub(crate) fn new_with_tracker(
|
||||
@@ -104,6 +110,7 @@ impl InsertExec {
|
||||
input: Arc<dyn ExecutionPlan>,
|
||||
write_params: WriteParams,
|
||||
tracker: Option<Arc<WriteProgressTracker>>,
|
||||
schema_metadata_updates: Option<SchemaMetadataUpdates>,
|
||||
) -> Self {
|
||||
let schema = COUNT_SCHEMA.clone();
|
||||
let num_partitions = input.output_partitioning().partition_count();
|
||||
@@ -120,6 +127,7 @@ impl InsertExec {
|
||||
input,
|
||||
write_params,
|
||||
tracker,
|
||||
schema_metadata_updates,
|
||||
properties: Arc::new(properties),
|
||||
partial_transactions: Arc::new(Mutex::new(Vec::with_capacity(num_partitions))),
|
||||
metrics: ExecutionPlanMetricsSet::new(),
|
||||
@@ -176,6 +184,7 @@ impl ExecutionPlan for InsertExec {
|
||||
children[0].clone(),
|
||||
self.write_params.clone(),
|
||||
self.tracker.clone(),
|
||||
self.schema_metadata_updates.clone(),
|
||||
)))
|
||||
}
|
||||
|
||||
@@ -191,6 +200,7 @@ impl ExecutionPlan for InsertExec {
|
||||
let total_partitions = self.input.output_partitioning().partition_count();
|
||||
let ds_wrapper = self.ds_wrapper.clone();
|
||||
let tracker = self.tracker.clone();
|
||||
let schema_metadata_updates = self.schema_metadata_updates.clone();
|
||||
|
||||
let output_bytes = MetricBuilder::new(&self.metrics).output_bytes(partition);
|
||||
let input_schema = input_stream.schema();
|
||||
@@ -220,6 +230,8 @@ impl ExecutionPlan for InsertExec {
|
||||
}));
|
||||
}
|
||||
|
||||
// Each partition stages an uncommitted data-only transaction.
|
||||
// Metadata invalidation is attached once on the merged commit.
|
||||
let transaction = InsertBuilder::new(dataset.clone())
|
||||
.with_params(&write_params)
|
||||
.execute_uncommitted_stream(input_stream)
|
||||
@@ -241,8 +253,15 @@ impl ExecutionPlan for InsertExec {
|
||||
};
|
||||
|
||||
if let Some(transactions) = to_commit
|
||||
&& let Some(merged_txn) = merge_transactions(transactions)
|
||||
&& let Some(mut merged_txn) = merge_transactions(transactions)
|
||||
{
|
||||
// Attach the precomputed patch only for non-empty writes, and
|
||||
// only once for the whole multi-partition transaction.
|
||||
if count_rows_from_operation(&merged_txn.operation) > 0
|
||||
&& let Some(updates) = schema_metadata_updates
|
||||
{
|
||||
merged_txn = merged_txn.with_schema_metadata_updates(updates)?;
|
||||
}
|
||||
let new_dataset = CommitBuilder::new(dataset.clone())
|
||||
.execute(merged_txn)
|
||||
.await?;
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use futures::FutureExt;
|
||||
use lance::dataset::DeleteBuilder;
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use lance::dataset::DeleteBuilder;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::{NativeTable, Predicate};
|
||||
@@ -29,34 +29,40 @@ pub(crate) async fn execute_delete(
|
||||
predicate: Predicate<'_>,
|
||||
) -> Result<DeleteResult> {
|
||||
table.dataset.ensure_mutable()?;
|
||||
match predicate {
|
||||
Predicate::String(s) => {
|
||||
let mut dataset = (*table.dataset.get().await?).clone();
|
||||
let delete_result = dataset.delete(s).boxed().await?;
|
||||
let num_deleted_rows = delete_result.num_deleted_rows;
|
||||
let version = dataset.version().version;
|
||||
table.dataset.update(dataset);
|
||||
Ok(DeleteResult {
|
||||
num_deleted_rows,
|
||||
version,
|
||||
})
|
||||
}
|
||||
Predicate::Expr(expr) => {
|
||||
let dataset = table.dataset.get().await?;
|
||||
let delete_result = DeleteBuilder::from_expr(Arc::clone(&dataset), expr.clone())
|
||||
.execute()
|
||||
.await?;
|
||||
let num_deleted_rows = delete_result.num_deleted_rows;
|
||||
let version = delete_result.new_dataset.version().version;
|
||||
table.dataset.update(
|
||||
Arc::try_unwrap(delete_result.new_dataset).unwrap_or_else(|arc| (*arc).clone()),
|
||||
);
|
||||
Ok(DeleteResult {
|
||||
num_deleted_rows,
|
||||
version,
|
||||
})
|
||||
}
|
||||
|
||||
// One exact dataset supplies binding-snapshot planning, the DeleteBuilder,
|
||||
// and its transaction basis. Do not call table schema()/version() or another
|
||||
// get(). Conflicts are not caught/replanned here.
|
||||
let dataset = table.dataset.get().await?;
|
||||
|
||||
// String preserves the legacy Dataset::delete zero-retry baseline; Expr
|
||||
// retains DeleteBuilder defaults until a generated patch is attached.
|
||||
let mut builder = match predicate {
|
||||
Predicate::String(s) => DeleteBuilder::new(Arc::clone(&dataset), s).conflict_retries(0),
|
||||
Predicate::Expr(expr) => DeleteBuilder::from_expr(Arc::clone(&dataset), expr.clone()),
|
||||
};
|
||||
|
||||
if let Some(schema_metadata_updates) =
|
||||
super::generated_column_invalidation::plan_native_delete_generated_column_invalidation(
|
||||
dataset.as_ref(),
|
||||
)?
|
||||
{
|
||||
// Exact-basis fence: never retry an old generated patch on latest.
|
||||
builder = builder
|
||||
.with_schema_metadata_updates(schema_metadata_updates)?
|
||||
.conflict_retries(0);
|
||||
}
|
||||
|
||||
let delete_result = builder.execute().await?;
|
||||
let num_deleted_rows = delete_result.num_deleted_rows;
|
||||
let version = delete_result.new_dataset.version().version;
|
||||
table
|
||||
.dataset
|
||||
.update(Arc::try_unwrap(delete_result.new_dataset).unwrap_or_else(|arc| (*arc).clone()));
|
||||
Ok(DeleteResult {
|
||||
num_deleted_rows,
|
||||
version,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,242 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Crate-private Native wiring for generated-column invalidation (B4b / B4c / B4d / B4e).
|
||||
//!
|
||||
//! Converts the B4a pure planner into one Lance field-metadata patch for Native
|
||||
//! append, update, and delete commits. Planning is strict-decode/validate;
|
||||
//! overwrite of a table with any generated-column definition, direct writes of
|
||||
//! generated outputs via Update, and Native merge-insert (standard and LSM)
|
||||
//! fail closed as [`Error::NotSupported`].
|
||||
|
||||
use std::collections::{BTreeSet, HashMap};
|
||||
|
||||
use lance::Dataset;
|
||||
use lance::dataset::transaction::{SchemaMetadataUpdates, UpdateMap, UpdateMapEntry};
|
||||
|
||||
use crate::Result;
|
||||
use crate::error::Error;
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
use crate::function::plan_generated_column_invalidation::{
|
||||
GeneratedColumnMutationImpact, PlannedGeneratedColumnMetadataUpdate,
|
||||
plan_generated_column_invalidation,
|
||||
};
|
||||
|
||||
use super::generated_column_binding_snapshot_from_dataset;
|
||||
|
||||
/// Plan Native append invalidation against one exact dataset snapshot.
|
||||
///
|
||||
/// Strict-decodes and validates every present generated-column metadata value
|
||||
/// through the B4a planner. When `is_overwrite` is true and any generated column
|
||||
/// is present, returns [`Error::NotSupported`] before mutation. Otherwise returns
|
||||
/// `Some(patch)` when at least one generated column would be invalidated, or
|
||||
/// `None` when the table has no generated columns.
|
||||
pub(super) fn plan_native_append_generated_column_invalidation(
|
||||
dataset: &Dataset,
|
||||
is_overwrite: bool,
|
||||
) -> Result<Option<SchemaMetadataUpdates>> {
|
||||
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snapshot,
|
||||
&GeneratedColumnMutationImpact::RowSetChanged,
|
||||
)?;
|
||||
if plan.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
if is_overwrite {
|
||||
return Err(Error::NotSupported {
|
||||
message: "Overwrite is not supported on tables with generated columns".to_string(),
|
||||
});
|
||||
}
|
||||
Ok(Some(planned_invalidation_to_schema_metadata_updates(plan)))
|
||||
}
|
||||
|
||||
/// Plan Native update invalidation against one exact dataset snapshot.
|
||||
///
|
||||
/// Strict-decodes and validates every present generated-column definition before
|
||||
/// impact calculation, even when `updated_field_ids` does not affect any
|
||||
/// generated output. After the global planner succeeds, a target whose snapshot
|
||||
/// entry contains generated metadata is rejected as a direct generated-output
|
||||
/// write ([`Error::NotSupported`]) before any Update file write. Returns
|
||||
/// `Some(patch)` when the impact closure is non-empty, otherwise `None`.
|
||||
pub(super) fn plan_native_update_generated_column_invalidation(
|
||||
dataset: &Dataset,
|
||||
updated_field_ids: BTreeSet<i32>,
|
||||
) -> Result<Option<SchemaMetadataUpdates>> {
|
||||
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snapshot,
|
||||
&GeneratedColumnMutationImpact::UpdatedFields(updated_field_ids.clone()),
|
||||
)?;
|
||||
|
||||
for field_id in &updated_field_ids {
|
||||
let Some(entry) = snapshot
|
||||
.entries()
|
||||
.iter()
|
||||
.find(|entry| entry.field_id() == *field_id)
|
||||
else {
|
||||
return Err(Error::InvalidInput {
|
||||
message: format!("updated field id {field_id} was not found in the table schema"),
|
||||
});
|
||||
};
|
||||
if entry
|
||||
.field()
|
||||
.metadata()
|
||||
.contains_key(GENERATED_COLUMN_METADATA_KEY)
|
||||
{
|
||||
return Err(Error::NotSupported {
|
||||
message: "Updating generated columns is not supported".to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if plan.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(Some(planned_invalidation_to_schema_metadata_updates(plan)))
|
||||
}
|
||||
|
||||
/// Plan Native delete invalidation against one exact dataset snapshot.
|
||||
///
|
||||
/// Strict-decodes and validates every present generated-column metadata value
|
||||
/// through the B4a `RowSetChanged` planner before any Delete scanner/file IO.
|
||||
/// Returns `Some(patch)` when at least one generated column would be invalidated,
|
||||
/// or `None` when the table has no generated columns. Actual zero-row Delete
|
||||
/// suppression is owned by Lance A4d, not this planner.
|
||||
pub(super) fn plan_native_delete_generated_column_invalidation(
|
||||
dataset: &Dataset,
|
||||
) -> Result<Option<SchemaMetadataUpdates>> {
|
||||
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snapshot,
|
||||
&GeneratedColumnMutationImpact::RowSetChanged,
|
||||
)?;
|
||||
if plan.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(Some(planned_invalidation_to_schema_metadata_updates(plan)))
|
||||
}
|
||||
|
||||
/// Fail closed before Native `merge_insert` when any generated column is present.
|
||||
///
|
||||
/// Strict-decodes and validates every present generated-column metadata value
|
||||
/// through the B4a `RowSetChanged` planner against one exact dataset snapshot.
|
||||
/// Malformed metadata returns the existing [`Error::InvalidInput`] validation
|
||||
/// category. When at least one valid generated column is present, returns
|
||||
/// [`Error::NotSupported`] before LSM dispatch or source iteration. Ordinary
|
||||
/// tables (no generated metadata) return `Ok(())`.
|
||||
pub(super) fn reject_native_merge_insert_if_generated_columns_present(
|
||||
dataset: &Dataset,
|
||||
) -> Result<()> {
|
||||
let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?;
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snapshot,
|
||||
&GeneratedColumnMutationImpact::RowSetChanged,
|
||||
)?;
|
||||
if plan.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
Err(Error::NotSupported {
|
||||
message: "Merge insert is not supported on tables with generated columns".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Convert planner replacements into one non-empty Lance field-metadata patch.
|
||||
///
|
||||
/// Each entry is keyed by stable output field ID, uses `replace: false`, and
|
||||
/// replaces only [`GENERATED_COLUMN_METADATA_KEY`].
|
||||
fn planned_invalidation_to_schema_metadata_updates(
|
||||
plan: Vec<PlannedGeneratedColumnMetadataUpdate>,
|
||||
) -> SchemaMetadataUpdates {
|
||||
SchemaMetadataUpdates {
|
||||
schema_metadata_updates: None,
|
||||
field_metadata_updates: plan
|
||||
.into_iter()
|
||||
.map(|update| {
|
||||
(
|
||||
update.output_field_id(),
|
||||
UpdateMap {
|
||||
update_entries: vec![UpdateMapEntry {
|
||||
key: GENERATED_COLUMN_METADATA_KEY.to_string(),
|
||||
value: Some(update.metadata_json().to_string()),
|
||||
}],
|
||||
replace: false,
|
||||
},
|
||||
)
|
||||
})
|
||||
.collect::<HashMap<_, _>>(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// Construct a planned update through the public accessors by planning a
|
||||
/// minimal in-memory snapshot, then assert the Lance patch shape.
|
||||
#[test]
|
||||
fn planned_replacements_become_non_replace_field_patch() {
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput,
|
||||
FunctionParameter, FunctionSignature, GeneratedColumnBindingSnapshot,
|
||||
GeneratedColumnDefinition,
|
||||
};
|
||||
use arrow_array::{ArrayRef, StringArray};
|
||||
use arrow_schema::{DataType, Field};
|
||||
use std::sync::Arc;
|
||||
|
||||
let field_id = 11;
|
||||
let function = Function::new(
|
||||
FunctionId::try_new("fn.exact.b4b.helper.patch").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("x")])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
let definition = GeneratedColumnDefinition::try_new(field_id, call, 3, 3).unwrap();
|
||||
let json = definition.to_metadata_json().unwrap();
|
||||
let snap = GeneratedColumnBindingSnapshot::try_new(
|
||||
1,
|
||||
vec![Arc::new(
|
||||
Field::new("gen_out", DataType::Int32, true)
|
||||
.with_metadata([(GENERATED_COLUMN_METADATA_KEY.to_string(), json)].into()),
|
||||
)],
|
||||
vec![field_id],
|
||||
)
|
||||
.unwrap();
|
||||
let plan = plan_generated_column_invalidation(
|
||||
&snap,
|
||||
&GeneratedColumnMutationImpact::RowSetChanged,
|
||||
)
|
||||
.unwrap();
|
||||
let patch = planned_invalidation_to_schema_metadata_updates(plan);
|
||||
assert!(!patch.is_empty());
|
||||
assert!(patch.schema_metadata_updates.is_none());
|
||||
let map = patch
|
||||
.field_metadata_updates
|
||||
.get(&field_id)
|
||||
.expect("stable field id must be present");
|
||||
assert!(!map.replace);
|
||||
assert_eq!(map.update_entries.len(), 1);
|
||||
assert_eq!(map.update_entries[0].key, GENERATED_COLUMN_METADATA_KEY);
|
||||
let decoded = GeneratedColumnDefinition::from_metadata_json(
|
||||
map.update_entries[0].value.as_deref().unwrap(),
|
||||
field_id,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(decoded.dependency_epoch(), 4);
|
||||
assert_eq!(decoded.materialized_epoch(), 3);
|
||||
}
|
||||
}
|
||||
@@ -233,10 +233,22 @@ pub(crate) async fn execute_merge_insert(
|
||||
params: MergeInsertBuilder,
|
||||
new_data: Box<dyn RecordBatchReader + Send>,
|
||||
) -> Result<MergeResult> {
|
||||
match lsm::lsm_dispatch_decision(table, ¶ms).await? {
|
||||
// One exact dataset supplies the generated-column fail-closed guard and
|
||||
// downstream standard/LSM routing/execution. Do not refetch after the guard.
|
||||
let dataset = table.dataset.get().await?;
|
||||
super::generated_column_invalidation::reject_native_merge_insert_if_generated_columns_present(
|
||||
dataset.as_ref(),
|
||||
)?;
|
||||
|
||||
match lsm::lsm_dispatch_decision(¶ms, dataset.as_ref()).await? {
|
||||
lsm::LsmDispatch::Lsm(plan) => {
|
||||
let future =
|
||||
lsm::execute_lsm_merge_insert(table, plan, params.validate_single_shard, new_data);
|
||||
let future = lsm::execute_lsm_merge_insert(
|
||||
table,
|
||||
plan,
|
||||
params.validate_single_shard,
|
||||
new_data,
|
||||
dataset,
|
||||
);
|
||||
return match params.timeout {
|
||||
Some(timeout) => match tokio::time::timeout(timeout, future).await {
|
||||
Ok(result) => result,
|
||||
@@ -250,7 +262,6 @@ pub(crate) async fn execute_merge_insert(
|
||||
lsm::LsmDispatch::Standard => {}
|
||||
}
|
||||
|
||||
let dataset = table.dataset.get().await?;
|
||||
let mut builder = LanceMergeInsertBuilder::try_new(dataset.clone(), params.on)?;
|
||||
match (
|
||||
params.when_matched_update_all,
|
||||
|
||||
@@ -531,18 +531,18 @@ pub(crate) enum LsmDispatch {
|
||||
}
|
||||
|
||||
/// Decide whether a `merge_insert` should be routed through the MemWAL write
|
||||
/// path, validating the builder against the installed spec.
|
||||
/// path, validating the builder against the installed spec on the exact
|
||||
/// caller-supplied dataset snapshot.
|
||||
#[allow(clippy::redundant_pub_crate)]
|
||||
pub(crate) async fn lsm_dispatch_decision(
|
||||
table: &NativeTable,
|
||||
params: &MergeInsertBuilder,
|
||||
dataset: &Dataset,
|
||||
) -> Result<LsmDispatch> {
|
||||
// Explicit opt-out: use the standard path regardless of any installed spec.
|
||||
if params.use_lsm == Some(false) {
|
||||
return Ok(LsmDispatch::Standard);
|
||||
}
|
||||
|
||||
let dataset = table.dataset.get().await?;
|
||||
let Some(details) = dataset.mem_wal_index_details().await? else {
|
||||
// No write spec installed. `use_lsm(true)` demanded MemWAL routing, so
|
||||
// that is an error; otherwise fall back to the standard path.
|
||||
@@ -646,14 +646,17 @@ fn resolve_lsm_mode(details: &MemWalIndexDetails) -> Result<LsmMode> {
|
||||
/// a validation failure (e.g. input spanning shards) never leaves a partial
|
||||
/// write behind. When `validate_single_shard` is set, every row is checked to
|
||||
/// route to one shard; when disabled, only the first row of the whole input is.
|
||||
///
|
||||
/// `dataset` must be the same exact snapshot used for the generated-column
|
||||
/// guard and [`lsm_dispatch_decision`].
|
||||
#[allow(clippy::redundant_pub_crate)]
|
||||
pub(crate) async fn execute_lsm_merge_insert(
|
||||
table: &NativeTable,
|
||||
plan: LsmPlan,
|
||||
validate_single_shard: bool,
|
||||
new_data: Box<dyn RecordBatchReader + Send>,
|
||||
dataset: Arc<Dataset>,
|
||||
) -> Result<MergeResult> {
|
||||
let dataset = table.dataset.get().await?;
|
||||
let target_schema: SchemaRef = Arc::new(ArrowSchema::from(dataset.schema()));
|
||||
|
||||
// Collect, align and shard-validate the whole input before writing
|
||||
|
||||
@@ -0,0 +1,531 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! RED runtime contract tests for Native merge-insert fail-closed guard (B4e).
|
||||
//!
|
||||
//! Tables with generated-column definitions cannot carry dependency-epoch
|
||||
//! metadata updates through Native merge-insert in this slice. Both the
|
||||
//! standard and MemWAL/LSM routes must reject before consuming source input or
|
||||
//! mutating the table. Ordinary tables keep existing merge-insert semantics.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use arrow_array::{Int32Array, RecordBatch, RecordBatchReader, StringArray};
|
||||
use arrow_schema::{ArrowError, DataType, Field, Schema, SchemaRef};
|
||||
use futures::TryStreamExt;
|
||||
use tempfile::TempDir;
|
||||
|
||||
use crate::connection::ConnectBuilder;
|
||||
use crate::error::Error;
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
|
||||
GeneratedColumnStatus,
|
||||
};
|
||||
use crate::query::{ExecutableQuery, QueryBase, Select};
|
||||
use crate::table::Table;
|
||||
|
||||
const ID: &str = "id";
|
||||
const ORDINARY: &str = "ordinary";
|
||||
const GEN_OUT: &str = "gen_out";
|
||||
const INITIAL_DEPENDENCY_EPOCH: u64 = 3;
|
||||
const INITIAL_MATERIALIZED_EPOCH: u64 = 3;
|
||||
const FN_ID: &str = "fn.exact.b4e.merge.literal";
|
||||
const MALFORMED_MARKER: &str = "SENSITIVE_B4E_MERGE_METADATA_MARKER_7c91_e2ab";
|
||||
|
||||
struct Fixture {
|
||||
_tmp: TempDir,
|
||||
table: Table,
|
||||
table_name: String,
|
||||
uri: String,
|
||||
}
|
||||
|
||||
/// RecordBatchReader that counts how many times [`Self::next`] is called.
|
||||
struct ObservableReader {
|
||||
inner: Box<dyn RecordBatchReader + Send>,
|
||||
next_calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
impl ObservableReader {
|
||||
fn wrap(
|
||||
inner: Box<dyn RecordBatchReader + Send>,
|
||||
next_calls: Arc<AtomicUsize>,
|
||||
) -> Box<dyn RecordBatchReader + Send> {
|
||||
Box::new(Self { inner, next_calls })
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for ObservableReader {
|
||||
type Item = Result<RecordBatch, ArrowError>;
|
||||
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.next_calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.inner.next()
|
||||
}
|
||||
}
|
||||
|
||||
impl RecordBatchReader for ObservableReader {
|
||||
fn schema(&self) -> SchemaRef {
|
||||
self.inner.schema()
|
||||
}
|
||||
}
|
||||
|
||||
fn literal_only_function() -> Function {
|
||||
Function::new(
|
||||
FunctionId::try_new(FN_ID).unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
)
|
||||
}
|
||||
|
||||
fn literal_only_definition(output_field_id: i32) -> GeneratedColumnDefinition {
|
||||
let function = literal_only_function();
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("literal-only")])) as arrow_array::ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
GeneratedColumnDefinition::try_new(
|
||||
output_field_id,
|
||||
call,
|
||||
INITIAL_DEPENDENCY_EPOCH,
|
||||
INITIAL_MATERIALIZED_EPOCH,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn seed_batch() -> RecordBatch {
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(ID, DataType::Int32, false),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
]));
|
||||
RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1, 2])),
|
||||
Arc::new(StringArray::from(vec![Some("a"), Some("b")])),
|
||||
Arc::new(Int32Array::from(vec![10, 20])),
|
||||
],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn source_batch(ids: &[i32], ordinary: &[&str], gen_values: &[i32]) -> RecordBatch {
|
||||
assert_eq!(ids.len(), ordinary.len());
|
||||
assert_eq!(ids.len(), gen_values.len());
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(ID, DataType::Int32, false),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
]));
|
||||
RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(ids.to_vec())),
|
||||
Arc::new(StringArray::from(
|
||||
ordinary
|
||||
.iter()
|
||||
.map(|value| Some(*value))
|
||||
.collect::<Vec<_>>(),
|
||||
)),
|
||||
Arc::new(Int32Array::from(gen_values.to_vec())),
|
||||
],
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn boxed_reader(batch: RecordBatch) -> Box<dyn RecordBatchReader + Send> {
|
||||
let schema = batch.schema();
|
||||
Box::new(arrow_array::RecordBatchIterator::new(
|
||||
vec![Ok(batch)].into_iter(),
|
||||
schema,
|
||||
))
|
||||
}
|
||||
|
||||
fn empty_reader() -> Box<dyn RecordBatchReader + Send> {
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(ID, DataType::Int32, false),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
]));
|
||||
Box::new(arrow_array::RecordBatchIterator::new(
|
||||
std::iter::empty::<Result<RecordBatch, ArrowError>>(),
|
||||
schema,
|
||||
))
|
||||
}
|
||||
|
||||
async fn create_ordinary_table(name: &str) -> Fixture {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().to_str().unwrap().to_string();
|
||||
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
|
||||
let table = conn
|
||||
.create_table(name, seed_batch())
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
Fixture {
|
||||
_tmp: tmp,
|
||||
table,
|
||||
table_name: name.to_string(),
|
||||
uri,
|
||||
}
|
||||
}
|
||||
|
||||
async fn create_table_with_complete_literal_generated(name: &str) -> Fixture {
|
||||
let fixture = create_ordinary_table(name).await;
|
||||
let snapshot = fixture
|
||||
.table
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.unwrap();
|
||||
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
|
||||
let definition = literal_only_definition(field_id);
|
||||
let json = definition.to_metadata_json().unwrap();
|
||||
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
|
||||
fixture
|
||||
.table
|
||||
.as_native()
|
||||
.expect("generated-column fixture planting requires a Native table"),
|
||||
GEN_OUT,
|
||||
json,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
fixture
|
||||
}
|
||||
|
||||
async fn read_generated_definition(table: &Table) -> GeneratedColumnDefinition {
|
||||
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
|
||||
snapshot
|
||||
.field(GEN_OUT)
|
||||
.expect(GEN_OUT)
|
||||
.generated_column_definition()
|
||||
.expect("generated metadata must decode")
|
||||
.expect("generated metadata must be present")
|
||||
}
|
||||
|
||||
async fn read_raw_generated_metadata(table: &Table) -> String {
|
||||
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
|
||||
snapshot
|
||||
.field(GEN_OUT)
|
||||
.expect(GEN_OUT)
|
||||
.field()
|
||||
.metadata()
|
||||
.get(GENERATED_COLUMN_METADATA_KEY)
|
||||
.expect("generated metadata key must be present")
|
||||
.clone()
|
||||
}
|
||||
|
||||
async fn ordinary_rows(table: &Table) -> Vec<(i32, String)> {
|
||||
let batches = table
|
||||
.query()
|
||||
.select(Select::columns(&[ID, ORDINARY]))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
let mut rows = Vec::new();
|
||||
for batch in batches {
|
||||
let ids = batch
|
||||
.column(0)
|
||||
.as_any()
|
||||
.downcast_ref::<Int32Array>()
|
||||
.unwrap();
|
||||
let ordinary = batch
|
||||
.column(1)
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.unwrap();
|
||||
for index in 0..batch.num_rows() {
|
||||
rows.push((ids.value(index), ordinary.value(index).to_string()));
|
||||
}
|
||||
}
|
||||
rows.sort_by_key(|(id, _)| *id);
|
||||
rows
|
||||
}
|
||||
|
||||
fn assert_not_supported(err: &Error, label: &str) {
|
||||
assert!(
|
||||
matches!(err, Error::NotSupported { .. }),
|
||||
"{label}: expected NotSupported, got {err:?}"
|
||||
);
|
||||
}
|
||||
|
||||
fn assert_invalid_input_redacted(err: &Error, planted_raw: &str, label: &str) {
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"{label}: expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let rendered = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!rendered.contains(MALFORMED_MARKER),
|
||||
"{label}: diagnostic echoed unique metadata marker: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(FN_ID),
|
||||
"{label}: diagnostic echoed Function ID: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: diagnostic echoed metadata wire key: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(planted_raw),
|
||||
"{label}: diagnostic echoed raw metadata JSON: {rendered}"
|
||||
);
|
||||
}
|
||||
|
||||
fn configure_standard_merge(builder: &mut crate::table::merge::MergeInsertBuilder) {
|
||||
builder
|
||||
.when_matched_update_all(None)
|
||||
.when_not_matched_insert_all()
|
||||
.when_not_matched_by_source_delete(None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standard_merge_insert_rejects_when_generated_column_present_before_input_consumption() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4e_standard_reject").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_rows(&fixture.table).await;
|
||||
let definition_before = read_generated_definition(&fixture.table).await;
|
||||
let raw_before = read_raw_generated_metadata(&fixture.table).await;
|
||||
assert_eq!(
|
||||
definition_before.function_call().function_id().as_str(),
|
||||
FN_ID
|
||||
);
|
||||
|
||||
let next_calls = Arc::new(AtomicUsize::new(0));
|
||||
let reader = ObservableReader::wrap(
|
||||
boxed_reader(source_batch(&[1, 3], &["updated", "inserted"], &[11, 30])),
|
||||
next_calls.clone(),
|
||||
);
|
||||
|
||||
let mut builder = fixture.table.merge_insert(&[ID]);
|
||||
configure_standard_merge(&mut builder);
|
||||
let err = builder
|
||||
.execute(reader)
|
||||
.await
|
||||
.expect_err("generated-column table must reject standard merge_insert");
|
||||
assert_not_supported(&err, "standard merge_insert generated reject");
|
||||
assert_eq!(
|
||||
next_calls.load(Ordering::SeqCst),
|
||||
0,
|
||||
"rejection must occur before consuming the RecordBatchReader"
|
||||
);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
read_generated_definition(&fixture.table).await,
|
||||
definition_before
|
||||
);
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fixture.table).await,
|
||||
raw_before
|
||||
);
|
||||
assert_eq!(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_status(GEN_OUT)
|
||||
.await
|
||||
.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn empty_standard_merge_insert_rejects_when_generated_column_present() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4e_empty_reject").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_rows(&fixture.table).await;
|
||||
let raw_before = read_raw_generated_metadata(&fixture.table).await;
|
||||
|
||||
let next_calls = Arc::new(AtomicUsize::new(0));
|
||||
let reader = ObservableReader::wrap(empty_reader(), next_calls.clone());
|
||||
|
||||
let mut builder = fixture.table.merge_insert(&[ID]);
|
||||
configure_standard_merge(&mut builder);
|
||||
let err = builder
|
||||
.execute(reader)
|
||||
.await
|
||||
.expect_err("empty merge_insert must still reject on generated-column tables");
|
||||
assert_not_supported(&err, "empty standard merge_insert generated reject");
|
||||
assert_eq!(next_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fixture.table).await,
|
||||
raw_before
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn forced_lsm_without_spec_rejects_generated_before_missing_spec_and_input() {
|
||||
let fixture = create_table_with_complete_literal_generated("b4e_lsm_force_reject").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_rows(&fixture.table).await;
|
||||
let raw_before = read_raw_generated_metadata(&fixture.table).await;
|
||||
|
||||
let next_calls = Arc::new(AtomicUsize::new(0));
|
||||
let reader = ObservableReader::wrap(
|
||||
boxed_reader(source_batch(&[1], &["must-not-land"], &[11])),
|
||||
next_calls.clone(),
|
||||
);
|
||||
|
||||
let mut builder = fixture.table.merge_insert(&[ID]);
|
||||
builder
|
||||
.when_matched_update_all(None)
|
||||
.when_not_matched_insert_all()
|
||||
.use_lsm(true);
|
||||
let err = builder
|
||||
.execute(reader)
|
||||
.await
|
||||
.expect_err("generated-column guard must run before LSM missing-spec validation");
|
||||
assert_not_supported(&err, "forced LSM generated reject");
|
||||
assert_eq!(next_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fixture.table).await,
|
||||
raw_before
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn malformed_generated_metadata_rejects_merge_insert_before_mutation_and_redacts() {
|
||||
let fixture = create_ordinary_table("b4e_malformed_preflight").await;
|
||||
let snapshot = fixture
|
||||
.table
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.unwrap();
|
||||
let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id();
|
||||
let planted_raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{{"function_id":"{FN_ID}","marker":"{MALFORMED_MARKER}"}},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(planted_raw.contains(MALFORMED_MARKER));
|
||||
assert!(planted_raw.contains(FN_ID));
|
||||
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
|
||||
fixture
|
||||
.table
|
||||
.as_native()
|
||||
.expect("generated-column fixture planting requires a Native table"),
|
||||
GEN_OUT,
|
||||
planted_raw.clone(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fixture.table).await,
|
||||
planted_raw,
|
||||
"planted malformed raw metadata must round-trip byte-for-byte"
|
||||
);
|
||||
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_rows(&fixture.table).await;
|
||||
let next_calls = Arc::new(AtomicUsize::new(0));
|
||||
let reader = ObservableReader::wrap(
|
||||
boxed_reader(source_batch(&[1], &["must-not-land"], &[99])),
|
||||
next_calls.clone(),
|
||||
);
|
||||
|
||||
let mut builder = fixture.table.merge_insert(&[ID]);
|
||||
configure_standard_merge(&mut builder);
|
||||
let err = builder
|
||||
.execute(reader)
|
||||
.await
|
||||
.expect_err("malformed generated metadata must fail closed before merge_insert");
|
||||
assert_invalid_input_redacted(&err, &planted_raw, "malformed merge_insert preflight");
|
||||
assert_eq!(next_calls.load(Ordering::SeqCst), 0);
|
||||
|
||||
let fresh = ConnectBuilder::new(&fixture.uri)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.open_table(&fixture.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(fresh.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_rows(&fresh).await, rows_before);
|
||||
assert_eq!(read_raw_generated_metadata(&fresh).await, planted_raw);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ordinary_table_standard_merge_insert_preserves_result_semantics() {
|
||||
let fixture = create_ordinary_table("b4e_ordinary_standard").await;
|
||||
let mut builder = fixture.table.merge_insert(&[ID]);
|
||||
configure_standard_merge(&mut builder);
|
||||
let result = builder
|
||||
.execute(boxed_reader(source_batch(
|
||||
&[1, 3],
|
||||
&["updated", "inserted"],
|
||||
&[11, 30],
|
||||
)))
|
||||
.await
|
||||
.expect("ordinary-table standard merge_insert must succeed");
|
||||
|
||||
assert_eq!(result.num_inserted_rows, 1);
|
||||
assert_eq!(result.num_updated_rows, 1);
|
||||
assert_eq!(result.num_deleted_rows, 1);
|
||||
assert_eq!(result.num_attempts, 1);
|
||||
assert_eq!(result.num_rows, 2);
|
||||
assert!(result.version > 0);
|
||||
|
||||
assert_eq!(
|
||||
ordinary_rows(&fixture.table).await,
|
||||
vec![(1, "updated".to_string()), (3, "inserted".to_string()),]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ordinary_table_forced_lsm_without_spec_keeps_missing_spec_error() {
|
||||
let fixture = create_ordinary_table("b4e_ordinary_lsm_missing_spec").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_rows(&fixture.table).await;
|
||||
|
||||
let mut builder = fixture.table.merge_insert(&[ID]);
|
||||
builder
|
||||
.when_matched_update_all(None)
|
||||
.when_not_matched_insert_all()
|
||||
.use_lsm(true);
|
||||
let err = builder
|
||||
.execute(boxed_reader(source_batch(&[1], &["x"], &[1])))
|
||||
.await
|
||||
.expect_err("ordinary table without MemWAL spec must keep missing-spec InvalidInput");
|
||||
match err {
|
||||
Error::InvalidInput { message } => {
|
||||
assert!(
|
||||
message.contains("no MemWAL write spec"),
|
||||
"expected missing-spec message, got {message}"
|
||||
);
|
||||
}
|
||||
other => panic!("expected InvalidInput missing-spec, got {other:?}"),
|
||||
}
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_rows(&fixture.table).await, rows_before);
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
+377
-19
@@ -5,12 +5,13 @@ use std::sync::Arc;
|
||||
|
||||
mod lsm;
|
||||
|
||||
use super::NativeTable;
|
||||
use super::{NativeTable, generated_column_binding_snapshot_from_dataset};
|
||||
use crate::connection::NamespaceClientPushdownOperation;
|
||||
use crate::error::{Error, Result};
|
||||
use crate::expr::expr_to_sql_string;
|
||||
use crate::query::{
|
||||
DEFAULT_TOP_K, QueryExecutionOptions, QueryFilter, QueryRequest, Select, VectorQueryRequest,
|
||||
validate_generated_column_query,
|
||||
};
|
||||
use crate::utils::{MaxBatchLengthStream, TimeoutStream, default_vector_column};
|
||||
use arrow::array::{AsArray, FixedSizeListBuilder, Float32Builder};
|
||||
@@ -22,6 +23,7 @@ use datafusion_physical_plan::projection::ProjectionExec;
|
||||
use datafusion_physical_plan::repartition::RepartitionExec;
|
||||
use datafusion_physical_plan::union::UnionExec;
|
||||
use futures::future::try_join_all;
|
||||
use lance::Dataset;
|
||||
use lance::dataset::mem_wal::DatasetMemWalExt;
|
||||
use lance::dataset::scanner::DatasetRecordBatchStream;
|
||||
use lance::dataset::scanner::Scanner;
|
||||
@@ -47,7 +49,7 @@ impl AnyQuery {
|
||||
}
|
||||
}
|
||||
|
||||
//Decide between namespace or local
|
||||
// Decide between namespace or local.
|
||||
pub async fn execute_query(
|
||||
table: &NativeTable,
|
||||
query: &AnyQuery,
|
||||
@@ -55,16 +57,25 @@ pub async fn execute_query(
|
||||
) -> Result<DatasetRecordBatchStream> {
|
||||
// QueryTable pushdown runs the query server-side, but only on the main
|
||||
// branch: the namespace request carries no branch yet, so a branch handle
|
||||
// must fall through to local execution.
|
||||
if can_execute_namespace_query(table, query).await?
|
||||
&& let Some(ref namespace_client) = table.namespace_client
|
||||
{
|
||||
return execute_namespace_query(table, namespace_client.clone(), query, options).await;
|
||||
// must fall through to local execution. Successful pushdown owns one
|
||||
// Dataset Arc for MemWAL eligibility, generated-column guard, and version
|
||||
// fencing. Obviously ineligible paths avoid an unused get(); MemWAL
|
||||
// fallthrough and other local paths acquire/guard/plan independently.
|
||||
if let Some(stream) = try_execute_namespace_query(table, query).await? {
|
||||
return Ok(stream);
|
||||
}
|
||||
execute_generic_query(table, query, options).await
|
||||
}
|
||||
|
||||
async fn can_execute_namespace_query(table: &NativeTable, query: &AnyQuery) -> Result<bool> {
|
||||
/// Attempt QueryTable pushdown against one exact Dataset snapshot.
|
||||
///
|
||||
/// Returns `Ok(None)` when pushdown is ineligible so the caller can fall through
|
||||
/// to the local exact-snapshot planner. Does not `dataset.get()` on paths that
|
||||
/// are obviously ineligible before the MemWAL check.
|
||||
async fn try_execute_namespace_query(
|
||||
table: &NativeTable,
|
||||
query: &AnyQuery,
|
||||
) -> Result<Option<DatasetRecordBatchStream>> {
|
||||
if !(table
|
||||
.pushdown_operations
|
||||
.contains(&NamespaceClientPushdownOperation::QueryTable)
|
||||
@@ -72,17 +83,38 @@ async fn can_execute_namespace_query(table: &NativeTable, query: &AnyQuery) -> R
|
||||
&& table.dataset.current_branch().is_none()
|
||||
&& !requires_local_namespace_execution(query))
|
||||
{
|
||||
return Ok(false);
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(namespace_client) = table.namespace_client.clone() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
// One Dataset Arc owns MemWAL eligibility, the generated-column guard, and
|
||||
// the version fence sent on the request.
|
||||
let dataset = table.dataset.get().await?;
|
||||
// A MemWAL write spec means reads auto-route through the LSM scanner in
|
||||
// `create_plan` even when `use_lsm` is unset. The namespace request has no
|
||||
// use_lsm field, so pushing the default query down would silently omit
|
||||
// un-compacted rows — force local execution whenever a spec is installed.
|
||||
let dataset = table.dataset.get().await?;
|
||||
// Do not guard here; the local planner acquires its own snapshot.
|
||||
if dataset.mem_wal_index_details().await?.is_some() {
|
||||
return Ok(false);
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(true)
|
||||
|
||||
let snapshot = generated_column_binding_snapshot_from_dataset(dataset.as_ref())?;
|
||||
validate_generated_column_query(&snapshot, query)?;
|
||||
|
||||
let version = i64::try_from(dataset.version().version).map_err(|_| Error::InvalidInput {
|
||||
message: format!(
|
||||
"dataset version {} exceeds i64::MAX and cannot be sent on QueryTable",
|
||||
dataset.version().version
|
||||
),
|
||||
})?;
|
||||
|
||||
Ok(Some(
|
||||
execute_namespace_query(table, namespace_client, query, version).await?,
|
||||
))
|
||||
}
|
||||
|
||||
fn requires_local_namespace_execution(query: &AnyQuery) -> bool {
|
||||
@@ -128,18 +160,38 @@ async fn execute_generic_query(
|
||||
Ok(DatasetRecordBatchStream::new(inner))
|
||||
}
|
||||
|
||||
/// Public/internal Native planner entry: one Dataset get, guard, then plan.
|
||||
///
|
||||
/// Acquires exactly one [`Arc<Dataset>`], builds/validates the generated-column
|
||||
/// snapshot from that object, then plans entirely against the same Arc.
|
||||
/// Multi-vector recursion clones the owned Arc and must not call this outer
|
||||
/// entry (no additional `dataset.get()`).
|
||||
pub async fn create_plan(
|
||||
table: &NativeTable,
|
||||
query: &AnyQuery,
|
||||
options: QueryExecutionOptions,
|
||||
) -> Result<Arc<dyn ExecutionPlan>> {
|
||||
let ds_ref = table.dataset.get().await?;
|
||||
let snapshot = generated_column_binding_snapshot_from_dataset(ds_ref.as_ref())?;
|
||||
// Pass the original full AnyQuery before VectorQuery conversion/splitting so
|
||||
// select/filter/order/FTS/vector references are all covered. check_filter
|
||||
// precedence is preserved inside validate_generated_column_query.
|
||||
validate_generated_column_query(&snapshot, query)?;
|
||||
create_plan_with_dataset(table, query, options, ds_ref).await
|
||||
}
|
||||
|
||||
/// Plan against an already-owned Dataset Arc. Used by the guarded outer entry
|
||||
/// and by multi-vector recursion (Arc clones only).
|
||||
async fn create_plan_with_dataset(
|
||||
table: &NativeTable,
|
||||
query: &AnyQuery,
|
||||
options: QueryExecutionOptions,
|
||||
ds_ref: Arc<Dataset>,
|
||||
) -> Result<Arc<dyn ExecutionPlan>> {
|
||||
let query = match query {
|
||||
AnyQuery::VectorQuery(query) => query.clone(),
|
||||
AnyQuery::Query(query) => VectorQueryRequest::from_plain_query(query.clone()),
|
||||
};
|
||||
query.base.check_filter()?;
|
||||
|
||||
let ds_ref = table.dataset.get().await?;
|
||||
|
||||
// MemWAL read routing driven by `use_lsm`:
|
||||
// * unset — route through the LSM scanner iff the table carries a write spec
|
||||
@@ -198,7 +250,9 @@ pub async fn create_plan(
|
||||
}
|
||||
query_vector = Some(Arc::new(fsl_builder.finish()));
|
||||
} else {
|
||||
// Multiple query vectors: create a plan for each and union them
|
||||
// Multiple query vectors: create a plan for each and union them.
|
||||
// Recurse with clones of the already-owned Arc — never the outer
|
||||
// create_plan entry (which would perform another dataset.get()).
|
||||
let query_vecs = query.query_vector.clone();
|
||||
let plan_futures = query_vecs
|
||||
.into_iter()
|
||||
@@ -206,8 +260,15 @@ pub async fn create_plan(
|
||||
let mut sub_query = query.clone();
|
||||
sub_query.query_vector = vec![query_vector];
|
||||
let options_ref = options.clone();
|
||||
let ds_ref = ds_ref.clone();
|
||||
async move {
|
||||
create_plan(table, &AnyQuery::VectorQuery(sub_query), options_ref).await
|
||||
create_plan_with_dataset(
|
||||
table,
|
||||
&AnyQuery::VectorQuery(sub_query),
|
||||
options_ref,
|
||||
ds_ref,
|
||||
)
|
||||
.await
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
@@ -381,11 +442,15 @@ pub(crate) fn create_multi_vector_plan(
|
||||
}
|
||||
|
||||
/// Execute a query on the namespace server instead of locally.
|
||||
///
|
||||
/// Caller must already have validated the generated-column guard against the
|
||||
/// exact Dataset snapshot whose `version` is passed here. The incomplete /
|
||||
/// malformed guard runs before this dispatch.
|
||||
async fn execute_namespace_query(
|
||||
table: &NativeTable,
|
||||
namespace_client: Arc<dyn LanceNamespace>,
|
||||
query: &AnyQuery,
|
||||
_options: QueryExecutionOptions,
|
||||
version: i64,
|
||||
) -> Result<DatasetRecordBatchStream> {
|
||||
// Build table_id from namespace + table name
|
||||
let mut table_id = table.namespace.clone();
|
||||
@@ -393,8 +458,9 @@ async fn execute_namespace_query(
|
||||
|
||||
// Convert AnyQuery to namespace QueryTableRequest
|
||||
let mut ns_request = convert_to_namespace_query(query)?;
|
||||
// Set the table ID on the request
|
||||
// Set the table ID and exact guarded Dataset version on the request.
|
||||
ns_request.id = Some(table_id);
|
||||
ns_request.version = Some(version);
|
||||
|
||||
// Call the namespace query_table API
|
||||
let response_bytes = namespace_client
|
||||
@@ -1163,4 +1229,296 @@ mod tests {
|
||||
Some(ApproxMode::Accurate)
|
||||
);
|
||||
}
|
||||
|
||||
/// Records namespace `query_table` requests and returns a valid Arrow IPC file.
|
||||
#[derive(Debug, Default)]
|
||||
struct RecordingNamespaceClient {
|
||||
requests: std::sync::Mutex<Vec<NsQueryTableRequest>>,
|
||||
}
|
||||
|
||||
impl RecordingNamespaceClient {
|
||||
fn ipc_response() -> bytes::Bytes {
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_ipc::writer::FileWriter;
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new("gen_out", DataType::Int32, true),
|
||||
Field::new("ordinary", DataType::Utf8, true),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema.clone(),
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])) as ArrayRef,
|
||||
Arc::new(StringArray::from(vec![Some("x")])) as ArrayRef,
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
let mut buf = Vec::new();
|
||||
{
|
||||
let mut writer = FileWriter::try_new(&mut buf, &schema).unwrap();
|
||||
writer.write(&batch).unwrap();
|
||||
writer.finish().unwrap();
|
||||
}
|
||||
bytes::Bytes::from(buf)
|
||||
}
|
||||
|
||||
fn call_count(&self) -> usize {
|
||||
self.requests.lock().unwrap().len()
|
||||
}
|
||||
|
||||
fn requests(&self) -> Vec<NsQueryTableRequest> {
|
||||
self.requests.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl LanceNamespace for RecordingNamespaceClient {
|
||||
fn namespace_id(&self) -> String {
|
||||
"recording".to_string()
|
||||
}
|
||||
|
||||
async fn query_table(&self, request: NsQueryTableRequest) -> lance::Result<bytes::Bytes> {
|
||||
self.requests.lock().unwrap().push(request);
|
||||
Ok(Self::ipc_response())
|
||||
}
|
||||
}
|
||||
|
||||
async fn runtime_namespace_table(
|
||||
name: &str,
|
||||
) -> (
|
||||
crate::table::Table,
|
||||
NativeTable,
|
||||
Arc<RecordingNamespaceClient>,
|
||||
) {
|
||||
use crate::connect;
|
||||
use crate::connection::NamespaceClientPushdownOperation;
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
|
||||
let conn = connect("memory://").execute().await.unwrap();
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new("gen_out", DataType::Int32, true),
|
||||
Field::new("ordinary", DataType::Utf8, true),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
Arc::new(StringArray::from(vec![Some("x")])),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
let table = conn.create_table(name, batch).execute().await.unwrap();
|
||||
let namespace_client = Arc::new(RecordingNamespaceClient::default());
|
||||
let mut native_table = table.as_native().unwrap().clone();
|
||||
native_table.namespace_client = Some(namespace_client.clone());
|
||||
native_table
|
||||
.pushdown_operations
|
||||
.insert(NamespaceClientPushdownOperation::QueryTable);
|
||||
(table, native_table, namespace_client)
|
||||
}
|
||||
|
||||
async fn plant_runtime_generated_column_metadata(
|
||||
table: &crate::table::Table,
|
||||
column: &str,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) {
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput,
|
||||
FunctionParameter, FunctionSignature, GeneratedColumnDefinition,
|
||||
};
|
||||
use arrow_array::StringArray;
|
||||
|
||||
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
|
||||
let field_id = snapshot.field(column).expect(column).field_id();
|
||||
let function = Function::new(
|
||||
FunctionId::try_new("fn.exact.status.native").unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("ok")])) as ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
let json = GeneratedColumnDefinition::try_new(
|
||||
field_id,
|
||||
call,
|
||||
dependency_epoch,
|
||||
materialized_epoch,
|
||||
)
|
||||
.unwrap()
|
||||
.to_metadata_json()
|
||||
.unwrap();
|
||||
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
|
||||
table
|
||||
.as_native()
|
||||
.expect("generated-column fixture planting requires a Native table"),
|
||||
column,
|
||||
json,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
fn assert_incomplete_runtime_error(err: &Error, label: &str) {
|
||||
use crate::error::FunctionErrorCode;
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
|
||||
match err {
|
||||
Error::Function {
|
||||
code: FunctionErrorCode::GeneratedColumnIncomplete,
|
||||
message,
|
||||
} => {
|
||||
let rendered = format!("{err}\n{err:?}\n{message}");
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: leaked metadata key: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains("function_call"),
|
||||
"{label}: leaked function_call: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains("fn.exact.status.native"),
|
||||
"{label}: leaked Function ID: {rendered}"
|
||||
);
|
||||
}
|
||||
other => panic!(
|
||||
"{label}: expected Error::Function(GeneratedColumnIncomplete), got {other:?}"
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn generated_column_query_runtime_namespace_pushdown_fenced() {
|
||||
use crate::function::GeneratedColumnStatus;
|
||||
use crate::query::Select;
|
||||
|
||||
let (table, native_table, namespace_client) =
|
||||
runtime_namespace_table("runtime_ns_fence").await;
|
||||
|
||||
plant_runtime_generated_column_metadata(&table, "gen_out", 3, 3).await;
|
||||
assert_eq!(
|
||||
table.generated_column_status("gen_out").await.unwrap(),
|
||||
GeneratedColumnStatus::Complete
|
||||
);
|
||||
let guarded_version = native_table.dataset.get().await.unwrap().version().version;
|
||||
|
||||
let complete_query = AnyQuery::Query(QueryRequest {
|
||||
select: Select::Columns(vec!["gen_out".to_string()]),
|
||||
limit: Some(10),
|
||||
..Default::default()
|
||||
});
|
||||
let stream = execute_query(
|
||||
&native_table,
|
||||
&complete_query,
|
||||
QueryExecutionOptions::default(),
|
||||
)
|
||||
.await
|
||||
.expect("complete generated-column query must be eligible for QueryTable pushdown");
|
||||
let batches = stream.try_collect::<Vec<_>>().await.unwrap();
|
||||
assert_eq!(batches.iter().map(|b| b.num_rows()).sum::<usize>(), 1);
|
||||
assert_eq!(namespace_client.call_count(), 1);
|
||||
let requests = namespace_client.requests();
|
||||
assert_eq!(
|
||||
requests[0].version,
|
||||
Some(guarded_version as i64),
|
||||
"pushed-down query must fence to the exact guarded Dataset version; None races to latest"
|
||||
);
|
||||
|
||||
// Incomplete referenced output must reject before any namespace dispatch.
|
||||
plant_runtime_generated_column_metadata(&table, "gen_out", 8, 2).await;
|
||||
assert_eq!(
|
||||
table.generated_column_status("gen_out").await.unwrap(),
|
||||
GeneratedColumnStatus::Incomplete
|
||||
);
|
||||
let before_incomplete = namespace_client.call_count();
|
||||
let incomplete_query = AnyQuery::Query(QueryRequest {
|
||||
select: Select::Columns(vec!["gen_out".to_string()]),
|
||||
limit: Some(10),
|
||||
..Default::default()
|
||||
});
|
||||
let Err(incomplete_err) = execute_query(
|
||||
&native_table,
|
||||
&incomplete_query,
|
||||
QueryExecutionOptions::default(),
|
||||
)
|
||||
.await
|
||||
else {
|
||||
panic!("incomplete generated-column query must reject before dispatch");
|
||||
};
|
||||
assert_incomplete_runtime_error(&incomplete_err, "namespace_incomplete");
|
||||
assert_eq!(
|
||||
namespace_client.call_count(),
|
||||
before_incomplete,
|
||||
"incomplete guard must not dispatch query_table"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn generated_column_query_runtime_malformed_rejects_before_namespace() {
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
use crate::query::Select;
|
||||
|
||||
let (table, native_table, namespace_client) =
|
||||
runtime_namespace_table("runtime_ns_malformed").await;
|
||||
|
||||
const MARKER: &str = "SENSITIVE_RUNTIME_MALFORMED_b3e2a_9c1d";
|
||||
let raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":0,"function_call":{MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
|
||||
table
|
||||
.as_native()
|
||||
.expect("generated-column fixture planting requires a Native table"),
|
||||
"gen_out",
|
||||
raw.clone(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let query = AnyQuery::Query(QueryRequest {
|
||||
select: Select::Columns(vec!["gen_out".to_string()]),
|
||||
limit: Some(10),
|
||||
..Default::default()
|
||||
});
|
||||
let Err(err) = execute_query(&native_table, &query, QueryExecutionOptions::default()).await
|
||||
else {
|
||||
panic!("malformed referenced metadata must fail closed before dispatch");
|
||||
};
|
||||
assert!(
|
||||
matches!(err, Error::InvalidInput { .. }),
|
||||
"expected InvalidInput, got {err:?}"
|
||||
);
|
||||
let rendered = format!("{err}\n{err:?}");
|
||||
assert!(
|
||||
!rendered.contains(MARKER),
|
||||
"malformed diagnostics must not echo raw marker: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(&raw),
|
||||
"malformed diagnostics must not echo raw metadata: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"malformed diagnostics must not echo metadata wire key: {rendered}"
|
||||
);
|
||||
assert_eq!(
|
||||
namespace_client.call_count(),
|
||||
0,
|
||||
"malformed guard must not dispatch query_table"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,12 +8,94 @@
|
||||
//! - [`alter_columns`](execute_alter_columns): Rename columns, change types, or modify nullability
|
||||
//! - [`drop_columns`](execute_drop_columns): Remove columns from the table
|
||||
|
||||
use arrow_array::RecordBatchReader;
|
||||
use lance::dataset::{ColumnAlteration, NewColumnTransform};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
use super::NativeTable;
|
||||
use crate::Result;
|
||||
use crate::function::GENERATED_COLUMN_METADATA_KEY;
|
||||
use crate::function::schema_admission::reject_caller_authored_generated_column_schema;
|
||||
use crate::{Error, Result};
|
||||
|
||||
/// Reject caller-authored schema-bearing `add_columns` transforms that carry
|
||||
/// reserved generated-column top-level field metadata.
|
||||
///
|
||||
/// Borrows without consuming the transform: Stream is not polled, Reader is
|
||||
/// not iterated, and BatchUDF mapper is not invoked. `SqlExpressions` cannot
|
||||
/// carry an Arrow output schema and is accepted.
|
||||
pub(crate) fn reject_caller_authored_generated_column_add_columns_transform(
|
||||
transforms: &NewColumnTransform,
|
||||
) -> Result<()> {
|
||||
match transforms {
|
||||
NewColumnTransform::BatchUDF(udf) => {
|
||||
reject_caller_authored_generated_column_schema(udf.output_schema.as_ref())
|
||||
}
|
||||
NewColumnTransform::Stream(stream) => {
|
||||
reject_caller_authored_generated_column_schema(stream.schema().as_ref())
|
||||
}
|
||||
NewColumnTransform::Reader(reader) => {
|
||||
reject_caller_authored_generated_column_schema(reader.schema().as_ref())
|
||||
}
|
||||
NewColumnTransform::AllNulls(schema) => {
|
||||
reject_caller_authored_generated_column_schema(schema.as_ref())
|
||||
}
|
||||
NewColumnTransform::SqlExpressions(_) => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared rejection for general-purpose field-metadata updates that name the
|
||||
/// reserved generated-column definition key.
|
||||
///
|
||||
/// Generated-column definitions are table-schema state owned by
|
||||
/// create/change/refresh Jobs. The public `update_field_metadata` API must not
|
||||
/// create, replace, or remove that reserved key. Both Native and Remote
|
||||
/// `BaseTable` implementations call this helper so direct trait calls cannot
|
||||
/// bypass the syntax guard.
|
||||
pub(crate) fn reject_reserved_generated_column_metadata_key_updates(
|
||||
updates: &[FieldMetadataUpdate],
|
||||
) -> Result<()> {
|
||||
for update in updates {
|
||||
if update.metadata.contains_key(GENERATED_COLUMN_METADATA_KEY) {
|
||||
return Err(reserved_generated_column_metadata_not_supported());
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn reserved_generated_column_metadata_not_supported() -> Error {
|
||||
Error::NotSupported {
|
||||
message: "generated column definitions are owned by create/change/refresh Jobs \
|
||||
and cannot be created, replaced, or removed through update_field_metadata"
|
||||
.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Native-only state-aware guard: whole-map `replace()` on a field whose exact
|
||||
/// Dataset snapshot metadata already contains the reserved generated-column key
|
||||
/// would wipe that Job-owned definition even when the replacement map omits the
|
||||
/// key. Detects raw key presence without decoding the payload.
|
||||
fn reject_replace_that_would_remove_generated_column_metadata(
|
||||
dataset: &lance::Dataset,
|
||||
updates: &[FieldMetadataUpdate],
|
||||
) -> Result<()> {
|
||||
let schema = dataset.schema();
|
||||
for update in updates {
|
||||
if !update.replace {
|
||||
continue;
|
||||
}
|
||||
let Some(fields) = schema.resolve_case_insensitive(&update.path) else {
|
||||
continue;
|
||||
};
|
||||
let field = fields
|
||||
.last()
|
||||
.expect("resolve_case_insensitive returns a non-empty path");
|
||||
if field.metadata.contains_key(GENERATED_COLUMN_METADATA_KEY) {
|
||||
return Err(reserved_generated_column_metadata_not_supported());
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// The result of an add columns operation.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
|
||||
@@ -99,6 +181,7 @@ pub(crate) async fn execute_add_columns(
|
||||
transforms: NewColumnTransform,
|
||||
read_columns: Option<Vec<String>>,
|
||||
) -> Result<AddColumnsResult> {
|
||||
reject_caller_authored_generated_column_add_columns_transform(&transforms)?;
|
||||
table.dataset.ensure_mutable()?;
|
||||
let mut dataset = (*table.dataset.get().await?).clone();
|
||||
dataset.add_columns(transforms, read_columns, None).await?;
|
||||
@@ -144,8 +227,10 @@ pub(crate) async fn execute_update_field_metadata(
|
||||
table: &NativeTable,
|
||||
updates: &[FieldMetadataUpdate],
|
||||
) -> Result<UpdateFieldMetadataResult> {
|
||||
reject_reserved_generated_column_metadata_key_updates(updates)?;
|
||||
table.dataset.ensure_mutable()?;
|
||||
let mut dataset = (*table.dataset.get().await?).clone();
|
||||
reject_replace_that_would_remove_generated_column_metadata(&dataset, updates)?;
|
||||
|
||||
let mut builder = dataset.update_field_metadata();
|
||||
for update in updates {
|
||||
@@ -163,6 +248,33 @@ pub(crate) async fn execute_update_field_metadata(
|
||||
Ok(UpdateFieldMetadataResult { version })
|
||||
}
|
||||
|
||||
/// Test-only raw installer for generated-column field metadata on Native tables.
|
||||
///
|
||||
/// Uses the Lance metadata commit path directly so contract fixtures can plant
|
||||
/// reserved-key bytes without going through the public `update_field_metadata`
|
||||
/// guard. Absent from non-test builds.
|
||||
#[cfg(test)]
|
||||
pub(crate) async fn install_raw_generated_column_metadata_for_tests(
|
||||
table: &NativeTable,
|
||||
path: impl AsRef<str>,
|
||||
raw: impl Into<String>,
|
||||
) -> Result<UpdateFieldMetadataResult> {
|
||||
table.dataset.ensure_mutable()?;
|
||||
let mut dataset = (*table.dataset.get().await?).clone();
|
||||
let path = path.as_ref();
|
||||
let raw = raw.into();
|
||||
dataset
|
||||
.update_field_metadata()
|
||||
.update(
|
||||
path,
|
||||
[(GENERATED_COLUMN_METADATA_KEY.to_string(), Some(raw))],
|
||||
)?
|
||||
.await?;
|
||||
let version = dataset.version().version;
|
||||
table.dataset.update(dataset);
|
||||
Ok(UpdateFieldMetadataResult { version })
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use arrow_array::{Int32Array, StringArray, record_batch};
|
||||
|
||||
@@ -0,0 +1,306 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Dependency-contract test for Lance A4 / A4u / A4d schema metadata attachment (B4p).
|
||||
//!
|
||||
//! Pins the exact generic Lance API shape LanceDB B4 will consume:
|
||||
//! [`SchemaMetadataUpdates`], [`UpdateMap`], [`UpdateMapEntry`],
|
||||
//! [`Transaction::with_schema_metadata_updates`], and the public
|
||||
//! `with_schema_metadata_updates` methods on insert/update/delete builders.
|
||||
//!
|
||||
//! Also pins:
|
||||
//! - A4u Update no-op: an attached field metadata patch must accompany a real
|
||||
//! data change; when a predicate matches zero rows, `rows_updated == 0` and
|
||||
//! the patch must not be published.
|
||||
//! - A4d Delete no-op: when a predicate scans but deletes zero rows,
|
||||
//! `num_deleted_rows == 0` and the attached patch must not be published.
|
||||
//!
|
||||
//! Neutral metadata keys only. No Function / UDF / Job semantics.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema as ArrowSchema};
|
||||
use lance::Result;
|
||||
use lance::dataset::transaction::{
|
||||
Operation, SchemaMetadataUpdates, Transaction, UpdateMap, UpdateMapEntry,
|
||||
};
|
||||
use lance::dataset::{Dataset, DeleteBuilder, InsertBuilder, UpdateBuilder};
|
||||
use lance_table::format::Fragment;
|
||||
|
||||
const FIELD_ID: i32 = 7;
|
||||
const META_KEY: &str = "b4p.dependency.meta";
|
||||
const META_VALUE: &str = "neutral-value";
|
||||
|
||||
fn field_metadata_patch(field_id: i32) -> SchemaMetadataUpdates {
|
||||
SchemaMetadataUpdates {
|
||||
schema_metadata_updates: None,
|
||||
field_metadata_updates: HashMap::from([(
|
||||
field_id,
|
||||
UpdateMap {
|
||||
update_entries: vec![UpdateMapEntry {
|
||||
key: META_KEY.to_string(),
|
||||
value: Some(META_VALUE.to_string()),
|
||||
}],
|
||||
replace: false,
|
||||
},
|
||||
)]),
|
||||
}
|
||||
}
|
||||
|
||||
async fn write_neutral_fixture(uri: &str) -> Dataset {
|
||||
let schema = Arc::new(ArrowSchema::new(vec![
|
||||
Field::new("id", DataType::Int32, false),
|
||||
Field::new("value", DataType::Utf8, false),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema.clone(),
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1, 2, 3])),
|
||||
Arc::new(StringArray::from(vec!["a", "b", "c"])),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
Dataset::write(RecordBatchIterator::new(vec![Ok(batch)], schema), uri, None)
|
||||
.await
|
||||
.expect("fixture dataset must write")
|
||||
}
|
||||
|
||||
/// Compile-time proof that InsertBuilder exposes the A4 attachment method.
|
||||
#[allow(dead_code)]
|
||||
fn typecheck_insert_builder_attachment<'a>(
|
||||
builder: InsertBuilder<'a>,
|
||||
updates: SchemaMetadataUpdates,
|
||||
) -> Result<InsertBuilder<'a>> {
|
||||
builder.with_schema_metadata_updates(updates)
|
||||
}
|
||||
|
||||
/// Compile-time proof that UpdateBuilder exposes the A4 attachment method.
|
||||
#[allow(dead_code)]
|
||||
fn typecheck_update_builder_attachment(
|
||||
builder: UpdateBuilder,
|
||||
updates: SchemaMetadataUpdates,
|
||||
) -> Result<UpdateBuilder> {
|
||||
builder.with_schema_metadata_updates(updates)
|
||||
}
|
||||
|
||||
/// Compile-time proof that DeleteBuilder exposes the A4 attachment method.
|
||||
#[allow(dead_code)]
|
||||
fn typecheck_delete_builder_attachment(
|
||||
builder: DeleteBuilder,
|
||||
updates: SchemaMetadataUpdates,
|
||||
) -> Result<DeleteBuilder> {
|
||||
builder.with_schema_metadata_updates(updates)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn append_transaction_retains_schema_metadata_updates_patch() {
|
||||
let updates = field_metadata_patch(FIELD_ID);
|
||||
assert!(
|
||||
!updates.is_empty(),
|
||||
"fixture must be a substantive non-empty field metadata patch"
|
||||
);
|
||||
|
||||
let transaction = Transaction::new(
|
||||
0,
|
||||
Operation::Append {
|
||||
fragments: vec![Fragment::new(1)],
|
||||
},
|
||||
None,
|
||||
)
|
||||
.with_schema_metadata_updates(updates.clone())
|
||||
.expect("non-empty field metadata patch must attach to Append");
|
||||
|
||||
assert_eq!(transaction.schema_metadata_updates.as_ref(), Some(&updates));
|
||||
|
||||
let field_map = transaction
|
||||
.schema_metadata_updates
|
||||
.as_ref()
|
||||
.expect("attached patch must be present")
|
||||
.field_metadata_updates
|
||||
.get(&FIELD_ID)
|
||||
.expect("stable field id 7 must be present");
|
||||
assert!(!field_map.replace);
|
||||
assert_eq!(field_map.update_entries.len(), 1);
|
||||
assert_eq!(field_map.update_entries[0].key, META_KEY);
|
||||
assert_eq!(
|
||||
field_map.update_entries[0].value.as_deref(),
|
||||
Some(META_VALUE)
|
||||
);
|
||||
}
|
||||
|
||||
/// A4u dependency: a no-op Update (predicate matches zero rows) must not
|
||||
/// publish an attached field metadata patch. Manifest version advancement is
|
||||
/// unconstrained.
|
||||
#[tokio::test]
|
||||
async fn noop_update_does_not_publish_attached_field_metadata() {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().join("noop_update.lance");
|
||||
let uri = uri.to_str().unwrap();
|
||||
|
||||
let dataset = write_neutral_fixture(uri).await;
|
||||
|
||||
let field = dataset
|
||||
.schema()
|
||||
.field("value")
|
||||
.expect("value column must exist");
|
||||
let field_id = field.id;
|
||||
assert!(
|
||||
!field.metadata.contains_key(META_KEY),
|
||||
"{META_KEY} must be initially absent, got {:?}",
|
||||
field.metadata
|
||||
);
|
||||
|
||||
let updates = field_metadata_patch(field_id);
|
||||
assert!(
|
||||
!updates.is_empty(),
|
||||
"fixture must be a substantive non-empty field metadata patch"
|
||||
);
|
||||
|
||||
let before_count = dataset.count_rows(None).await.unwrap();
|
||||
assert_eq!(before_count, 3);
|
||||
|
||||
let result = UpdateBuilder::new(Arc::new(dataset))
|
||||
.update_where("id < 0")
|
||||
.unwrap()
|
||||
.set("value", "'changed'")
|
||||
.unwrap()
|
||||
.with_schema_metadata_updates(updates)
|
||||
.expect("Update attachment must construct")
|
||||
.build()
|
||||
.unwrap()
|
||||
.execute()
|
||||
.await
|
||||
.expect("no-op attached Update must complete");
|
||||
|
||||
assert_eq!(result.rows_updated, 0, "predicate must match zero rows");
|
||||
assert_eq!(
|
||||
result.new_dataset.count_rows(None).await.unwrap(),
|
||||
before_count,
|
||||
"row count must remain unchanged"
|
||||
);
|
||||
assert_eq!(
|
||||
result
|
||||
.new_dataset
|
||||
.count_rows(Some("value = 'changed'".into()))
|
||||
.await
|
||||
.unwrap(),
|
||||
0,
|
||||
"SET expression must not rewrite any rows"
|
||||
);
|
||||
assert_eq!(
|
||||
result
|
||||
.new_dataset
|
||||
.count_rows(Some("value IN ('a', 'b', 'c')".into()))
|
||||
.await
|
||||
.unwrap(),
|
||||
before_count,
|
||||
"original values must remain unchanged"
|
||||
);
|
||||
|
||||
let reopened = Dataset::open(uri).await.unwrap();
|
||||
assert_eq!(reopened.count_rows(None).await.unwrap(), before_count);
|
||||
assert_eq!(
|
||||
reopened
|
||||
.count_rows(Some("value = 'changed'".into()))
|
||||
.await
|
||||
.unwrap(),
|
||||
0
|
||||
);
|
||||
let reopened_field = reopened
|
||||
.schema()
|
||||
.field_by_id(field_id)
|
||||
.expect("stable field id must still exist");
|
||||
assert!(
|
||||
!reopened_field.metadata.contains_key(META_KEY),
|
||||
"no-op Update must not publish attached field metadata; got {:?}",
|
||||
reopened_field.metadata.get(META_KEY)
|
||||
);
|
||||
}
|
||||
|
||||
/// A4d dependency: a no-op Delete (predicate scans but matches zero rows) must
|
||||
/// not publish an attached field metadata patch. Manifest version advancement
|
||||
/// is unconstrained.
|
||||
#[tokio::test]
|
||||
async fn noop_delete_does_not_publish_attached_field_metadata() {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().join("noop_delete.lance");
|
||||
let uri = uri.to_str().unwrap();
|
||||
|
||||
let dataset = write_neutral_fixture(uri).await;
|
||||
|
||||
let field = dataset
|
||||
.schema()
|
||||
.field("value")
|
||||
.expect("value column must exist");
|
||||
let field_id = field.id;
|
||||
assert!(
|
||||
!field.metadata.contains_key(META_KEY),
|
||||
"{META_KEY} must be initially absent, got {:?}",
|
||||
field.metadata
|
||||
);
|
||||
|
||||
let updates = field_metadata_patch(field_id);
|
||||
assert!(
|
||||
!updates.is_empty(),
|
||||
"fixture must be a substantive non-empty field metadata patch"
|
||||
);
|
||||
|
||||
let before_count = dataset.count_rows(None).await.unwrap();
|
||||
assert_eq!(before_count, 3);
|
||||
|
||||
let result = DeleteBuilder::new(Arc::new(dataset), "id < 0")
|
||||
.with_schema_metadata_updates(updates)
|
||||
.expect("Delete attachment must construct")
|
||||
.execute()
|
||||
.await
|
||||
.expect("no-op attached Delete must complete");
|
||||
|
||||
assert_eq!(result.num_deleted_rows, 0, "predicate must match zero rows");
|
||||
assert_eq!(
|
||||
result.new_dataset.count_rows(None).await.unwrap(),
|
||||
before_count,
|
||||
"row count must remain unchanged"
|
||||
);
|
||||
assert_eq!(
|
||||
result
|
||||
.new_dataset
|
||||
.count_rows(Some("value IN ('a', 'b', 'c')".into()))
|
||||
.await
|
||||
.unwrap(),
|
||||
before_count,
|
||||
"original values must remain unchanged"
|
||||
);
|
||||
|
||||
let returned_field = result
|
||||
.new_dataset
|
||||
.schema()
|
||||
.field_by_id(field_id)
|
||||
.expect("stable field id must still exist on returned dataset");
|
||||
assert!(
|
||||
!returned_field.metadata.contains_key(META_KEY),
|
||||
"no-op Delete must not publish attached field metadata on returned dataset; got {:?}",
|
||||
returned_field.metadata.get(META_KEY)
|
||||
);
|
||||
|
||||
let reopened = Dataset::open(uri).await.unwrap();
|
||||
assert_eq!(reopened.count_rows(None).await.unwrap(), before_count);
|
||||
assert_eq!(
|
||||
reopened
|
||||
.count_rows(Some("value IN ('a', 'b', 'c')".into()))
|
||||
.await
|
||||
.unwrap(),
|
||||
before_count,
|
||||
"fresh open must preserve all original rows"
|
||||
);
|
||||
let reopened_field = reopened
|
||||
.schema()
|
||||
.field_by_id(field_id)
|
||||
.expect("stable field id must still exist");
|
||||
assert!(
|
||||
!reopened_field.metadata.contains_key(META_KEY),
|
||||
"no-op Delete must not publish attached field metadata; got {:?}",
|
||||
reopened_field.metadata.get(META_KEY)
|
||||
);
|
||||
}
|
||||
@@ -1,8 +1,10 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
use std::sync::Arc;
|
||||
|
||||
use lance::Dataset;
|
||||
use lance::dataset::UpdateBuilder as LanceUpdateBuilder;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
@@ -80,27 +82,37 @@ pub(crate) async fn execute_update(
|
||||
) -> Result<UpdateResult> {
|
||||
table.dataset.ensure_mutable()?;
|
||||
|
||||
// 1. Snapshot the current dataset
|
||||
// One exact dataset supplies SET/filter planning, stable target field IDs,
|
||||
// generated-column invalidation planning, the Lance UpdateBuilder, and its
|
||||
// transaction basis. Do not call table schema()/version() or another get().
|
||||
let dataset = table.dataset.get().await?;
|
||||
|
||||
// 2. Initialize the Lance Core builder
|
||||
let mut builder = LanceUpdateBuilder::new(dataset);
|
||||
let mut builder = LanceUpdateBuilder::new(dataset.clone());
|
||||
|
||||
// 3. Apply the filter (WHERE clause)
|
||||
if let Some(predicate) = update.filter {
|
||||
builder = builder.update_where(&predicate)?;
|
||||
}
|
||||
|
||||
// 4. Apply the columns (SET clause)
|
||||
for (column, value) in update.columns {
|
||||
builder = builder.set(column, &value)?;
|
||||
let columns = update.columns;
|
||||
for (column, value) in &columns {
|
||||
builder = builder.set(column, value)?;
|
||||
}
|
||||
|
||||
// After Lance SET validation, resolve stable field IDs from the same
|
||||
// snapshot and plan invalidation before UpdateJob writes files.
|
||||
let updated_field_ids = updated_stable_field_ids(dataset.as_ref(), &columns)?;
|
||||
if let Some(schema_metadata_updates) =
|
||||
super::generated_column_invalidation::plan_native_update_generated_column_invalidation(
|
||||
dataset.as_ref(),
|
||||
updated_field_ids,
|
||||
)?
|
||||
{
|
||||
builder = builder.with_schema_metadata_updates(schema_metadata_updates)?;
|
||||
}
|
||||
|
||||
// 5. Execute the operation (Write new files)
|
||||
let operation = builder.build()?;
|
||||
let res = operation.execute().await?;
|
||||
|
||||
// 6. Update the table's view of the latest version
|
||||
table.dataset.update(res.new_dataset.as_ref().clone());
|
||||
|
||||
Ok(UpdateResult {
|
||||
@@ -109,6 +121,24 @@ pub(crate) async fn execute_update(
|
||||
})
|
||||
}
|
||||
|
||||
/// Resolve exact top-level SET targets to a deterministic set of stable field IDs.
|
||||
fn updated_stable_field_ids(
|
||||
dataset: &Dataset,
|
||||
columns: &[(String, String)],
|
||||
) -> Result<BTreeSet<i32>> {
|
||||
let mut ids = BTreeSet::new();
|
||||
for (column, _) in columns {
|
||||
let field = dataset
|
||||
.schema()
|
||||
.field(column)
|
||||
.ok_or_else(|| Error::InvalidInput {
|
||||
message: format!("Column '{column}' does not exist in dataset schema"),
|
||||
})?;
|
||||
ids.insert(field.id);
|
||||
}
|
||||
Ok(ids)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::connect;
|
||||
|
||||
@@ -0,0 +1,627 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Runtime contract tests for the B4f reserved generated-metadata update guard.
|
||||
//!
|
||||
//! Pins that the general-purpose [`crate::table::Table::update_field_metadata`]
|
||||
//! API cannot create, replace, or remove `GENERATED_COLUMN_METADATA_KEY`, and
|
||||
//! that Native `replace()` cannot wipe an existing generated definition by
|
||||
//! omitting the reserved key. Remote explicit-key attempts must reject before
|
||||
//! transport.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use futures::TryStreamExt;
|
||||
use tempfile::TempDir;
|
||||
|
||||
use crate::connection::ConnectBuilder;
|
||||
use crate::error::Error;
|
||||
use crate::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition,
|
||||
};
|
||||
use crate::query::{ExecutableQuery, QueryBase, Select};
|
||||
use crate::table::Table;
|
||||
use crate::table::schema_evolution::FieldMetadataUpdate;
|
||||
|
||||
const GEN_OUT: &str = "gen_out";
|
||||
const ORDINARY: &str = "ordinary";
|
||||
const CATEGORY: &str = "category";
|
||||
const FN_ID: &str = "fn.exact.b4f.guard.literal";
|
||||
const MALFORMED_MARKER: &str = "SENSITIVE_B4F_GUARD_METADATA_MARKER_7c1e_d04b";
|
||||
|
||||
struct Fixture {
|
||||
_tmp: TempDir,
|
||||
table: Table,
|
||||
uri: String,
|
||||
}
|
||||
|
||||
fn literal_definition(
|
||||
output_field_id: i32,
|
||||
dependency_epoch: u64,
|
||||
materialized_epoch: u64,
|
||||
) -> GeneratedColumnDefinition {
|
||||
let function = Function::new(
|
||||
FunctionId::try_new(FN_ID).unwrap(),
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("label", DataType::Utf8)],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"label".to_string(),
|
||||
FunctionArgument::try_literal(
|
||||
Arc::new(StringArray::from(vec![Some("b4f-guard")])) as arrow_array::ArrayRef
|
||||
)
|
||||
.unwrap(),
|
||||
)],
|
||||
)
|
||||
.unwrap();
|
||||
GeneratedColumnDefinition::try_new(output_field_id, call, dependency_epoch, materialized_epoch)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
async fn create_ordinary_table(name: &str) -> Fixture {
|
||||
let tmp = tempfile::tempdir().unwrap();
|
||||
let uri = tmp.path().to_str().unwrap().to_string();
|
||||
let conn = ConnectBuilder::new(&uri).execute().await.unwrap();
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new(GEN_OUT, DataType::Int32, true),
|
||||
Field::new(ORDINARY, DataType::Utf8, true),
|
||||
Field::new(CATEGORY, DataType::Utf8, true),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
Arc::new(StringArray::from(vec![Some("seed")])),
|
||||
Arc::new(StringArray::from(vec![Some("A")])),
|
||||
],
|
||||
)
|
||||
.unwrap();
|
||||
let table = conn.create_table(name, batch).execute().await.unwrap();
|
||||
Fixture {
|
||||
_tmp: tmp,
|
||||
table,
|
||||
uri,
|
||||
}
|
||||
}
|
||||
|
||||
async fn plant_generated_raw(table: &Table, column: &str, raw: String) {
|
||||
crate::table::schema_evolution::install_raw_generated_column_metadata_for_tests(
|
||||
table
|
||||
.as_native()
|
||||
.expect("generated-column fixture planting requires a Native table"),
|
||||
column,
|
||||
raw,
|
||||
)
|
||||
.await
|
||||
.expect("fixture raw generated-column metadata install must succeed");
|
||||
}
|
||||
|
||||
async fn plant_valid_generated(table: &Table, column: &str) -> String {
|
||||
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
|
||||
let field_id = snapshot.field(column).expect(column).field_id();
|
||||
let raw = literal_definition(field_id, 3, 3)
|
||||
.to_metadata_json()
|
||||
.unwrap();
|
||||
plant_generated_raw(table, column, raw.clone()).await;
|
||||
raw
|
||||
}
|
||||
|
||||
async fn read_raw_generated_metadata(table: &Table, column: &str) -> Option<String> {
|
||||
let snapshot = table.generated_column_binding_snapshot().await.unwrap();
|
||||
snapshot
|
||||
.field(column)
|
||||
.expect(column)
|
||||
.field()
|
||||
.metadata()
|
||||
.get(GENERATED_COLUMN_METADATA_KEY)
|
||||
.cloned()
|
||||
}
|
||||
|
||||
async fn ordinary_values(table: &Table) -> Vec<String> {
|
||||
let batches = table
|
||||
.query()
|
||||
.select(Select::columns(&[ORDINARY]))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
batches
|
||||
.iter()
|
||||
.flat_map(|batch| {
|
||||
batch
|
||||
.column(0)
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.map(|v| v.unwrap().to_string())
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn reopen(uri: &str, name: &str) -> Table {
|
||||
ConnectBuilder::new(uri)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.open_table(name)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn assert_not_supported_redacted(err: &Error, label: &str, forbidden_substrings: &[&str]) {
|
||||
match err {
|
||||
Error::NotSupported { message } => {
|
||||
let rendered = format!("{err}\n{err:?}\n{message}");
|
||||
assert!(
|
||||
!rendered.contains(GENERATED_COLUMN_METADATA_KEY),
|
||||
"{label}: leaked metadata wire key: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(FN_ID),
|
||||
"{label}: leaked Function ID: {rendered}"
|
||||
);
|
||||
assert!(
|
||||
!rendered.contains(MALFORMED_MARKER),
|
||||
"{label}: leaked malformed marker: {rendered}"
|
||||
);
|
||||
for needle in forbidden_substrings {
|
||||
assert!(
|
||||
!rendered.contains(needle),
|
||||
"{label}: leaked forbidden substring `{needle}`: {rendered}"
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
message.to_lowercase().contains("generated")
|
||||
|| message.to_lowercase().contains("job"),
|
||||
"{label}: message must describe Job-owned generated-column boundary: {message}"
|
||||
);
|
||||
}
|
||||
other => panic!("{label}: expected Error::NotSupported, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn native_ordinary_field_explicit_reserved_key_set_rejects_and_preserves_state() {
|
||||
let fixture = create_ordinary_table("b4f_ordinary_set").await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let rows_before = ordinary_values(&fixture.table).await;
|
||||
let schema_before = fixture.table.schema().await.unwrap();
|
||||
let category_md_before = schema_before
|
||||
.field_with_name(CATEGORY)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.clone();
|
||||
|
||||
let snapshot = fixture
|
||||
.table
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.unwrap();
|
||||
let field_id = snapshot.field(CATEGORY).expect(CATEGORY).field_id();
|
||||
let payload = literal_definition(field_id, 1, 1)
|
||||
.to_metadata_json()
|
||||
.unwrap();
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.update_field_metadata(&[
|
||||
FieldMetadataUpdate::new(CATEGORY).set(GENERATED_COLUMN_METADATA_KEY, payload.clone())
|
||||
])
|
||||
.await
|
||||
.expect_err("explicit reserved-key set on ordinary field must reject");
|
||||
assert_not_supported_redacted(&err, "ordinary reserved set", &[&payload]);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(ordinary_values(&fixture.table).await, rows_before);
|
||||
let schema_after = fixture.table.schema().await.unwrap();
|
||||
assert_eq!(
|
||||
schema_after.field_with_name(CATEGORY).unwrap().metadata(),
|
||||
&category_md_before
|
||||
);
|
||||
assert!(
|
||||
read_raw_generated_metadata(&fixture.table, CATEGORY)
|
||||
.await
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn native_generated_field_explicit_remove_rejects_and_preserves_raw() {
|
||||
let fixture = create_ordinary_table("b4f_gen_remove").await;
|
||||
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.update_field_metadata(&[
|
||||
FieldMetadataUpdate::new(GEN_OUT).remove(GENERATED_COLUMN_METADATA_KEY)
|
||||
])
|
||||
.await
|
||||
.expect_err("explicit reserved-key remove must reject");
|
||||
assert_not_supported_redacted(&err, "generated remove", &[&planted]);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fixture.table, GEN_OUT)
|
||||
.await
|
||||
.as_deref(),
|
||||
Some(planted.as_str())
|
||||
);
|
||||
let fresh = reopen(&fixture.uri, "b4f_gen_remove").await;
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fresh, GEN_OUT)
|
||||
.await
|
||||
.as_deref(),
|
||||
Some(planted.as_str())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn native_generated_field_explicit_replacement_rejects_and_preserves_raw() {
|
||||
let fixture = create_ordinary_table("b4f_gen_replace_value").await;
|
||||
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
let replacement = literal_definition(
|
||||
fixture
|
||||
.table
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.unwrap()
|
||||
.field(GEN_OUT)
|
||||
.unwrap()
|
||||
.field_id(),
|
||||
9,
|
||||
1,
|
||||
)
|
||||
.to_metadata_json()
|
||||
.unwrap();
|
||||
assert_ne!(replacement, planted);
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(GEN_OUT)
|
||||
.set(GENERATED_COLUMN_METADATA_KEY, replacement.clone())])
|
||||
.await
|
||||
.expect_err("explicit reserved-key replacement must reject");
|
||||
assert_not_supported_redacted(&err, "generated replace value", &[&planted, &replacement]);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
let fresh = reopen(&fixture.uri, "b4f_gen_replace_value").await;
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fresh, GEN_OUT)
|
||||
.await
|
||||
.as_deref(),
|
||||
Some(planted.as_str())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn native_generated_field_replace_with_ordinary_metadata_rejects_and_preserves_raw() {
|
||||
let fixture = create_ordinary_table("b4f_gen_replace_map").await;
|
||||
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(GEN_OUT)
|
||||
.replace()
|
||||
.set("unit", "label")])
|
||||
.await
|
||||
.expect_err("replace() that would wipe generated definition must reject");
|
||||
assert_not_supported_redacted(&err, "generated replace map", &[&planted]);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
let fresh = reopen(&fixture.uri, "b4f_gen_replace_map").await;
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fresh, GEN_OUT)
|
||||
.await
|
||||
.as_deref(),
|
||||
Some(planted.as_str())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn native_mixed_batch_rejects_atomically_no_partial_commit() {
|
||||
let fixture = create_ordinary_table("b4f_mixed_batch").await;
|
||||
let planted = plant_valid_generated(&fixture.table, GEN_OUT).await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.update_field_metadata(&[
|
||||
FieldMetadataUpdate::new(CATEGORY).set("unit", "label"),
|
||||
FieldMetadataUpdate::new(GEN_OUT).remove(GENERATED_COLUMN_METADATA_KEY),
|
||||
])
|
||||
.await
|
||||
.expect_err("mixed batch with forbidden update must reject all-or-none");
|
||||
assert_not_supported_redacted(&err, "mixed forbidden second", &[&planted]);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
let schema = fixture.table.schema().await.unwrap();
|
||||
assert!(
|
||||
!schema
|
||||
.field_with_name(CATEGORY)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key("unit"),
|
||||
"ordinary metadata must not partially commit"
|
||||
);
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.update_field_metadata(&[
|
||||
FieldMetadataUpdate::new(GEN_OUT).set(GENERATED_COLUMN_METADATA_KEY, planted.clone()),
|
||||
FieldMetadataUpdate::new(CATEGORY).set("unit", "label"),
|
||||
])
|
||||
.await
|
||||
.expect_err("mixed batch with forbidden update first must reject all-or-none");
|
||||
assert_not_supported_redacted(&err, "mixed forbidden first", &[&planted]);
|
||||
|
||||
let fresh = reopen(&fixture.uri, "b4f_mixed_batch").await;
|
||||
assert_eq!(fresh.version().await.unwrap(), version_before);
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fresh, GEN_OUT)
|
||||
.await
|
||||
.as_deref(),
|
||||
Some(planted.as_str())
|
||||
);
|
||||
let fresh_schema = fresh.schema().await.unwrap();
|
||||
assert!(
|
||||
!fresh_schema
|
||||
.field_with_name(CATEGORY)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.contains_key("unit")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn native_malformed_generated_raw_replace_rejects_redacted_and_preserves_raw() {
|
||||
let fixture = create_ordinary_table("b4f_malformed_replace").await;
|
||||
let field_id = fixture
|
||||
.table
|
||||
.generated_column_binding_snapshot()
|
||||
.await
|
||||
.unwrap()
|
||||
.field(GEN_OUT)
|
||||
.unwrap()
|
||||
.field_id();
|
||||
let planted_raw = format!(
|
||||
r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{MALFORMED_MARKER},"dependency_epoch":1,"materialized_epoch":1}}"#
|
||||
);
|
||||
assert!(planted_raw.contains(MALFORMED_MARKER));
|
||||
plant_generated_raw(&fixture.table, GEN_OUT, planted_raw.clone()).await;
|
||||
let version_before = fixture.table.version().await.unwrap();
|
||||
|
||||
let err = fixture
|
||||
.table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(GEN_OUT)
|
||||
.replace()
|
||||
.set("unit", "label")])
|
||||
.await
|
||||
.expect_err("malformed generated raw must still block replace()");
|
||||
assert_not_supported_redacted(&err, "malformed replace", &[&planted_raw]);
|
||||
|
||||
assert_eq!(fixture.table.version().await.unwrap(), version_before);
|
||||
let fresh = reopen(&fixture.uri, "b4f_malformed_replace").await;
|
||||
assert_eq!(
|
||||
read_raw_generated_metadata(&fresh, GEN_OUT)
|
||||
.await
|
||||
.as_deref(),
|
||||
Some(planted_raw.as_str())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn native_ordinary_metadata_merge_set_remove_replace_still_work() {
|
||||
let fixture = create_ordinary_table("b4f_ordinary_controls").await;
|
||||
|
||||
fixture
|
||||
.table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
|
||||
.set("unit", "label")
|
||||
.set("pii", "false")])
|
||||
.await
|
||||
.unwrap();
|
||||
let md = fixture
|
||||
.table
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name(CATEGORY)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.clone();
|
||||
assert_eq!(md.get("unit").map(String::as_str), Some("label"));
|
||||
assert_eq!(md.get("pii").map(String::as_str), Some("false"));
|
||||
|
||||
fixture
|
||||
.table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
|
||||
.set("source", "import")
|
||||
.remove("pii")])
|
||||
.await
|
||||
.unwrap();
|
||||
let md = fixture
|
||||
.table
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name(CATEGORY)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.clone();
|
||||
assert_eq!(md.get("unit").map(String::as_str), Some("label"));
|
||||
assert_eq!(md.get("source").map(String::as_str), Some("import"));
|
||||
assert!(!md.contains_key("pii"));
|
||||
|
||||
fixture
|
||||
.table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
|
||||
.replace()
|
||||
.set("only", "kept")])
|
||||
.await
|
||||
.unwrap();
|
||||
let md = fixture
|
||||
.table
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name(CATEGORY)
|
||||
.unwrap()
|
||||
.metadata()
|
||||
.clone();
|
||||
assert_eq!(md.len(), 1);
|
||||
assert_eq!(md.get("only").map(String::as_str), Some("kept"));
|
||||
assert!(
|
||||
read_raw_generated_metadata(&fixture.table, CATEGORY)
|
||||
.await
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(feature = "remote")]
|
||||
mod remote_explicit_key_guard {
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::*;
|
||||
use crate::Error;
|
||||
use crate::remote::{ClientConfig, HeaderProvider};
|
||||
|
||||
#[derive(Debug)]
|
||||
struct CountingHeaderProvider {
|
||||
calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl HeaderProvider for CountingHeaderProvider {
|
||||
async fn get_headers(&self) -> crate::Result<HashMap<String, String>> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(HashMap::from([(
|
||||
"X-Test-Header".to_string(),
|
||||
"must-not-be-requested".to_string(),
|
||||
)]))
|
||||
}
|
||||
}
|
||||
|
||||
fn panic_handler(
|
||||
calls: Arc<AtomicUsize>,
|
||||
) -> impl Fn(reqwest::Request) -> http::Response<String> + Clone + Send + Sync + 'static {
|
||||
move |_request| {
|
||||
calls.fetch_add(1, Ordering::SeqCst);
|
||||
panic!("remote reserved-key update must not invoke the HTTP handler");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_explicit_set_rejects_before_handler_and_header_provider() {
|
||||
let handler_calls = Arc::new(AtomicUsize::new(0));
|
||||
let header_calls = Arc::new(AtomicUsize::new(0));
|
||||
let config = ClientConfig {
|
||||
header_provider: Some(Arc::new(CountingHeaderProvider {
|
||||
calls: header_calls.clone(),
|
||||
}) as Arc<dyn HeaderProvider>),
|
||||
..Default::default()
|
||||
};
|
||||
let table = Table::new_with_handler_and_config(
|
||||
"my_table",
|
||||
panic_handler(handler_calls.clone()),
|
||||
config,
|
||||
);
|
||||
|
||||
let err = table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY)
|
||||
.set(GENERATED_COLUMN_METADATA_KEY, r#"{"format_version":1}"#)])
|
||||
.await
|
||||
.expect_err("remote explicit reserved-key set must reject");
|
||||
assert!(
|
||||
matches!(err, Error::NotSupported { .. }),
|
||||
"expected NotSupported, got {err:?}"
|
||||
);
|
||||
assert_eq!(handler_calls.load(Ordering::SeqCst), 0);
|
||||
assert_eq!(header_calls.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_explicit_remove_rejects_before_handler_and_header_provider() {
|
||||
let handler_calls = Arc::new(AtomicUsize::new(0));
|
||||
let header_calls = Arc::new(AtomicUsize::new(0));
|
||||
let config = ClientConfig {
|
||||
header_provider: Some(Arc::new(CountingHeaderProvider {
|
||||
calls: header_calls.clone(),
|
||||
}) as Arc<dyn HeaderProvider>),
|
||||
..Default::default()
|
||||
};
|
||||
let table = Table::new_with_handler_and_config(
|
||||
"my_table",
|
||||
panic_handler(handler_calls.clone()),
|
||||
config,
|
||||
);
|
||||
|
||||
let err = table
|
||||
.update_field_metadata(&[
|
||||
FieldMetadataUpdate::new(CATEGORY).remove(GENERATED_COLUMN_METADATA_KEY)
|
||||
])
|
||||
.await
|
||||
.expect_err("remote explicit reserved-key remove must reject");
|
||||
assert!(
|
||||
matches!(err, Error::NotSupported { .. }),
|
||||
"expected NotSupported, got {err:?}"
|
||||
);
|
||||
assert_eq!(handler_calls.load(Ordering::SeqCst), 0);
|
||||
assert_eq!(header_calls.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn remote_ordinary_metadata_update_sends_exact_body_and_succeeds() {
|
||||
let table = Table::new_with_handler("my_table", |request| {
|
||||
assert_eq!(request.method(), "POST");
|
||||
assert_eq!(
|
||||
request.url().path(),
|
||||
"/v1/table/my_table/update_field_metadata/"
|
||||
);
|
||||
let body = request
|
||||
.body()
|
||||
.expect("ordinary update must send a body")
|
||||
.as_bytes()
|
||||
.expect("body is in-memory");
|
||||
let parsed: serde_json::Value = serde_json::from_slice(body).unwrap();
|
||||
assert_eq!(
|
||||
parsed,
|
||||
serde_json::json!({
|
||||
"updates": [{
|
||||
"path": "category",
|
||||
"metadata": { "unit": "label" },
|
||||
"replace": false
|
||||
}]
|
||||
})
|
||||
);
|
||||
http::Response::builder()
|
||||
.status(200)
|
||||
.body(r#"{"version": 7}"#.to_string())
|
||||
.unwrap()
|
||||
});
|
||||
|
||||
let result = table
|
||||
.update_field_metadata(&[FieldMetadataUpdate::new(CATEGORY).set("unit", "label")])
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(result.version, 7);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -10,7 +10,7 @@ use arrow_array::{
|
||||
use arrow_schema::{DataType, Field, Fields, Schema};
|
||||
use futures::TryStreamExt;
|
||||
use lance::Dataset;
|
||||
use lance_file::version::LanceFileVersion;
|
||||
use lance_file::version::{ConcreteFileVersion, LanceFileVersion};
|
||||
use lancedb::{
|
||||
Connection, Error, Result, Table,
|
||||
blob::{BlobRangeRequest, blob},
|
||||
@@ -61,7 +61,7 @@ async fn create_inline_blob_table(
|
||||
Ok(table)
|
||||
}
|
||||
|
||||
async fn storage_format_version(table: &Table) -> LanceFileVersion {
|
||||
async fn storage_format_version(table: &Table) -> ConcreteFileVersion {
|
||||
table
|
||||
.as_native()
|
||||
.unwrap()
|
||||
@@ -69,9 +69,18 @@ async fn storage_format_version(table: &Table) -> LanceFileVersion {
|
||||
.await
|
||||
.unwrap()
|
||||
.data_storage_format
|
||||
.lance_file_version()
|
||||
.unwrap()
|
||||
.resolve()
|
||||
.lance_file_format()
|
||||
}
|
||||
|
||||
/// Blob v2 storage capability for the current concrete formats.
|
||||
///
|
||||
/// Exact formats deliberately have no Ord: capability is not implied by release
|
||||
/// order. Enumerate every current concrete variant explicitly.
|
||||
fn supports_blob_v2_storage(version: ConcreteFileVersion) -> bool {
|
||||
match version {
|
||||
ConcreteFileVersion::V2_2 | ConcreteFileVersion::V2_3 => true,
|
||||
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1 => false,
|
||||
}
|
||||
}
|
||||
|
||||
async fn uses_stable_row_ids(table: &Table) -> bool {
|
||||
@@ -112,7 +121,9 @@ async fn declaring_blob_column_bumps_format_and_enables_stable_row_ids() -> Resu
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
|
||||
assert!(supports_blob_v2_storage(
|
||||
storage_format_version(&table).await
|
||||
));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -127,7 +138,9 @@ async fn explicit_stable_row_id_setting_wins_over_blob_default() -> Result<()> {
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
|
||||
assert!(supports_blob_v2_storage(
|
||||
storage_format_version(&table).await
|
||||
));
|
||||
assert!(!uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -139,7 +152,9 @@ async fn non_blob_table_keeps_default_format_and_row_id_setting() -> Result<()>
|
||||
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int64, false)]));
|
||||
let table = db.create_empty_table("t", schema).execute().await?;
|
||||
|
||||
assert!(storage_format_version(&table).await < LanceFileVersion::V2_2);
|
||||
assert!(!supports_blob_v2_storage(
|
||||
storage_format_version(&table).await
|
||||
));
|
||||
assert!(!uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -171,7 +186,9 @@ async fn creating_with_blob_data_bumps_format() -> Result<()> {
|
||||
.unwrap();
|
||||
let table = db.create_table("t", batch).execute().await?;
|
||||
|
||||
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
|
||||
assert!(supports_blob_v2_storage(
|
||||
storage_format_version(&table).await
|
||||
));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
assert_eq!(table.count_rows(None).await?, 1);
|
||||
Ok(())
|
||||
@@ -281,7 +298,9 @@ async fn connection_level_stable_row_id_setting_wins_over_blob_default() -> Resu
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
|
||||
assert!(supports_blob_v2_storage(
|
||||
storage_format_version(&table).await
|
||||
));
|
||||
assert!(!uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -297,7 +316,9 @@ async fn namespace_create_applies_blob_defaults() -> Result<()> {
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
|
||||
assert!(supports_blob_v2_storage(
|
||||
storage_format_version(&table).await
|
||||
));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -474,7 +495,9 @@ async fn fetch_blobs_round_trips_nested_blob_column() -> Result<()> {
|
||||
let batch = RecordBatch::try_new(schema, vec![Arc::new(info_array) as ArrayRef]).unwrap();
|
||||
let table = db.create_table("t", batch).execute().await?;
|
||||
|
||||
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
|
||||
assert!(supports_blob_v2_storage(
|
||||
storage_format_version(&table).await
|
||||
));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
|
||||
let ids = collect_row_ids(&table).await?;
|
||||
@@ -1305,7 +1328,7 @@ async fn optimize_preserves_blob_v2_null_and_empty_distinction() -> Result<()> {
|
||||
.await?;
|
||||
table.add(null_empty_input_batch()).execute().await?;
|
||||
assert!(
|
||||
storage_format_version(&table).await >= LanceFileVersion::V2_2,
|
||||
supports_blob_v2_storage(storage_format_version(&table).await),
|
||||
"blob v2 columns require storage >= 2.2"
|
||||
);
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,837 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Contract tests for FunctionDefinition registration input (FF-007 / B1c).
|
||||
//!
|
||||
//! These tests pin the intended public surface under [`lancedb::function`] for
|
||||
//! Python definition transport only. They intentionally fail to compile until
|
||||
//! that API exists.
|
||||
//!
|
||||
//! Rejection cases are judged by `Result` structure (`is_err` / `is_ok`), never
|
||||
//! by diagnostic message substrings.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use arrow_schema::DataType;
|
||||
use lancedb::Result;
|
||||
use lancedb::function::{
|
||||
Function, FunctionCapability, FunctionDefinition, FunctionId, FunctionOutput,
|
||||
FunctionParameter, FunctionSignature, PythonFunctionDefinition,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
fn sample_signature() -> Result<FunctionSignature> {
|
||||
FunctionSignature::try_new(
|
||||
vec![
|
||||
FunctionParameter::new("text", DataType::Utf8),
|
||||
FunctionParameter::new("limit", DataType::Int32),
|
||||
],
|
||||
FunctionOutput::new(DataType::Utf8, true),
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_source() -> &'static str {
|
||||
"def normalize(text, limit):\n return text[:limit]\n"
|
||||
}
|
||||
|
||||
fn sample_python_definition() -> Result<PythonFunctionDefinition> {
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"normalize",
|
||||
sample_source(),
|
||||
"3.12",
|
||||
vec!["Unidecode==1.3.8".to_string()],
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_capabilities() -> Result<Vec<FunctionCapability>> {
|
||||
Ok(vec![
|
||||
FunctionCapability::try_network("https://api.example.com")?,
|
||||
FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?,
|
||||
])
|
||||
}
|
||||
|
||||
fn sample_definition() -> Result<FunctionDefinition> {
|
||||
FunctionDefinition::try_new(
|
||||
sample_signature()?,
|
||||
sample_python_definition()?,
|
||||
sample_capabilities()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn assert_json_object_keys_exact(value: &Value, expected: &[&str]) {
|
||||
let object = value
|
||||
.as_object()
|
||||
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
|
||||
let keys: BTreeSet<&str> = object.keys().map(|k| k.as_str()).collect();
|
||||
let expected: BTreeSet<&str> = expected.iter().copied().collect();
|
||||
assert_eq!(
|
||||
keys, expected,
|
||||
"JSON object key set must match exactly (iteration order is not a contract); got {keys:?}, expected {expected:?} in {value}"
|
||||
);
|
||||
}
|
||||
|
||||
fn assert_json_object_keys_subset(value: &Value, allowed: &[&str]) {
|
||||
let object = value
|
||||
.as_object()
|
||||
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
|
||||
for key in object.keys() {
|
||||
assert!(
|
||||
allowed.contains(&key.as_str()),
|
||||
"unexpected JSON key `{key}` in {value}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Identity / lineage / artifact / runtime fields that must not appear as public
|
||||
/// object keys on FunctionDefinition wire. Parameter object key `name` is not
|
||||
/// listed here: FunctionSignature legitimately uses it under `signature.parameters`.
|
||||
const FORBIDDEN_DEFINITION_KEYS: &[&str] = &[
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionId",
|
||||
"catalog",
|
||||
"catalog_name",
|
||||
"version",
|
||||
"function_version",
|
||||
"FunctionVersion",
|
||||
"lineage",
|
||||
"user_version",
|
||||
"idempotency_key",
|
||||
"digest",
|
||||
"artifact",
|
||||
"artifact_digest",
|
||||
"storage",
|
||||
"storage_location",
|
||||
"location",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"nullPolicy",
|
||||
"timestamp",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"worker",
|
||||
"scheduler",
|
||||
"attempt",
|
||||
"attempt_id",
|
||||
"replica",
|
||||
"placement",
|
||||
];
|
||||
|
||||
fn assert_object_keys_not_forbidden(value: &Value, context: &str) {
|
||||
let object = value.as_object().unwrap_or_else(|| {
|
||||
panic!("expected JSON object at {context}, got {value}");
|
||||
});
|
||||
for key in object.keys() {
|
||||
assert!(
|
||||
!FORBIDDEN_DEFINITION_KEYS.contains(&key.as_str()),
|
||||
"FunctionDefinition wire must not contain forbidden key `{key}` at {context}: {value}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_forbidden_definition_keys_absent(value: &Value) {
|
||||
// Top-level definition key set (order-independent).
|
||||
assert_json_object_keys_exact(
|
||||
value,
|
||||
&[
|
||||
"format_version",
|
||||
"signature",
|
||||
"implementation",
|
||||
"capabilities",
|
||||
],
|
||||
);
|
||||
assert_object_keys_not_forbidden(value, "definition");
|
||||
// Catalog / function identity name is absent at definition root; parameter
|
||||
// `name` is allowed only under signature.parameters.
|
||||
assert!(
|
||||
value.get("name").is_none(),
|
||||
"top-level FunctionDefinition wire must not contain catalog/function identity key `name`: {value}"
|
||||
);
|
||||
assert!(
|
||||
value.get("catalog_name").is_none(),
|
||||
"top-level FunctionDefinition wire must not contain `catalog_name`: {value}"
|
||||
);
|
||||
|
||||
let implementation = value.get("implementation").expect("implementation object");
|
||||
assert_json_object_keys_exact(
|
||||
implementation,
|
||||
&["kind", "module", "callable", "source", "python", "packages"],
|
||||
);
|
||||
assert_object_keys_not_forbidden(implementation, "implementation");
|
||||
assert!(
|
||||
implementation.get("name").is_none(),
|
||||
"implementation must not contain catalog/function identity key `name`: {implementation}"
|
||||
);
|
||||
assert!(
|
||||
implementation.get("catalog_name").is_none(),
|
||||
"implementation must not contain `catalog_name`: {implementation}"
|
||||
);
|
||||
|
||||
let capabilities = value
|
||||
.get("capabilities")
|
||||
.and_then(Value::as_array)
|
||||
.expect("capabilities array");
|
||||
for (idx, capability) in capabilities.iter().enumerate() {
|
||||
let kind = capability
|
||||
.get("kind")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_else(|| panic!("capabilities[{idx}] missing kind"));
|
||||
match kind {
|
||||
"network" => assert_json_object_keys_exact(capability, &["kind", "origin"]),
|
||||
"secret" => assert_json_object_keys_exact(
|
||||
capability,
|
||||
&["kind", "reference", "environment_variable"],
|
||||
),
|
||||
other => panic!("unexpected capability kind `{other}` in contract fixture"),
|
||||
}
|
||||
let context = format!("capabilities[{idx}]");
|
||||
assert_object_keys_not_forbidden(capability, &context);
|
||||
assert!(
|
||||
capability.get("name").is_none(),
|
||||
"{context} must not contain catalog/function identity key `name`: {capability}"
|
||||
);
|
||||
assert!(
|
||||
capability.get("catalog_name").is_none(),
|
||||
"{context} must not contain `catalog_name`: {capability}"
|
||||
);
|
||||
}
|
||||
|
||||
// Signature may carry parameter objects with key `name`. Still reject
|
||||
// identity/lineage/runtime keys and function-identity `name`/`catalog_name`
|
||||
// on the signature and output objects themselves.
|
||||
let signature = value.get("signature").expect("signature object");
|
||||
assert_object_keys_not_forbidden(signature, "signature");
|
||||
assert!(
|
||||
signature.get("name").is_none(),
|
||||
"signature object must not contain catalog/function identity key `name`: {signature}"
|
||||
);
|
||||
assert!(
|
||||
signature.get("catalog_name").is_none(),
|
||||
"signature object must not contain `catalog_name`: {signature}"
|
||||
);
|
||||
if let Some(parameters) = signature.get("parameters").and_then(Value::as_array) {
|
||||
for (idx, parameter) in parameters.iter().enumerate() {
|
||||
let context = format!("signature.parameters[{idx}]");
|
||||
assert_object_keys_not_forbidden(parameter, &context);
|
||||
assert!(
|
||||
parameter.get("catalog_name").is_none(),
|
||||
"{context} must not contain `catalog_name`: {parameter}"
|
||||
);
|
||||
// `name` is intentionally allowed on parameter objects.
|
||||
assert!(
|
||||
parameter.get("name").is_some(),
|
||||
"{context} must include parameter `name`"
|
||||
);
|
||||
}
|
||||
}
|
||||
if let Some(output) = signature.get("output") {
|
||||
assert_object_keys_not_forbidden(output, "signature.output");
|
||||
assert!(
|
||||
output.get("name").is_none(),
|
||||
"signature.output must not contain catalog/function identity key `name`: {output}"
|
||||
);
|
||||
assert!(
|
||||
output.get("catalog_name").is_none(),
|
||||
"signature.output must not contain `catalog_name`: {output}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn definition_json_round_trip_pins_exact_wire_shape_and_order() -> Result<()> {
|
||||
let definition = sample_definition()?;
|
||||
|
||||
assert_eq!(definition.signature().parameters().len(), 2);
|
||||
assert_eq!(definition.signature().parameters()[0].name(), "text");
|
||||
assert_eq!(
|
||||
definition.signature().parameters()[0].data_type(),
|
||||
&DataType::Utf8
|
||||
);
|
||||
assert_eq!(definition.signature().parameters()[1].name(), "limit");
|
||||
assert_eq!(
|
||||
definition.signature().parameters()[1].data_type(),
|
||||
&DataType::Int32
|
||||
);
|
||||
assert_eq!(definition.signature().output().data_type(), &DataType::Utf8);
|
||||
assert!(definition.signature().output().nullable());
|
||||
|
||||
let python = definition.python_definition();
|
||||
assert_eq!(python.module(), "normalize_mod");
|
||||
assert_eq!(python.callable(), "normalize");
|
||||
assert_eq!(python.source(), sample_source());
|
||||
assert_eq!(python.python(), "3.12");
|
||||
assert_eq!(python.packages(), &["Unidecode==1.3.8".to_string()]);
|
||||
|
||||
let capabilities = definition.capabilities();
|
||||
assert_eq!(capabilities.len(), 2);
|
||||
assert_eq!(capabilities[0].origin(), Some("https://api.example.com"));
|
||||
assert_eq!(capabilities[0].reference(), None);
|
||||
assert_eq!(capabilities[0].environment_variable(), None);
|
||||
assert_eq!(capabilities[1].reference(), Some("secret://team/api-token"));
|
||||
assert_eq!(capabilities[1].environment_variable(), Some("API_TOKEN"));
|
||||
assert_eq!(capabilities[1].origin(), None);
|
||||
|
||||
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
|
||||
assert_json_object_keys_exact(
|
||||
&json,
|
||||
&[
|
||||
"format_version",
|
||||
"signature",
|
||||
"implementation",
|
||||
"capabilities",
|
||||
],
|
||||
);
|
||||
assert_eq!(json["format_version"], 1);
|
||||
|
||||
let signature = json
|
||||
.get("signature")
|
||||
.and_then(Value::as_object)
|
||||
.expect("signature object");
|
||||
assert_json_object_keys_subset(&Value::Object(signature.clone()), &["parameters", "output"]);
|
||||
let parameters = signature
|
||||
.get("parameters")
|
||||
.and_then(Value::as_array)
|
||||
.expect("parameters array");
|
||||
assert_eq!(parameters.len(), 2);
|
||||
assert_eq!(parameters[0]["name"], Value::String("text".into()));
|
||||
assert_eq!(parameters[1]["name"], Value::String("limit".into()));
|
||||
for parameter in parameters {
|
||||
assert_json_object_keys_subset(parameter, &["name", "data_type_ipc"]);
|
||||
assert!(
|
||||
parameter
|
||||
.get("data_type_ipc")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|s| !s.is_empty()),
|
||||
"parameter data_type_ipc must be non-empty base64"
|
||||
);
|
||||
}
|
||||
let output = signature
|
||||
.get("output")
|
||||
.and_then(Value::as_object)
|
||||
.expect("output object");
|
||||
assert_json_object_keys_subset(
|
||||
&Value::Object(output.clone()),
|
||||
&["data_type_ipc", "nullable"],
|
||||
);
|
||||
assert_eq!(output.get("nullable"), Some(&Value::Bool(true)));
|
||||
|
||||
let implementation = json.get("implementation").expect("implementation object");
|
||||
assert_json_object_keys_exact(
|
||||
implementation,
|
||||
&["kind", "module", "callable", "source", "python", "packages"],
|
||||
);
|
||||
assert_eq!(implementation["kind"], Value::String("python".into()));
|
||||
assert_eq!(
|
||||
implementation["module"],
|
||||
Value::String("normalize_mod".into())
|
||||
);
|
||||
assert_eq!(
|
||||
implementation["callable"],
|
||||
Value::String("normalize".into())
|
||||
);
|
||||
assert_eq!(
|
||||
implementation["source"],
|
||||
Value::String(sample_source().into())
|
||||
);
|
||||
assert_eq!(implementation["python"], Value::String("3.12".into()));
|
||||
assert_eq!(
|
||||
implementation["packages"],
|
||||
Value::Array(vec![Value::String("Unidecode==1.3.8".into())])
|
||||
);
|
||||
|
||||
let capabilities_json = json
|
||||
.get("capabilities")
|
||||
.and_then(Value::as_array)
|
||||
.expect("capabilities array");
|
||||
assert_eq!(capabilities_json.len(), 2);
|
||||
assert_json_object_keys_exact(&capabilities_json[0], &["kind", "origin"]);
|
||||
assert_eq!(
|
||||
capabilities_json[0]["kind"],
|
||||
Value::String("network".into())
|
||||
);
|
||||
assert_eq!(
|
||||
capabilities_json[0]["origin"],
|
||||
Value::String("https://api.example.com".into())
|
||||
);
|
||||
assert_json_object_keys_exact(
|
||||
&capabilities_json[1],
|
||||
&["kind", "reference", "environment_variable"],
|
||||
);
|
||||
assert_eq!(capabilities_json[1]["kind"], Value::String("secret".into()));
|
||||
assert_eq!(
|
||||
capabilities_json[1]["reference"],
|
||||
Value::String("secret://team/api-token".into())
|
||||
);
|
||||
assert_eq!(
|
||||
capabilities_json[1]["environment_variable"],
|
||||
Value::String("API_TOKEN".into())
|
||||
);
|
||||
|
||||
// Same ordered signature IPC representation as Function handle transport.
|
||||
let function = Function::new(FunctionId::try_new("fn.wire.compare")?, sample_signature()?);
|
||||
let function_json = serde_json::to_value(&function).expect("serialize Function");
|
||||
assert_eq!(json["signature"], function_json["signature"]);
|
||||
|
||||
let restored: FunctionDefinition =
|
||||
serde_json::from_value(json.clone()).expect("deserialize FunctionDefinition");
|
||||
assert_eq!(
|
||||
restored.signature().parameters()[0].name(),
|
||||
definition.signature().parameters()[0].name()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.signature().parameters()[0].data_type(),
|
||||
definition.signature().parameters()[0].data_type()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.signature().parameters()[1].name(),
|
||||
definition.signature().parameters()[1].name()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.signature().parameters()[1].data_type(),
|
||||
definition.signature().parameters()[1].data_type()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.signature().output().data_type(),
|
||||
definition.signature().output().data_type()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.signature().output().nullable(),
|
||||
definition.signature().output().nullable()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.python_definition().module(),
|
||||
definition.python_definition().module()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.python_definition().callable(),
|
||||
definition.python_definition().callable()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.python_definition().source(),
|
||||
definition.python_definition().source()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.python_definition().python(),
|
||||
definition.python_definition().python()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.python_definition().packages(),
|
||||
definition.python_definition().packages()
|
||||
);
|
||||
assert_eq!(restored.capabilities().len(), 2);
|
||||
assert_eq!(
|
||||
restored.capabilities()[0].origin(),
|
||||
definition.capabilities()[0].origin()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.capabilities()[1].reference(),
|
||||
definition.capabilities()[1].reference()
|
||||
);
|
||||
assert_eq!(
|
||||
restored.capabilities()[1].environment_variable(),
|
||||
definition.capabilities()[1].environment_variable()
|
||||
);
|
||||
|
||||
// Package and capability order are part of the structural wire.
|
||||
let multi_pkg = FunctionDefinition::try_new(
|
||||
sample_signature()?,
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"normalize",
|
||||
sample_source(),
|
||||
"3.12",
|
||||
vec![
|
||||
"Unidecode==1.3.8".to_string(),
|
||||
"requests==2.32.3".to_string(),
|
||||
],
|
||||
)?,
|
||||
vec![
|
||||
FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?,
|
||||
FunctionCapability::try_network("https://api.example.com")?,
|
||||
FunctionCapability::try_network("https://other.example.com")?,
|
||||
],
|
||||
)?;
|
||||
let multi_json = serde_json::to_value(&multi_pkg).expect("serialize multi-order definition");
|
||||
assert_eq!(
|
||||
multi_json["implementation"]["packages"],
|
||||
Value::Array(vec![
|
||||
Value::String("Unidecode==1.3.8".into()),
|
||||
Value::String("requests==2.32.3".into()),
|
||||
])
|
||||
);
|
||||
assert_eq!(
|
||||
multi_json["capabilities"][0]["kind"],
|
||||
Value::String("secret".into())
|
||||
);
|
||||
assert_eq!(
|
||||
multi_json["capabilities"][1]["origin"],
|
||||
Value::String("https://api.example.com".into())
|
||||
);
|
||||
assert_eq!(
|
||||
multi_json["capabilities"][2]["origin"],
|
||||
Value::String("https://other.example.com".into())
|
||||
);
|
||||
let multi_restored: FunctionDefinition =
|
||||
serde_json::from_value(multi_json.clone()).expect("deserialize multi-order definition");
|
||||
assert_eq!(
|
||||
multi_restored.python_definition().packages(),
|
||||
&[
|
||||
"Unidecode==1.3.8".to_string(),
|
||||
"requests==2.32.3".to_string()
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
multi_restored.capabilities()[0].reference(),
|
||||
Some("secret://team/api-token")
|
||||
);
|
||||
assert_eq!(
|
||||
multi_restored.capabilities()[1].origin(),
|
||||
Some("https://api.example.com")
|
||||
);
|
||||
assert_eq!(
|
||||
multi_restored.capabilities()[2].origin(),
|
||||
Some("https://other.example.com")
|
||||
);
|
||||
|
||||
// Byte-for-byte repeated serde_json encoding for the same value.
|
||||
let encoded_a = serde_json::to_string(&definition).expect("encode a");
|
||||
let encoded_b = serde_json::to_string(&definition).expect("encode b");
|
||||
assert_eq!(encoded_a, encoded_b);
|
||||
assert_eq!(
|
||||
serde_json::to_value(&restored).expect("re-serialize restored"),
|
||||
json
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_string(&multi_restored).expect("re-encode multi"),
|
||||
serde_json::to_string(&multi_pkg).expect("encode multi")
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn definition_wire_excludes_identity_lineage_artifact_and_runtime_fields() -> Result<()> {
|
||||
let definition = sample_definition()?;
|
||||
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
|
||||
// Forbidden public fields are object keys at structural levels only.
|
||||
// Do not substring-scan encoded JSON: user source/reference/package text
|
||||
// may legitimately contain those tokens.
|
||||
assert_forbidden_definition_keys_absent(&json);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn definition_decode_fails_closed_for_unknown_version_fields_and_kinds() -> Result<()> {
|
||||
let definition = sample_definition()?;
|
||||
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
|
||||
|
||||
let mut unknown_version = json.clone();
|
||||
unknown_version["format_version"] = Value::from(2);
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(unknown_version).is_err(),
|
||||
"format_version other than 1 must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_outer = json.clone();
|
||||
unknown_outer
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("unexpected_field".into(), Value::Bool(true));
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(unknown_outer).is_err(),
|
||||
"unknown outer field must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_implementation_field = json.clone();
|
||||
unknown_implementation_field["implementation"]
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("entrypoint".into(), Value::String("main".into()));
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(unknown_implementation_field).is_err(),
|
||||
"unknown nested implementation field must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_capability_field = json.clone();
|
||||
unknown_capability_field["capabilities"][0]
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("headers".into(), Value::Object(Default::default()));
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(unknown_capability_field).is_err(),
|
||||
"unknown nested capability field must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_implementation_kind = json.clone();
|
||||
unknown_implementation_kind["implementation"]["kind"] = Value::String("builtin".into());
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(unknown_implementation_kind).is_err(),
|
||||
"unknown implementation kind must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_capability_kind = json.clone();
|
||||
unknown_capability_kind["capabilities"][0]["kind"] = Value::String("filesystem".into());
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(unknown_capability_kind).is_err(),
|
||||
"unknown capability kind must fail closed"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn constructors_and_decode_reject_empty_fields_and_duplicate_packages() -> Result<()> {
|
||||
let signature = sample_signature()?;
|
||||
let packages = vec!["Unidecode==1.3.8".to_string()];
|
||||
let capabilities = sample_capabilities()?;
|
||||
|
||||
assert!(
|
||||
PythonFunctionDefinition::try_new(
|
||||
"",
|
||||
"normalize",
|
||||
sample_source(),
|
||||
"3.12",
|
||||
packages.clone()
|
||||
)
|
||||
.is_err(),
|
||||
"empty module must be rejected"
|
||||
);
|
||||
assert!(
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"",
|
||||
sample_source(),
|
||||
"3.12",
|
||||
packages.clone()
|
||||
)
|
||||
.is_err(),
|
||||
"empty callable must be rejected"
|
||||
);
|
||||
assert!(
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"normalize",
|
||||
"",
|
||||
"3.12",
|
||||
packages.clone()
|
||||
)
|
||||
.is_err(),
|
||||
"empty source must be rejected"
|
||||
);
|
||||
assert!(
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"normalize",
|
||||
sample_source(),
|
||||
"",
|
||||
packages.clone()
|
||||
)
|
||||
.is_err(),
|
||||
"empty python runtime request must be rejected"
|
||||
);
|
||||
assert!(
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"normalize",
|
||||
sample_source(),
|
||||
"3.12",
|
||||
vec!["".to_string()],
|
||||
)
|
||||
.is_err(),
|
||||
"empty package requirement must be rejected"
|
||||
);
|
||||
assert!(
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"normalize",
|
||||
sample_source(),
|
||||
"3.12",
|
||||
vec![
|
||||
"Unidecode==1.3.8".to_string(),
|
||||
"Unidecode==1.3.8".to_string(),
|
||||
],
|
||||
)
|
||||
.is_err(),
|
||||
"duplicate package requirements must be rejected"
|
||||
);
|
||||
|
||||
assert!(
|
||||
FunctionCapability::try_network("").is_err(),
|
||||
"empty network origin must be rejected"
|
||||
);
|
||||
assert!(
|
||||
FunctionCapability::try_secret("", "API_TOKEN").is_err(),
|
||||
"empty secret reference must be rejected"
|
||||
);
|
||||
assert!(
|
||||
FunctionCapability::try_secret("secret://team/api-token", "").is_err(),
|
||||
"empty secret environment variable must be rejected"
|
||||
);
|
||||
|
||||
// Decode path must enforce the same emptiness / uniqueness rules.
|
||||
let definition = FunctionDefinition::try_new(
|
||||
signature.clone(),
|
||||
sample_python_definition()?,
|
||||
capabilities.clone(),
|
||||
)?;
|
||||
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
|
||||
|
||||
for (pointer, empty) in [
|
||||
("/implementation/module", ""),
|
||||
("/implementation/callable", ""),
|
||||
("/implementation/source", ""),
|
||||
("/implementation/python", ""),
|
||||
("/capabilities/0/origin", ""),
|
||||
("/capabilities/1/reference", ""),
|
||||
("/capabilities/1/environment_variable", ""),
|
||||
] {
|
||||
let mut invalid = json.clone();
|
||||
let target = invalid
|
||||
.pointer_mut(pointer)
|
||||
.unwrap_or_else(|| panic!("missing pointer {pointer}"));
|
||||
*target = Value::String(empty.into());
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(invalid).is_err(),
|
||||
"decode must reject empty value at {pointer}"
|
||||
);
|
||||
}
|
||||
|
||||
let mut empty_package = json.clone();
|
||||
empty_package["implementation"]["packages"] = Value::Array(vec![Value::String("".into())]);
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(empty_package).is_err(),
|
||||
"decode must reject empty package requirement"
|
||||
);
|
||||
|
||||
let mut duplicate_packages = json.clone();
|
||||
duplicate_packages["implementation"]["packages"] = Value::Array(vec![
|
||||
Value::String("Unidecode==1.3.8".into()),
|
||||
Value::String("Unidecode==1.3.8".into()),
|
||||
]);
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(duplicate_packages).is_err(),
|
||||
"decode must reject duplicate package requirements"
|
||||
);
|
||||
|
||||
// Keep the constructor path for FunctionDefinition itself structurally valid
|
||||
// when children are valid; emptiness is owned by child constructors above.
|
||||
assert!(
|
||||
FunctionDefinition::try_new(signature, sample_python_definition()?, capabilities).is_ok()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn secret_capability_wire_rejects_plaintext_value_fields() -> Result<()> {
|
||||
let definition = sample_definition()?;
|
||||
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
|
||||
let secret = &json["capabilities"][1];
|
||||
assert_json_object_keys_exact(secret, &["kind", "reference", "environment_variable"]);
|
||||
assert!(secret.get("value").is_none());
|
||||
assert!(secret.get("plaintext_secret").is_none());
|
||||
|
||||
let mut with_value = json.clone();
|
||||
with_value["capabilities"][1]
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("value".into(), Value::String("super-secret".into()));
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(with_value).is_err(),
|
||||
"secret capability must reject `value`"
|
||||
);
|
||||
|
||||
let mut with_plaintext = json.clone();
|
||||
with_plaintext["capabilities"][1]
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert(
|
||||
"plaintext_secret".into(),
|
||||
Value::String("super-secret".into()),
|
||||
);
|
||||
assert!(
|
||||
serde_json::from_value::<FunctionDefinition>(with_plaintext).is_err(),
|
||||
"secret capability must reject `plaintext_secret`"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn debug_redacts_source_and_secret_reference_while_getters_remain_exact() -> Result<()> {
|
||||
let python = sample_python_definition()?;
|
||||
let source = sample_source();
|
||||
assert_eq!(python.source(), source);
|
||||
let python_debug = format!("{python:?}");
|
||||
assert!(
|
||||
!python_debug.contains(source),
|
||||
"PythonFunctionDefinition Debug must not contain source body: {python_debug}"
|
||||
);
|
||||
assert!(
|
||||
!python_debug.contains("return text[:limit]"),
|
||||
"PythonFunctionDefinition Debug must not leak source fragments: {python_debug}"
|
||||
);
|
||||
|
||||
let secret = FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?;
|
||||
assert_eq!(secret.reference(), Some("secret://team/api-token"));
|
||||
assert_eq!(secret.environment_variable(), Some("API_TOKEN"));
|
||||
let secret_debug = format!("{secret:?}");
|
||||
assert!(
|
||||
!secret_debug.contains("secret://team/api-token"),
|
||||
"FunctionCapability secret Debug must not contain reference: {secret_debug}"
|
||||
);
|
||||
|
||||
let definition = sample_definition()?;
|
||||
assert_eq!(definition.python_definition().source(), source);
|
||||
assert_eq!(
|
||||
definition.capabilities()[1].reference(),
|
||||
Some("secret://team/api-token")
|
||||
);
|
||||
let definition_debug = format!("{definition:?}");
|
||||
assert!(
|
||||
!definition_debug.contains(source),
|
||||
"FunctionDefinition Debug must not contain source body: {definition_debug}"
|
||||
);
|
||||
assert!(
|
||||
!definition_debug.contains("secret://team/api-token"),
|
||||
"FunctionDefinition Debug must not contain secret reference: {definition_debug}"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn definition_has_no_identity_before_registration_and_is_not_a_function_handle() -> Result<()> {
|
||||
let definition = sample_definition()?;
|
||||
let json = serde_json::to_value(&definition).expect("serialize FunctionDefinition");
|
||||
|
||||
assert!(json.get("id").is_none());
|
||||
assert!(json.get("function_id").is_none());
|
||||
assert_json_object_keys_exact(
|
||||
&json,
|
||||
&[
|
||||
"format_version",
|
||||
"signature",
|
||||
"implementation",
|
||||
"capabilities",
|
||||
],
|
||||
);
|
||||
|
||||
// Identity exists only on the immutable Function handle after registration.
|
||||
// Definition remains a separate authoring value and does not borrow or mint an ID.
|
||||
let registered = Function::new(
|
||||
FunctionId::try_new("fn.published.after.registration")?,
|
||||
definition.signature().clone(),
|
||||
);
|
||||
assert_eq!(registered.id().as_str(), "fn.published.after.registration");
|
||||
assert_eq!(
|
||||
registered.signature().parameters().len(),
|
||||
definition.signature().parameters().len()
|
||||
);
|
||||
|
||||
let definition_again = serde_json::to_value(&definition).expect("re-serialize definition");
|
||||
assert!(definition_again.get("id").is_none());
|
||||
assert!(definition_again.get("function_id").is_none());
|
||||
assert_ne!(
|
||||
serde_json::to_value(®istered).expect("serialize Function"),
|
||||
definition_again,
|
||||
"Function handle wire must remain distinct from FunctionDefinition wire"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Public contract tests for first-class Function error codes (FF-006).
|
||||
//!
|
||||
//! These tests pin the stable `FunctionErrorCode` wire strings, direct
|
||||
//! `Error::Function` projection, and optional `JobFailure.error_code`.
|
||||
//! They intentionally fail to compile until that public API exists.
|
||||
//!
|
||||
//! Categories are judged only by structural enum matching / equality, never
|
||||
//! by parsing diagnostic message text.
|
||||
|
||||
use lancedb::error::FunctionErrorCode;
|
||||
use lancedb::{Error, JobFailure};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
/// Exact stable wire strings for the eight known Function error categories.
|
||||
const KNOWN_WIRE_CODES: &[(&str, FunctionErrorCode)] = &[
|
||||
(
|
||||
"definition_validation_failure",
|
||||
FunctionErrorCode::DefinitionValidationFailure,
|
||||
),
|
||||
(
|
||||
"name_or_function_not_found",
|
||||
FunctionErrorCode::NameOrFunctionNotFound,
|
||||
),
|
||||
("name_conflict", FunctionErrorCode::NameConflict),
|
||||
(
|
||||
"unsupported_runtime_or_capability",
|
||||
FunctionErrorCode::UnsupportedRuntimeOrCapability,
|
||||
),
|
||||
("revoked_function", FunctionErrorCode::RevokedFunction),
|
||||
(
|
||||
"udf_execution_failure",
|
||||
FunctionErrorCode::UdfExecutionFailure,
|
||||
),
|
||||
(
|
||||
"generated_column_incomplete",
|
||||
FunctionErrorCode::GeneratedColumnIncomplete,
|
||||
),
|
||||
(
|
||||
"stale_or_conflicting_input",
|
||||
FunctionErrorCode::StaleOrConflictingInput,
|
||||
),
|
||||
];
|
||||
|
||||
fn assert_known_variant(code: &FunctionErrorCode, expected: &FunctionErrorCode) {
|
||||
assert_eq!(
|
||||
code, expected,
|
||||
"FunctionErrorCode must match structurally; got {code:?}, expected {expected:?}"
|
||||
);
|
||||
assert!(
|
||||
!matches!(code, FunctionErrorCode::Unrecognized(_)),
|
||||
"known wire string must not deserialize as Unrecognized: {code:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn function_error_code_known_variants_use_exact_stable_json_strings() {
|
||||
for (wire, expected) in KNOWN_WIRE_CODES {
|
||||
let encoded = serde_json::to_value(expected).expect("serialize FunctionErrorCode");
|
||||
assert_eq!(
|
||||
encoded,
|
||||
Value::String((*wire).to_string()),
|
||||
"stable JSON string for {expected:?}"
|
||||
);
|
||||
|
||||
let decoded: FunctionErrorCode = serde_json::from_value(Value::String((*wire).to_string()))
|
||||
.unwrap_or_else(|e| panic!("deserialize `{wire}`: {e}"));
|
||||
assert_known_variant(&decoded, expected);
|
||||
|
||||
let round_trip = serde_json::to_value(&decoded).expect("re-serialize");
|
||||
assert_eq!(round_trip, Value::String((*wire).to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unrecognized_error_code_preserves_exact_string_and_does_not_become_known() {
|
||||
let raw = "enterprise_future_category_xyz";
|
||||
let decoded: FunctionErrorCode = serde_json::from_value(json!(raw))
|
||||
.unwrap_or_else(|e| panic!("unknown code must deserialize, not fail: {e}"));
|
||||
|
||||
match &decoded {
|
||||
FunctionErrorCode::Unrecognized(preserved) => {
|
||||
assert_eq!(preserved, raw, "unknown code must be preserved verbatim");
|
||||
}
|
||||
other => panic!("expected FunctionErrorCode::Unrecognized, got {other:?}"),
|
||||
}
|
||||
|
||||
for (_, known) in KNOWN_WIRE_CODES {
|
||||
assert_ne!(
|
||||
&decoded, known,
|
||||
"unrecognized code must not equal known variant {known:?}"
|
||||
);
|
||||
}
|
||||
|
||||
let encoded = serde_json::to_value(&decoded).expect("serialize Unrecognized");
|
||||
assert_eq!(encoded, json!(raw));
|
||||
|
||||
let again: FunctionErrorCode =
|
||||
serde_json::from_value(encoded).expect("Unrecognized must round-trip");
|
||||
match again {
|
||||
FunctionErrorCode::Unrecognized(preserved) => assert_eq!(preserved, raw),
|
||||
other => panic!("round-trip must stay Unrecognized, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_function_carries_code_plus_diagnostic_message() {
|
||||
let err = Error::Function {
|
||||
code: FunctionErrorCode::NameConflict,
|
||||
message: "sanitized diagnostic only".to_string(),
|
||||
};
|
||||
|
||||
match err {
|
||||
Error::Function { code, message } => {
|
||||
assert_known_variant(&code, &FunctionErrorCode::NameConflict);
|
||||
assert_eq!(message, "sanitized diagnostic only");
|
||||
}
|
||||
other => panic!("expected Error::Function, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_function_category_is_the_code_field_not_the_message() {
|
||||
// Message text deliberately names a different category; structural code wins.
|
||||
let err = Error::Function {
|
||||
code: FunctionErrorCode::GeneratedColumnIncomplete,
|
||||
message: "looks like udf_execution_failure to a string parser".to_string(),
|
||||
};
|
||||
|
||||
match err {
|
||||
Error::Function { code, .. } => {
|
||||
assert_known_variant(&code, &FunctionErrorCode::GeneratedColumnIncomplete);
|
||||
assert_ne!(code, FunctionErrorCode::UdfExecutionFailure);
|
||||
}
|
||||
other => panic!("expected Error::Function, got {other:?}"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn job_failure_has_optional_error_code() {
|
||||
let with_code = JobFailure {
|
||||
error_code: Some(FunctionErrorCode::RevokedFunction),
|
||||
phase: Some("execute".to_string()),
|
||||
message: Some("revoked".to_string()),
|
||||
retryable: Some(false),
|
||||
source: None,
|
||||
};
|
||||
match &with_code.error_code {
|
||||
Some(code) => assert_known_variant(code, &FunctionErrorCode::RevokedFunction),
|
||||
None => panic!("error_code must be present when set"),
|
||||
}
|
||||
|
||||
let without_code = JobFailure {
|
||||
phase: Some("execute".to_string()),
|
||||
message: Some("older backend failure without a category".to_string()),
|
||||
retryable: Some(true),
|
||||
..Default::default()
|
||||
};
|
||||
assert!(
|
||||
without_code.error_code.is_none(),
|
||||
"missing error_code must stay None; diagnostics must not invent a category"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn job_failure_diagnostics_do_not_overwrite_error_code() {
|
||||
let failure = JobFailure {
|
||||
error_code: Some(FunctionErrorCode::StaleOrConflictingInput),
|
||||
phase: Some("commit".to_string()),
|
||||
message: Some("definition_validation_failure in worker logs".to_string()),
|
||||
retryable: Some(true),
|
||||
source: None,
|
||||
};
|
||||
|
||||
match &failure.error_code {
|
||||
Some(code) => {
|
||||
assert_known_variant(code, &FunctionErrorCode::StaleOrConflictingInput);
|
||||
assert_ne!(code, &FunctionErrorCode::DefinitionValidationFailure);
|
||||
}
|
||||
None => panic!("explicit error_code must remain set"),
|
||||
}
|
||||
assert_eq!(failure.phase.as_deref(), Some("commit"));
|
||||
assert_eq!(failure.retryable, Some(true));
|
||||
}
|
||||
@@ -0,0 +1,366 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Contract tests for JobResult value/wire (FF-012).
|
||||
//!
|
||||
//! These tests pin the intended public non-resource [`lancedb::JobResult`]
|
||||
//! surface under [`lancedb::job`]. They intentionally fail to compile until
|
||||
//! that API exists.
|
||||
//!
|
||||
//! Scope is JobResult value and JSON wire only. Job::wait behavior, remote
|
||||
//! describe shape, missing-result handling, local outcome, Python, Node, and
|
||||
//! Sophon are out of scope.
|
||||
//!
|
||||
//! Rejection cases are judged by `Result` structure (`is_err` / `is_ok`) or
|
||||
//! serde decode failure, never by diagnostic message substrings.
|
||||
//! JSON map iteration order is not a contract; exact key sets are compared
|
||||
//! independently. Byte reproducibility means repeated encoding of the same
|
||||
//! in-memory value.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use arrow_schema::DataType;
|
||||
use lancedb::JobResult;
|
||||
use lancedb::Result;
|
||||
use lancedb::function::{
|
||||
Function, FunctionId, FunctionOutput, FunctionParameter, FunctionSignature,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
fn sample_function() -> Result<Function> {
|
||||
let id = FunctionId::try_new("fn.exact.job-result")?;
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![
|
||||
FunctionParameter::new("x", DataType::Int32),
|
||||
FunctionParameter::new("label", DataType::Utf8),
|
||||
],
|
||||
FunctionOutput::new(DataType::Int32, true),
|
||||
)?;
|
||||
Ok(Function::new(id, signature))
|
||||
}
|
||||
|
||||
fn assert_json_object_keys_exact(value: &Value, expected: &[&str]) {
|
||||
let object = value
|
||||
.as_object()
|
||||
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
|
||||
let keys: BTreeSet<&str> = object.keys().map(|k| k.as_str()).collect();
|
||||
let expected: BTreeSet<&str> = expected.iter().copied().collect();
|
||||
assert_eq!(
|
||||
keys, expected,
|
||||
"JSON object key set must match exactly (iteration order is not a contract); got {keys:?}, expected {expected:?} in {value}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Outer JobResult object keys that must not appear. Nested Function signature
|
||||
/// keys (including parameter `name` and Function `id`) are legitimate and are
|
||||
/// not scanned here. Opaque string contents are not recursively searched.
|
||||
const FORBIDDEN_OUTER_RESULT_KEYS: &[&str] = &[
|
||||
"name",
|
||||
"definition",
|
||||
"FunctionDefinition",
|
||||
"source",
|
||||
"runtime",
|
||||
"packages",
|
||||
"capability",
|
||||
"capabilities",
|
||||
"artifact",
|
||||
"artifact_digest",
|
||||
"digest",
|
||||
"storage",
|
||||
"storage_location",
|
||||
"location",
|
||||
"table",
|
||||
"table_name",
|
||||
"table_ref",
|
||||
"version",
|
||||
"function_version",
|
||||
"FunctionVersion",
|
||||
"user_version",
|
||||
"lineage",
|
||||
"id",
|
||||
"job_id",
|
||||
"jobId",
|
||||
"type",
|
||||
"job_type",
|
||||
"state",
|
||||
"status",
|
||||
"lifecycle",
|
||||
"failure",
|
||||
"attempt",
|
||||
"attempt_id",
|
||||
"timestamp",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"retry",
|
||||
"retry_key",
|
||||
"idempotency",
|
||||
"idempotency_key",
|
||||
"commit_token",
|
||||
"secret",
|
||||
"compatibility",
|
||||
"deterministic",
|
||||
"null_policy",
|
||||
"nullPolicy",
|
||||
];
|
||||
|
||||
fn assert_outer_forbidden_keys_absent(value: &Value) {
|
||||
let object = value
|
||||
.as_object()
|
||||
.unwrap_or_else(|| panic!("expected outer JobResult JSON object, got {value}"));
|
||||
for key in object.keys() {
|
||||
assert!(
|
||||
!FORBIDDEN_OUTER_RESULT_KEYS.contains(&key.as_str()),
|
||||
"outer JobResult wire must not contain forbidden key `{key}`: {value}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_function_handle_exact(actual: &Function, expected: &Function) {
|
||||
assert_eq!(actual.id().as_str(), expected.id().as_str());
|
||||
assert_eq!(
|
||||
actual.signature().parameters().len(),
|
||||
expected.signature().parameters().len()
|
||||
);
|
||||
for (actual_param, expected_param) in actual
|
||||
.signature()
|
||||
.parameters()
|
||||
.iter()
|
||||
.zip(expected.signature().parameters().iter())
|
||||
{
|
||||
assert_eq!(actual_param.name(), expected_param.name());
|
||||
assert_eq!(actual_param.data_type(), expected_param.data_type());
|
||||
}
|
||||
assert_eq!(
|
||||
actual.signature().output().data_type(),
|
||||
expected.signature().output().data_type()
|
||||
);
|
||||
assert_eq!(
|
||||
actual.signature().output().nullable(),
|
||||
expected.signature().output().nullable()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn none_and_function_exact_key_sets_helpers_round_trip_and_bytes() -> Result<()> {
|
||||
let none = JobResult::None;
|
||||
assert_eq!(none.format_version(), 1);
|
||||
assert!(none.function().is_none());
|
||||
assert!(matches!(none, JobResult::None));
|
||||
|
||||
let none_json = serde_json::to_value(&none).expect("serialize JobResult::None");
|
||||
assert_json_object_keys_exact(&none_json, &["format_version", "kind"]);
|
||||
assert_eq!(none_json["format_version"], 1);
|
||||
assert_eq!(none_json["kind"], Value::String("none".into()));
|
||||
assert!(none_json.get("function").is_none());
|
||||
|
||||
let none_restored: JobResult =
|
||||
serde_json::from_value(none_json.clone()).expect("deserialize JobResult::None");
|
||||
assert_eq!(none_restored.format_version(), 1);
|
||||
assert!(none_restored.function().is_none());
|
||||
assert!(matches!(none_restored, JobResult::None));
|
||||
assert_eq!(none_restored, none);
|
||||
assert_eq!(
|
||||
serde_json::to_value(&none_restored).expect("re-serialize None"),
|
||||
none_json
|
||||
);
|
||||
|
||||
let none_a = serde_json::to_string(&none).expect("encode None a");
|
||||
let none_b = serde_json::to_string(&none).expect("encode None b");
|
||||
assert_eq!(
|
||||
none_a, none_b,
|
||||
"repeated None encoding must be byte-identical"
|
||||
);
|
||||
|
||||
let function = sample_function()?;
|
||||
let expected_function_wire =
|
||||
serde_json::to_value(&function).expect("serialize nested Function");
|
||||
let function_result = JobResult::Function(function.clone());
|
||||
assert_eq!(function_result.format_version(), 1);
|
||||
assert!(matches!(function_result, JobResult::Function(_)));
|
||||
assert_function_handle_exact(
|
||||
function_result.function().expect("Function variant"),
|
||||
&function,
|
||||
);
|
||||
|
||||
let function_json =
|
||||
serde_json::to_value(&function_result).expect("serialize JobResult::Function");
|
||||
assert_json_object_keys_exact(&function_json, &["format_version", "kind", "function"]);
|
||||
assert_eq!(function_json["format_version"], 1);
|
||||
assert_eq!(function_json["kind"], Value::String("function".into()));
|
||||
assert_eq!(
|
||||
function_json["function"], expected_function_wire,
|
||||
"nested function must be the exact existing Function wire"
|
||||
);
|
||||
assert_json_object_keys_exact(
|
||||
&function_json["function"],
|
||||
&["format_version", "id", "signature"],
|
||||
);
|
||||
assert_eq!(
|
||||
function_json["function"]["id"],
|
||||
Value::String("fn.exact.job-result".into())
|
||||
);
|
||||
assert_eq!(function_json["function"]["format_version"], 1);
|
||||
|
||||
let function_restored: JobResult =
|
||||
serde_json::from_value(function_json.clone()).expect("deserialize JobResult::Function");
|
||||
assert_eq!(function_restored.format_version(), 1);
|
||||
assert_function_handle_exact(
|
||||
function_restored.function().expect("Function variant"),
|
||||
&function,
|
||||
);
|
||||
assert_eq!(
|
||||
function_restored
|
||||
.function()
|
||||
.expect("Function variant")
|
||||
.id()
|
||||
.as_str(),
|
||||
"fn.exact.job-result"
|
||||
);
|
||||
assert_eq!(function_restored, function_result);
|
||||
assert_eq!(
|
||||
serde_json::to_value(&function_restored).expect("re-serialize Function"),
|
||||
function_json
|
||||
);
|
||||
|
||||
let function_a = serde_json::to_string(&function_result).expect("encode Function a");
|
||||
let function_b = serde_json::to_string(&function_result).expect("encode Function b");
|
||||
assert_eq!(
|
||||
function_a, function_b,
|
||||
"repeated Function encoding must be byte-identical"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn function_and_into_function_accessors_for_both_variants() -> Result<()> {
|
||||
let none = JobResult::None;
|
||||
assert!(none.function().is_none());
|
||||
assert!(none.into_function().is_none());
|
||||
|
||||
let function = sample_function()?;
|
||||
let function_result = JobResult::Function(function.clone());
|
||||
assert_function_handle_exact(
|
||||
function_result.function().expect("borrowed Function"),
|
||||
&function,
|
||||
);
|
||||
let owned = function_result
|
||||
.into_function()
|
||||
.expect("owned Function from Function variant");
|
||||
assert_function_handle_exact(&owned, &function);
|
||||
assert_eq!(owned.id().as_str(), "fn.exact.job-result");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_kind_field_version_and_malformed_function_fail_closed() -> Result<()> {
|
||||
let none = JobResult::None;
|
||||
let none_json = serde_json::to_value(&none).expect("serialize None");
|
||||
|
||||
let mut unknown_kind = none_json.clone();
|
||||
unknown_kind["kind"] = Value::String("artifact".into());
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(unknown_kind).is_err(),
|
||||
"unknown kind must fail closed and must not become None"
|
||||
);
|
||||
|
||||
let mut unknown_field = none_json.clone();
|
||||
unknown_field
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("unexpected_field".into(), Value::Bool(true));
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(unknown_field).is_err(),
|
||||
"unknown outer field must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_version = none_json.clone();
|
||||
unknown_version["format_version"] = Value::from(2);
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(unknown_version).is_err(),
|
||||
"unsupported format_version must fail closed"
|
||||
);
|
||||
|
||||
let mut unexpected_function_on_none = none_json.clone();
|
||||
unexpected_function_on_none.as_object_mut().unwrap().insert(
|
||||
"function".into(),
|
||||
serde_json::to_value(&sample_function()?).expect("nested Function"),
|
||||
);
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(unexpected_function_on_none).is_err(),
|
||||
"kind=none with unexpected function field must fail closed"
|
||||
);
|
||||
|
||||
let function = sample_function()?;
|
||||
let function_json =
|
||||
serde_json::to_value(JobResult::Function(function)).expect("serialize Function");
|
||||
|
||||
let mut missing_function = function_json.clone();
|
||||
missing_function.as_object_mut().unwrap().remove("function");
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(missing_function).is_err(),
|
||||
"kind=function without function field must fail closed"
|
||||
);
|
||||
|
||||
let mut empty_function_id = function_json.clone();
|
||||
empty_function_id["function"]["id"] = Value::String("".into());
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(empty_function_id).is_err(),
|
||||
"empty nested Function ID must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_nested_function_field = function_json.clone();
|
||||
unknown_nested_function_field["function"]
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("unexpected_field".into(), Value::Bool(true));
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(unknown_nested_function_field).is_err(),
|
||||
"unknown nested Function field must fail closed"
|
||||
);
|
||||
|
||||
let mut malformed_nested_version = function_json.clone();
|
||||
malformed_nested_version["function"]["format_version"] = Value::from(2);
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(malformed_nested_version).is_err(),
|
||||
"malformed nested Function must fail closed"
|
||||
);
|
||||
|
||||
// Unknown wire must not be preserved as a public variant or downgraded to None.
|
||||
let unknown_raw = serde_json::json!({
|
||||
"format_version": 1,
|
||||
"kind": "future_result_kind",
|
||||
"raw": {"keep": true}
|
||||
});
|
||||
assert!(
|
||||
serde_json::from_value::<JobResult>(unknown_raw).is_err(),
|
||||
"unknown kind must fail closed without a public raw/unknown variant"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn outer_result_excludes_forbidden_fields() -> Result<()> {
|
||||
let none_json = serde_json::to_value(&JobResult::None).expect("serialize None");
|
||||
assert_json_object_keys_exact(&none_json, &["format_version", "kind"]);
|
||||
assert_outer_forbidden_keys_absent(&none_json);
|
||||
|
||||
let function = sample_function()?;
|
||||
let function_json =
|
||||
serde_json::to_value(JobResult::Function(function)).expect("serialize Function");
|
||||
assert_json_object_keys_exact(&function_json, &["format_version", "kind", "function"]);
|
||||
assert_outer_forbidden_keys_absent(&function_json);
|
||||
|
||||
// Nested Function may carry `id` and signature parameter `name`; those are
|
||||
// not outer result keys and must remain present on the nested object.
|
||||
assert_eq!(
|
||||
function_json["function"]["id"],
|
||||
Value::String("fn.exact.job-result".into())
|
||||
);
|
||||
assert!(
|
||||
function_json["function"]["signature"]["parameters"][0]
|
||||
.get("name")
|
||||
.is_some(),
|
||||
"nested signature parameter `name` remains legitimate"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,537 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Contract tests for RegisterFunctionJobSpec (FF-008 / B1d).
|
||||
//!
|
||||
//! These tests pin the intended public surface under [`lancedb::function`] for
|
||||
//! registration Job operation input only. They intentionally fail to compile
|
||||
//! until that API exists.
|
||||
//!
|
||||
//! Rejection cases are judged by `Result` structure (`is_err` / `is_ok`), never
|
||||
//! by diagnostic message substrings. Catalog execution, typed Job wait, and
|
||||
//! result Function publication are out of scope.
|
||||
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use arrow_schema::DataType;
|
||||
use lancedb::Result;
|
||||
use lancedb::function::{
|
||||
FunctionCapability, FunctionDefinition, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, PythonFunctionDefinition, RegisterFunctionJobSpec,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
fn sample_signature() -> Result<FunctionSignature> {
|
||||
FunctionSignature::try_new(
|
||||
vec![
|
||||
FunctionParameter::new("text", DataType::Utf8),
|
||||
FunctionParameter::new("limit", DataType::Int32),
|
||||
],
|
||||
FunctionOutput::new(DataType::Utf8, true),
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_source() -> &'static str {
|
||||
"def normalize(text, limit):\n return text[:limit]\n"
|
||||
}
|
||||
|
||||
fn sample_python_definition() -> Result<PythonFunctionDefinition> {
|
||||
PythonFunctionDefinition::try_new(
|
||||
"normalize_mod",
|
||||
"normalize",
|
||||
sample_source(),
|
||||
"3.12",
|
||||
vec!["Unidecode==1.3.8".to_string()],
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_capabilities() -> Result<Vec<FunctionCapability>> {
|
||||
Ok(vec![
|
||||
FunctionCapability::try_network("https://api.example.com")?,
|
||||
FunctionCapability::try_secret("secret://team/api-token", "API_TOKEN")?,
|
||||
])
|
||||
}
|
||||
|
||||
fn sample_definition() -> Result<FunctionDefinition> {
|
||||
FunctionDefinition::try_new(
|
||||
sample_signature()?,
|
||||
sample_python_definition()?,
|
||||
sample_capabilities()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn sample_create_spec() -> Result<RegisterFunctionJobSpec> {
|
||||
RegisterFunctionJobSpec::try_new("text.normalize", sample_definition()?, None)
|
||||
}
|
||||
|
||||
fn sample_replace_spec() -> Result<RegisterFunctionJobSpec> {
|
||||
RegisterFunctionJobSpec::try_new(
|
||||
"text.normalize",
|
||||
sample_definition()?,
|
||||
Some(FunctionId::try_new("fn.existing.exact")?),
|
||||
)
|
||||
}
|
||||
|
||||
fn assert_json_object_keys_exact(value: &Value, expected: &[&str]) {
|
||||
let object = value
|
||||
.as_object()
|
||||
.unwrap_or_else(|| panic!("expected JSON object, got {value}"));
|
||||
let keys: BTreeSet<&str> = object.keys().map(|k| k.as_str()).collect();
|
||||
let expected: BTreeSet<&str> = expected.iter().copied().collect();
|
||||
assert_eq!(
|
||||
keys, expected,
|
||||
"JSON object key set must match exactly (iteration order is not a contract); got {keys:?}, expected {expected:?} in {value}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Lifecycle / identity / artifact / runtime keys that must not appear as public
|
||||
/// object keys on RegisterFunctionJobSpec wire. Exact key matching only: do not
|
||||
/// substring-scan encoded JSON (source text, parameter `name`, and
|
||||
/// `expected_current_function_id` must not false-match).
|
||||
const FORBIDDEN_SPEC_KEYS: &[&str] = &[
|
||||
"id",
|
||||
"function_id",
|
||||
"FunctionId",
|
||||
"new_function_id",
|
||||
"generated_function_id",
|
||||
"result_function_id",
|
||||
"version",
|
||||
"function_version",
|
||||
"FunctionVersion",
|
||||
"user_version",
|
||||
"lineage",
|
||||
"idempotency_key",
|
||||
"retry_key",
|
||||
"idempotency",
|
||||
"job_id",
|
||||
"jobId",
|
||||
"state",
|
||||
"status",
|
||||
"attempt",
|
||||
"attempt_id",
|
||||
"timestamp",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"table",
|
||||
"table_name",
|
||||
"table_ref",
|
||||
"executor",
|
||||
"environment",
|
||||
"artifact",
|
||||
"artifact_digest",
|
||||
"digest",
|
||||
"storage",
|
||||
"storage_location",
|
||||
"location",
|
||||
"worker",
|
||||
"scheduler",
|
||||
"replica",
|
||||
"placement",
|
||||
];
|
||||
|
||||
fn assert_object_keys_not_forbidden(value: &Value, context: &str) {
|
||||
let object = value.as_object().unwrap_or_else(|| {
|
||||
panic!("expected JSON object at {context}, got {value}");
|
||||
});
|
||||
for key in object.keys() {
|
||||
assert!(
|
||||
!FORBIDDEN_SPEC_KEYS.contains(&key.as_str()),
|
||||
"RegisterFunctionJobSpec wire must not contain forbidden key `{key}` at {context}: {value}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_forbidden_spec_keys_absent(value: &Value) {
|
||||
assert_json_object_keys_exact(
|
||||
value,
|
||||
&[
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
],
|
||||
);
|
||||
assert_object_keys_not_forbidden(value, "RegisterFunctionJobSpec");
|
||||
|
||||
// Precondition field is allowed; a generated/new Function ID key is not.
|
||||
assert!(
|
||||
value.get("function_id").is_none(),
|
||||
"spec must not carry generated/new `function_id`; use expected_current_function_id only: {value}"
|
||||
);
|
||||
assert!(
|
||||
value.get("id").is_none(),
|
||||
"spec must not carry generated/new Function `id`: {value}"
|
||||
);
|
||||
|
||||
let definition = value.get("definition").expect("definition object");
|
||||
assert_json_object_keys_exact(
|
||||
definition,
|
||||
&[
|
||||
"format_version",
|
||||
"signature",
|
||||
"implementation",
|
||||
"capabilities",
|
||||
],
|
||||
);
|
||||
assert_object_keys_not_forbidden(definition, "definition");
|
||||
// Catalog/function identity name belongs on the spec, not the nested definition.
|
||||
assert!(
|
||||
definition.get("name").is_none(),
|
||||
"nested definition must not contain catalog name: {definition}"
|
||||
);
|
||||
|
||||
let implementation = definition
|
||||
.get("implementation")
|
||||
.expect("implementation object");
|
||||
assert_json_object_keys_exact(
|
||||
implementation,
|
||||
&["kind", "module", "callable", "source", "python", "packages"],
|
||||
);
|
||||
assert_object_keys_not_forbidden(implementation, "definition.implementation");
|
||||
|
||||
let capabilities = definition
|
||||
.get("capabilities")
|
||||
.and_then(Value::as_array)
|
||||
.expect("capabilities array");
|
||||
for (idx, capability) in capabilities.iter().enumerate() {
|
||||
let context = format!("definition.capabilities[{idx}]");
|
||||
assert_object_keys_not_forbidden(capability, &context);
|
||||
}
|
||||
|
||||
let signature = definition.get("signature").expect("signature object");
|
||||
assert_object_keys_not_forbidden(signature, "definition.signature");
|
||||
if let Some(parameters) = signature.get("parameters").and_then(Value::as_array) {
|
||||
for (idx, parameter) in parameters.iter().enumerate() {
|
||||
let context = format!("definition.signature.parameters[{idx}]");
|
||||
assert_object_keys_not_forbidden(parameter, &context);
|
||||
// Parameter object key `name` is legitimate and must not be treated
|
||||
// as a forbidden catalog/function identity field.
|
||||
assert!(
|
||||
parameter.get("name").is_some(),
|
||||
"{context} must include parameter `name`"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn create_and_replace_round_trip_pins_name_definition_and_precondition() -> Result<()> {
|
||||
let create = sample_create_spec()?;
|
||||
assert_eq!(create.format_version(), 1);
|
||||
assert_eq!(create.name(), "text.normalize");
|
||||
assert!(create.expected_current_function_id().is_none());
|
||||
assert_eq!(
|
||||
create.definition().python_definition().source(),
|
||||
sample_source()
|
||||
);
|
||||
assert_eq!(create.definition().capabilities().len(), 2);
|
||||
assert_eq!(
|
||||
create.definition().capabilities()[0].origin(),
|
||||
Some("https://api.example.com")
|
||||
);
|
||||
assert_eq!(
|
||||
create.definition().capabilities()[1].reference(),
|
||||
Some("secret://team/api-token")
|
||||
);
|
||||
|
||||
let create_json = serde_json::to_value(&create).expect("serialize create spec");
|
||||
assert_eq!(create_json["format_version"], 1);
|
||||
assert_eq!(create_json["name"], Value::String("text.normalize".into()));
|
||||
assert_eq!(create_json["expected_current_function_id"], Value::Null);
|
||||
|
||||
let create_restored: RegisterFunctionJobSpec =
|
||||
serde_json::from_value(create_json.clone()).expect("deserialize create spec");
|
||||
assert_eq!(create_restored.format_version(), 1);
|
||||
assert_eq!(create_restored.name(), "text.normalize");
|
||||
assert!(create_restored.expected_current_function_id().is_none());
|
||||
assert_eq!(
|
||||
create_restored.definition().python_definition().module(),
|
||||
create.definition().python_definition().module()
|
||||
);
|
||||
assert_eq!(
|
||||
create_restored.definition().python_definition().callable(),
|
||||
create.definition().python_definition().callable()
|
||||
);
|
||||
assert_eq!(
|
||||
create_restored.definition().python_definition().source(),
|
||||
create.definition().python_definition().source()
|
||||
);
|
||||
assert_eq!(
|
||||
create_restored.definition().python_definition().python(),
|
||||
create.definition().python_definition().python()
|
||||
);
|
||||
assert_eq!(
|
||||
create_restored.definition().python_definition().packages(),
|
||||
create.definition().python_definition().packages()
|
||||
);
|
||||
assert_eq!(
|
||||
create_restored.definition().capabilities()[1].reference(),
|
||||
create.definition().capabilities()[1].reference()
|
||||
);
|
||||
assert_eq!(
|
||||
create_restored.definition().capabilities()[1].environment_variable(),
|
||||
create.definition().capabilities()[1].environment_variable()
|
||||
);
|
||||
|
||||
let replace = sample_replace_spec()?;
|
||||
assert_eq!(replace.name(), "text.normalize");
|
||||
assert_eq!(
|
||||
replace
|
||||
.expected_current_function_id()
|
||||
.map(FunctionId::as_str),
|
||||
Some("fn.existing.exact")
|
||||
);
|
||||
|
||||
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
|
||||
assert_eq!(
|
||||
replace_json["expected_current_function_id"],
|
||||
Value::String("fn.existing.exact".into())
|
||||
);
|
||||
let replace_restored: RegisterFunctionJobSpec =
|
||||
serde_json::from_value(replace_json.clone()).expect("deserialize replace spec");
|
||||
assert_eq!(
|
||||
replace_restored
|
||||
.expected_current_function_id()
|
||||
.map(FunctionId::as_str),
|
||||
Some("fn.existing.exact")
|
||||
);
|
||||
assert_eq!(replace_restored.name(), replace.name());
|
||||
assert_eq!(
|
||||
replace_restored.definition().python_definition().source(),
|
||||
replace.definition().python_definition().source()
|
||||
);
|
||||
|
||||
// Byte-for-byte repeated serde_json encoding for the same value.
|
||||
let create_a = serde_json::to_string(&create).expect("encode create a");
|
||||
let create_b = serde_json::to_string(&create).expect("encode create b");
|
||||
assert_eq!(create_a, create_b);
|
||||
let replace_a = serde_json::to_string(&replace).expect("encode replace a");
|
||||
let replace_b = serde_json::to_string(&replace).expect("encode replace b");
|
||||
assert_eq!(replace_a, replace_b);
|
||||
assert_eq!(
|
||||
serde_json::to_value(&create_restored).expect("re-serialize create"),
|
||||
create_json
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_value(&replace_restored).expect("re-serialize replace"),
|
||||
replace_json
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn create_wire_pins_exact_key_set_and_null_precondition() -> Result<()> {
|
||||
let create = sample_create_spec()?;
|
||||
let json = serde_json::to_value(&create).expect("serialize create spec");
|
||||
|
||||
// Exact object key set; Map iteration order is not part of the contract.
|
||||
assert_json_object_keys_exact(
|
||||
&json,
|
||||
&[
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
],
|
||||
);
|
||||
assert_eq!(json["format_version"], 1);
|
||||
assert_eq!(json["name"], Value::String("text.normalize".into()));
|
||||
// Create-if-absent always serializes the precondition key as JSON null.
|
||||
assert_eq!(json["expected_current_function_id"], Value::Null);
|
||||
assert!(json["expected_current_function_id"].is_null());
|
||||
|
||||
let replace = sample_replace_spec()?;
|
||||
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
|
||||
assert_json_object_keys_exact(
|
||||
&replace_json,
|
||||
&[
|
||||
"format_version",
|
||||
"name",
|
||||
"definition",
|
||||
"expected_current_function_id",
|
||||
],
|
||||
);
|
||||
assert_eq!(
|
||||
replace_json["expected_current_function_id"],
|
||||
Value::String("fn.existing.exact".into())
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_name_and_expected_id_unknown_field_and_version_fail_closed() -> Result<()> {
|
||||
let definition = sample_definition()?;
|
||||
|
||||
assert!(
|
||||
RegisterFunctionJobSpec::try_new("", definition.clone(), None).is_err(),
|
||||
"empty name must be rejected by constructor"
|
||||
);
|
||||
assert!(
|
||||
RegisterFunctionJobSpec::try_new(
|
||||
"text.normalize",
|
||||
definition,
|
||||
Some(FunctionId::try_new("fn.existing.exact")?),
|
||||
)
|
||||
.is_ok(),
|
||||
"non-empty name with exact expected ID must construct"
|
||||
);
|
||||
|
||||
let create = sample_create_spec()?;
|
||||
let json = serde_json::to_value(&create).expect("serialize create spec");
|
||||
|
||||
let mut empty_name = json.clone();
|
||||
empty_name["name"] = Value::String("".into());
|
||||
assert!(
|
||||
serde_json::from_value::<RegisterFunctionJobSpec>(empty_name).is_err(),
|
||||
"decode must reject empty name"
|
||||
);
|
||||
|
||||
let mut empty_expected_id = json.clone();
|
||||
empty_expected_id["expected_current_function_id"] = Value::String("".into());
|
||||
assert!(
|
||||
serde_json::from_value::<RegisterFunctionJobSpec>(empty_expected_id).is_err(),
|
||||
"decode must reject empty expected_current_function_id"
|
||||
);
|
||||
|
||||
let mut unknown_field = json.clone();
|
||||
unknown_field
|
||||
.as_object_mut()
|
||||
.unwrap()
|
||||
.insert("unexpected_field".into(), Value::Bool(true));
|
||||
assert!(
|
||||
serde_json::from_value::<RegisterFunctionJobSpec>(unknown_field).is_err(),
|
||||
"unknown outer field must fail closed"
|
||||
);
|
||||
|
||||
let mut unknown_version = json.clone();
|
||||
unknown_version["format_version"] = Value::from(2);
|
||||
assert!(
|
||||
serde_json::from_value::<RegisterFunctionJobSpec>(unknown_version).is_err(),
|
||||
"format_version other than 1 must fail closed"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn spec_wire_excludes_generated_identity_job_lifecycle_and_artifact_fields() -> Result<()> {
|
||||
let create = sample_create_spec()?;
|
||||
let create_json = serde_json::to_value(&create).expect("serialize create spec");
|
||||
assert_forbidden_spec_keys_absent(&create_json);
|
||||
|
||||
let replace = sample_replace_spec()?;
|
||||
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
|
||||
assert_forbidden_spec_keys_absent(&replace_json);
|
||||
// Replace may carry the exact opaque precondition string only.
|
||||
assert_eq!(
|
||||
replace_json["expected_current_function_id"],
|
||||
Value::String("fn.existing.exact".into())
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn debug_redacts_source_and_secret_reference_while_nested_getters_remain_exact() -> Result<()> {
|
||||
let source = sample_source();
|
||||
let secret_reference = "secret://team/api-token";
|
||||
|
||||
let create = sample_create_spec()?;
|
||||
assert_eq!(create.definition().python_definition().source(), source);
|
||||
assert_eq!(
|
||||
create.definition().capabilities()[1].reference(),
|
||||
Some(secret_reference)
|
||||
);
|
||||
let create_debug = format!("{create:?}");
|
||||
assert!(
|
||||
!create_debug.contains(source),
|
||||
"create RegisterFunctionJobSpec Debug must not contain source body: {create_debug}"
|
||||
);
|
||||
assert!(
|
||||
!create_debug.contains("return text[:limit]"),
|
||||
"create RegisterFunctionJobSpec Debug must not leak source fragments: {create_debug}"
|
||||
);
|
||||
assert!(
|
||||
!create_debug.contains(secret_reference),
|
||||
"create RegisterFunctionJobSpec Debug must not contain secret reference: {create_debug}"
|
||||
);
|
||||
|
||||
let replace = sample_replace_spec()?;
|
||||
assert_eq!(replace.definition().python_definition().source(), source);
|
||||
assert_eq!(
|
||||
replace.definition().capabilities()[1].reference(),
|
||||
Some(secret_reference)
|
||||
);
|
||||
assert_eq!(
|
||||
replace
|
||||
.expected_current_function_id()
|
||||
.map(FunctionId::as_str),
|
||||
Some("fn.existing.exact")
|
||||
);
|
||||
let replace_debug = format!("{replace:?}");
|
||||
assert!(
|
||||
!replace_debug.contains(source),
|
||||
"replace RegisterFunctionJobSpec Debug must not contain source body: {replace_debug}"
|
||||
);
|
||||
assert!(
|
||||
!replace_debug.contains(secret_reference),
|
||||
"replace RegisterFunctionJobSpec Debug must not contain secret reference: {replace_debug}"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn nested_definition_wire_is_exact_ff007_function_definition() -> Result<()> {
|
||||
let definition = sample_definition()?;
|
||||
let expected_definition_wire =
|
||||
serde_json::to_value(&definition).expect("serialize FunctionDefinition");
|
||||
|
||||
let create = RegisterFunctionJobSpec::try_new("text.normalize", definition.clone(), None)?;
|
||||
let create_json = serde_json::to_value(&create).expect("serialize create spec");
|
||||
assert_eq!(
|
||||
create_json["definition"], expected_definition_wire,
|
||||
"nested definition must be the exact FF-007 FunctionDefinition wire"
|
||||
);
|
||||
assert_json_object_keys_exact(
|
||||
&create_json["definition"],
|
||||
&[
|
||||
"format_version",
|
||||
"signature",
|
||||
"implementation",
|
||||
"capabilities",
|
||||
],
|
||||
);
|
||||
assert_eq!(
|
||||
create_json["definition"]["implementation"]["source"],
|
||||
Value::String(sample_source().into())
|
||||
);
|
||||
assert_eq!(
|
||||
create_json["definition"]["capabilities"][1]["reference"],
|
||||
Value::String("secret://team/api-token".into())
|
||||
);
|
||||
// Not a summarized / digested / artifact reference indirection.
|
||||
assert!(create_json["definition"].get("digest").is_none());
|
||||
assert!(create_json["definition"].get("artifact").is_none());
|
||||
assert!(create_json["definition"].get("artifact_digest").is_none());
|
||||
assert!(create_json["definition"].get("storage").is_none());
|
||||
assert!(
|
||||
create_json["definition"]["implementation"]
|
||||
.get("digest")
|
||||
.is_none()
|
||||
);
|
||||
assert!(
|
||||
create_json["definition"]["implementation"]
|
||||
.get("artifact")
|
||||
.is_none()
|
||||
);
|
||||
|
||||
let replace = RegisterFunctionJobSpec::try_new(
|
||||
"text.normalize",
|
||||
definition,
|
||||
Some(FunctionId::try_new("fn.existing.exact")?),
|
||||
)?;
|
||||
let replace_json = serde_json::to_value(&replace).expect("serialize replace spec");
|
||||
assert_eq!(
|
||||
replace_json["definition"], expected_definition_wire,
|
||||
"replace nested definition must remain the exact FF-007 wire"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,495 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Contract tests for GeneratedColumnBindingSnapshot (FF-029 / FF-030).
|
||||
//!
|
||||
//! Pins the hidden value projection, Table seam, and bound-call field
|
||||
//! validation used by generated-column call binding. These tests intentionally
|
||||
//! fail to compile until that API exists. They do not submit Jobs, mutate
|
||||
//! generated-column state, or resolve authored Function calls.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::{ArrayRef, Int32Array, RecordBatch, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use lance::dataset::NewColumnTransform;
|
||||
use lancedb::connect;
|
||||
use lancedb::function::{
|
||||
Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter,
|
||||
FunctionSignature, GeneratedColumnBindingEntry, GeneratedColumnBindingSnapshot,
|
||||
};
|
||||
use lancedb::table::ColumnAlteration;
|
||||
use lancedb::{Error, Result};
|
||||
use tempfile::tempdir;
|
||||
|
||||
fn sample_fields() -> Vec<arrow_schema::FieldRef> {
|
||||
vec![
|
||||
Arc::new(Field::new("text", DataType::Utf8, true)),
|
||||
Arc::new(Field::new("score", DataType::Int32, false)),
|
||||
Arc::new(Field::new("a.b", DataType::Utf8, true)),
|
||||
]
|
||||
}
|
||||
|
||||
fn sample_output() -> FunctionOutput {
|
||||
FunctionOutput::new(DataType::Int32, true)
|
||||
}
|
||||
|
||||
/// Parameter names intentionally differ from table column names so any
|
||||
/// name-based validation would fail these fixtures.
|
||||
fn two_field_function() -> Result<Function> {
|
||||
let id = FunctionId::try_new("fn.exact.binding.validate")?;
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![
|
||||
FunctionParameter::new("input_payload", DataType::Utf8),
|
||||
FunctionParameter::new("metric_value", DataType::Int32),
|
||||
],
|
||||
sample_output(),
|
||||
)?;
|
||||
Ok(Function::new(id, signature))
|
||||
}
|
||||
|
||||
fn one_field_function() -> Result<Function> {
|
||||
let id = FunctionId::try_new("fn.exact.binding.one-field")?;
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("payload_arg", DataType::Utf8)],
|
||||
sample_output(),
|
||||
)?;
|
||||
Ok(Function::new(id, signature))
|
||||
}
|
||||
|
||||
fn literal_only_function() -> Result<Function> {
|
||||
let id = FunctionId::try_new("fn.exact.binding.literal-only")?;
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("constant_arg", DataType::Int32)],
|
||||
sample_output(),
|
||||
)?;
|
||||
Ok(Function::new(id, signature))
|
||||
}
|
||||
|
||||
fn mixed_function() -> Result<Function> {
|
||||
let id = FunctionId::try_new("fn.exact.binding.mixed")?;
|
||||
let signature = FunctionSignature::try_new(
|
||||
vec![
|
||||
FunctionParameter::new("payload_arg", DataType::Utf8),
|
||||
FunctionParameter::new("constant_arg", DataType::Int32),
|
||||
],
|
||||
sample_output(),
|
||||
)?;
|
||||
Ok(Function::new(id, signature))
|
||||
}
|
||||
|
||||
fn int_literal(value: Option<i32>) -> Result<FunctionArgument> {
|
||||
FunctionArgument::try_literal(Arc::new(Int32Array::from(vec![value])) as ArrayRef)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn try_new_preserves_version_order_and_exact_lookup() -> Result<()> {
|
||||
let fields = sample_fields();
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(7, fields.clone(), vec![3, 5, 9])?;
|
||||
|
||||
assert_eq!(snapshot.version(), 7);
|
||||
let entries = snapshot.entries();
|
||||
assert_eq!(entries.len(), 3);
|
||||
assert_eq!(entries[0].field_id(), 3);
|
||||
assert_eq!(entries[0].field().name(), "text");
|
||||
assert_eq!(entries[0].field().data_type(), &DataType::Utf8);
|
||||
assert_eq!(entries[1].field_id(), 5);
|
||||
assert_eq!(entries[1].field().name(), "score");
|
||||
assert_eq!(entries[2].field_id(), 9);
|
||||
assert_eq!(entries[2].field().name(), "a.b");
|
||||
|
||||
let by_name = snapshot.field("score").expect("exact name");
|
||||
assert_eq!(by_name.field_id(), 5);
|
||||
assert!(snapshot.field("Score").is_none());
|
||||
assert!(snapshot.field("a").is_none());
|
||||
let dotted = snapshot
|
||||
.field("a.b")
|
||||
.expect("literal dotted top-level name");
|
||||
assert_eq!(dotted.field_id(), 9);
|
||||
assert_eq!(dotted.field().as_ref(), fields[2].as_ref());
|
||||
|
||||
// Type existence pin for the entry surface used by the next binding slice.
|
||||
let _: &GeneratedColumnBindingEntry = by_name;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn try_new_rejects_invalid_projections() {
|
||||
let fields = sample_fields();
|
||||
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, fields.clone(), vec![1, 2]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, fields.clone(), vec![1, 2, -1]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, fields.clone(), vec![1, 2, 1]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
|
||||
let duplicate_names = vec![
|
||||
Arc::new(Field::new("text", DataType::Utf8, true)),
|
||||
Arc::new(Field::new("text", DataType::Int32, false)),
|
||||
];
|
||||
assert!(matches!(
|
||||
GeneratedColumnBindingSnapshot::try_new(1, duplicate_names, vec![1, 2]),
|
||||
Err(Error::InvalidInput { .. })
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_field_arguments_accepts_valid_and_mixed_bindings() -> Result<()> {
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(3, sample_fields(), vec![3, 5, 9])?;
|
||||
|
||||
let one = one_field_function()?;
|
||||
let valid = FunctionCall::try_new(
|
||||
&one,
|
||||
vec![(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(3, DataType::Utf8)?,
|
||||
)],
|
||||
)?;
|
||||
snapshot.validate_field_arguments(&valid)?;
|
||||
|
||||
let two = two_field_function()?;
|
||||
let multi = FunctionCall::try_new(
|
||||
&two,
|
||||
vec![
|
||||
(
|
||||
"input_payload".to_string(),
|
||||
FunctionArgument::try_field(3, DataType::Utf8)?,
|
||||
),
|
||||
(
|
||||
"metric_value".to_string(),
|
||||
FunctionArgument::try_field(5, DataType::Int32)?,
|
||||
),
|
||||
],
|
||||
)?;
|
||||
snapshot.validate_field_arguments(&multi)?;
|
||||
|
||||
let mixed_fn = mixed_function()?;
|
||||
let mixed = FunctionCall::try_new(
|
||||
&mixed_fn,
|
||||
vec![
|
||||
(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(3, DataType::Utf8)?,
|
||||
),
|
||||
("constant_arg".to_string(), int_literal(Some(42))?),
|
||||
],
|
||||
)?;
|
||||
snapshot.validate_field_arguments(&mixed)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_field_arguments_rejects_missing_id_and_type_mismatch() -> Result<()> {
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(3, sample_fields(), vec![3, 5, 9])?;
|
||||
let one = one_field_function()?;
|
||||
|
||||
let missing = FunctionCall::try_new(
|
||||
&one,
|
||||
vec![(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(99, DataType::Utf8)?,
|
||||
)],
|
||||
)?;
|
||||
let err = snapshot
|
||||
.validate_field_arguments(&missing)
|
||||
.expect_err("missing stable field id");
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains("99"),
|
||||
"diagnostics may name field id: {message}"
|
||||
);
|
||||
assert!(
|
||||
!message.contains("text") && !message.contains("score") && !message.contains("a.b"),
|
||||
"diagnostics must not invent or use a column name: {message}"
|
||||
);
|
||||
|
||||
// Same stable ID, different Arrow type: exact-type equality must reject.
|
||||
// This covers Remote/other producer projections that keep the ID.
|
||||
let type_mismatch = FunctionCall::try_new(
|
||||
&one,
|
||||
vec![(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(5, DataType::Utf8)?,
|
||||
)],
|
||||
)?;
|
||||
let err = snapshot
|
||||
.validate_field_arguments(&type_mismatch)
|
||||
.expect_err("same-id exact type mismatch");
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains("5"),
|
||||
"diagnostics may name field id: {message}"
|
||||
);
|
||||
assert!(
|
||||
message.contains("Utf8") && message.contains("Int32"),
|
||||
"diagnostics may identify expected/current types: {message}"
|
||||
);
|
||||
assert!(
|
||||
!message.contains("score") && !message.contains("text"),
|
||||
"diagnostics must not invent or use a column name: {message}"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validate_field_arguments_literal_only_ignores_table_fields() -> Result<()> {
|
||||
// Snapshot has no field that a name-based binder could match to "constant_arg".
|
||||
let snapshot = GeneratedColumnBindingSnapshot::try_new(
|
||||
1,
|
||||
vec![Arc::new(Field::new("unrelated", DataType::Utf8, true))],
|
||||
vec![11],
|
||||
)?;
|
||||
let function = literal_only_function()?;
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![("constant_arg".to_string(), int_literal(Some(7))?)],
|
||||
)?;
|
||||
snapshot.validate_field_arguments(&call)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn table_seam_returns_atomic_native_snapshot() -> Result<()> {
|
||||
let tmp = tempdir().unwrap();
|
||||
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new("text", DataType::Utf8, true),
|
||||
Field::new("score", DataType::Int32, false),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(StringArray::from(vec![Some("a")])),
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
],
|
||||
)?;
|
||||
let table = db.create_table("binding", batch).execute().await?;
|
||||
|
||||
let snapshot = table.generated_column_binding_snapshot().await?;
|
||||
let public_schema = table.schema().await?;
|
||||
let version = table.version().await?;
|
||||
|
||||
assert_eq!(snapshot.version(), version);
|
||||
assert_eq!(snapshot.entries().len(), public_schema.fields().len());
|
||||
for (entry, field) in snapshot.entries().iter().zip(public_schema.fields()) {
|
||||
assert_eq!(entry.field().name(), field.name());
|
||||
assert_eq!(entry.field().data_type(), field.data_type());
|
||||
assert!(entry.field_id() >= 0);
|
||||
assert!(!field.metadata().contains_key("lance:field_id"));
|
||||
assert!(!entry.field().metadata().contains_key("lance:field_id"));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn validate_field_arguments_survives_rename_on_real_table() -> Result<()> {
|
||||
let tmp = tempdir().unwrap();
|
||||
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
|
||||
// Column names deliberately differ from Function parameter names.
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new("source_text", DataType::Utf8, true),
|
||||
Field::new("source_score", DataType::Int32, false),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(StringArray::from(vec![Some("hello")])),
|
||||
Arc::new(Int32Array::from(vec![7])),
|
||||
],
|
||||
)?;
|
||||
let table = db.create_table("binding_rename", batch).execute().await?;
|
||||
|
||||
let before = table.generated_column_binding_snapshot().await?;
|
||||
let text_entry = before.field("source_text").expect("source_text");
|
||||
let score_entry = before.field("source_score").expect("source_score");
|
||||
let text_id = text_entry.field_id();
|
||||
let score_id = score_entry.field_id();
|
||||
assert_eq!(text_entry.field().data_type(), &DataType::Utf8);
|
||||
assert_eq!(score_entry.field().data_type(), &DataType::Int32);
|
||||
|
||||
let function = two_field_function()?;
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![
|
||||
(
|
||||
"input_payload".to_string(),
|
||||
FunctionArgument::try_field(text_id, DataType::Utf8)?,
|
||||
),
|
||||
(
|
||||
"metric_value".to_string(),
|
||||
FunctionArgument::try_field(score_id, DataType::Int32)?,
|
||||
),
|
||||
],
|
||||
)?;
|
||||
before.validate_field_arguments(&call)?;
|
||||
|
||||
table
|
||||
.alter_columns(&[ColumnAlteration::new("source_text".into()).rename("renamed_text".into())])
|
||||
.await?;
|
||||
|
||||
let after = table.generated_column_binding_snapshot().await?;
|
||||
assert!(after.field("source_text").is_none());
|
||||
let renamed = after.field("renamed_text").expect("renamed_text");
|
||||
assert_eq!(renamed.field_id(), text_id);
|
||||
assert_eq!(renamed.field().data_type(), &DataType::Utf8);
|
||||
assert_eq!(
|
||||
after
|
||||
.field("source_score")
|
||||
.expect("source_score")
|
||||
.field_id(),
|
||||
score_id
|
||||
);
|
||||
after.validate_field_arguments(&call)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn validate_field_arguments_rejects_drop_recreate_same_name_type() -> Result<()> {
|
||||
let tmp = tempdir().unwrap();
|
||||
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new("keep_col", DataType::Int32, false),
|
||||
Field::new("bound_col", DataType::Utf8, true),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(Int32Array::from(vec![1])),
|
||||
Arc::new(StringArray::from(vec![Some("v")])),
|
||||
],
|
||||
)?;
|
||||
let table = db
|
||||
.create_table("binding_drop_recreate", batch)
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
let before = table.generated_column_binding_snapshot().await?;
|
||||
let bound = before.field("bound_col").expect("bound_col");
|
||||
let old_id = bound.field_id();
|
||||
assert_eq!(bound.field().data_type(), &DataType::Utf8);
|
||||
|
||||
let function = one_field_function()?;
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"payload_arg".to_string(),
|
||||
FunctionArgument::try_field(old_id, DataType::Utf8)?,
|
||||
)],
|
||||
)?;
|
||||
before.validate_field_arguments(&call)?;
|
||||
|
||||
table.drop_columns(&["bound_col"]).await?;
|
||||
table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::SqlExpressions(vec![(
|
||||
"bound_col".into(),
|
||||
"cast(NULL as string)".into(),
|
||||
)]))
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
let after = table.generated_column_binding_snapshot().await?;
|
||||
let recreated = after.field("bound_col").expect("recreated bound_col");
|
||||
assert_eq!(recreated.field().data_type(), &DataType::Utf8);
|
||||
assert_ne!(
|
||||
recreated.field_id(),
|
||||
old_id,
|
||||
"drop/recreate must allocate a new stable field id"
|
||||
);
|
||||
let err = after
|
||||
.validate_field_arguments(&call)
|
||||
.expect_err("old call must not bind by name");
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains(&old_id.to_string()),
|
||||
"diagnostics may name missing field id: {message}"
|
||||
);
|
||||
assert!(
|
||||
!message.contains("bound_col") && !message.contains("keep_col"),
|
||||
"diagnostics must not invent or use a column name: {message}"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn validate_field_arguments_rejects_cast_that_allocates_new_field_id() -> Result<()> {
|
||||
// Native Lance cast_to allocates a new stable field ID. The old bound call
|
||||
// must fail because that ID is absent. Same-ID exact-type mismatch is proved
|
||||
// separately via manually constructed snapshots (Remote/other producers).
|
||||
let tmp = tempdir().unwrap();
|
||||
let db = connect(tmp.path().to_str().unwrap()).execute().await?;
|
||||
let schema = Arc::new(Schema::new(vec![
|
||||
Field::new("label_col", DataType::Utf8, true),
|
||||
Field::new("metric_col", DataType::Int32, false),
|
||||
]));
|
||||
let batch = RecordBatch::try_new(
|
||||
schema,
|
||||
vec![
|
||||
Arc::new(StringArray::from(vec![Some("x")])),
|
||||
Arc::new(Int32Array::from(vec![3])),
|
||||
],
|
||||
)?;
|
||||
let table = db
|
||||
.create_table("binding_type_change", batch)
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
let before = table.generated_column_binding_snapshot().await?;
|
||||
let metric = before.field("metric_col").expect("metric_col");
|
||||
let old_metric_id = metric.field_id();
|
||||
assert_eq!(metric.field().data_type(), &DataType::Int32);
|
||||
|
||||
let id = FunctionId::try_new("fn.exact.binding.type-change")?;
|
||||
let function = Function::new(
|
||||
id,
|
||||
FunctionSignature::try_new(
|
||||
vec![FunctionParameter::new("metric_value", DataType::Int32)],
|
||||
sample_output(),
|
||||
)?,
|
||||
);
|
||||
let call = FunctionCall::try_new(
|
||||
&function,
|
||||
vec![(
|
||||
"metric_value".to_string(),
|
||||
FunctionArgument::try_field(old_metric_id, DataType::Int32)?,
|
||||
)],
|
||||
)?;
|
||||
before.validate_field_arguments(&call)?;
|
||||
|
||||
table
|
||||
.alter_columns(&[ColumnAlteration::new("metric_col".into()).cast_to(DataType::Int64)])
|
||||
.await?;
|
||||
|
||||
let after = table.generated_column_binding_snapshot().await?;
|
||||
let casted = after.field("metric_col").expect("metric_col");
|
||||
assert_eq!(casted.field().data_type(), &DataType::Int64);
|
||||
assert_ne!(
|
||||
casted.field_id(),
|
||||
old_metric_id,
|
||||
"Native Lance cast_to must allocate a new stable field id"
|
||||
);
|
||||
let err = after
|
||||
.validate_field_arguments(&call)
|
||||
.expect_err("old call must fail because the prior stable field id is absent");
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
let message = err.to_string();
|
||||
assert!(
|
||||
message.contains(&old_metric_id.to_string()),
|
||||
"diagnostics may name missing field id: {message}"
|
||||
);
|
||||
assert!(
|
||||
!message.contains("metric_col") && !message.contains("label_col"),
|
||||
"diagnostics must not invent or use a column name: {message}"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user