mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-26 07:58:31 +00:00
Compare commits
14 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c463ca1503 | |||
| 79e9dffd06 | |||
| c3efc320a6 | |||
| 8d6dea6313 | |||
| 5a2d3f39e2 | |||
| 219f41339d | |||
| 5982eebbb3 | |||
| 5a1c839382 | |||
| e40e073a9d | |||
| a47c22b26e | |||
| 6fb976cf89 | |||
| a615306f39 | |||
| 920fc0e455 | |||
| 5acce6782e |
@@ -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
+42
-42
@@ -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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
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.7"
|
||||
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.7#e581c49338bc83baf1ea50c5e235bd702f3fbeea"
|
||||
dependencies = [
|
||||
"frostem",
|
||||
"icu_segmenter",
|
||||
|
||||
+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.7", default-features = false, "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { "version" = "=11.0.0-beta.7", default-features = false, "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { "version" = "=11.0.0-beta.7", default-features = false, "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { "version" = "=11.0.0-beta.7", "tag" = "v11.0.0-beta.7", "git" = "https://github.com/lance-format/lance.git" }
|
||||
ahash = "0.8"
|
||||
# Note that this one does not include pyarrow
|
||||
arrow = { version = "58.0.0", optional = false }
|
||||
|
||||
@@ -101,6 +101,13 @@ ignore = [
|
||||
# https://rustsec.org/advisories/RUSTSEC-2026-0195
|
||||
{ id = "RUSTSEC-2026-0194", reason = "transitive via inferno/lance/opendal; XML from trusted cloud endpoints, not attacker-controlled" },
|
||||
{ id = "RUSTSEC-2026-0195", reason = "transitive via inferno/lance/opendal; XML from trusted cloud endpoints, not attacker-controlled" },
|
||||
# smartstring: unmaintained — the repository was archived by its author on
|
||||
# 2026-05-03. Not a vulnerability. Reached only transitively through polars
|
||||
# (polars-core/-io/-ops/-time/-utils); nothing in LanceDB depends on it directly.
|
||||
# The advisory states no safe upgrade is available: upstream recommends
|
||||
# compact_str/smol_str, so clearing this requires polars to migrate.
|
||||
# https://rustsec.org/advisories/RUSTSEC-2026-0249
|
||||
{ id = "RUSTSEC-2026-0249", reason = "smartstring unmaintained via polars; no fixed upstream release" },
|
||||
]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -69,14 +69,33 @@ abstract addColumns(newColumnTransforms): Promise<AddColumnsResult>
|
||||
|
||||
Add new columns with defined values.
|
||||
|
||||
The `{ computed }` form stores the expression rather than evaluating it
|
||||
now: the column is committed with no values, and rows get them from
|
||||
[Table#refreshColumn](Table.md#refreshcolumn). Declaring one therefore costs the same on a
|
||||
large table as on an empty one.
|
||||
|
||||
A refresh does not revisit rows it has already filled, so mutating an
|
||||
input leaves the value computed at fill time; recomputing means dropping
|
||||
the column and declaring it again. While a declaration reads a column,
|
||||
that column cannot be renamed, retyped or dropped.
|
||||
|
||||
Computed columns are local-only: LanceDB Cloud and Enterprise reject a
|
||||
declaration.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **newColumnTransforms**: `Field`<`any`> \| `Field`<`any`>[] \| `Schema`<`any`> \| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
|
||||
* **newColumnTransforms**:
|
||||
\| `Field`<`any`>
|
||||
\| `Field`<`any`>[]
|
||||
\| `Schema`<`any`>
|
||||
\| [`AddColumnsSql`](../interfaces/AddColumnsSql.md)[]
|
||||
\| `object`
|
||||
Either:
|
||||
- An array of objects with column names and SQL expressions to calculate values
|
||||
- A single Arrow Field defining one column with its data type (column will be initialized with null values)
|
||||
- An array of Arrow Fields defining columns with their data types (columns will be initialized with null values)
|
||||
- An Arrow Schema defining columns with their data types (columns will be initialized with null values)
|
||||
- `{ computed }`, declaring columns defined by a SQL expression whose type and inputs are derived from it
|
||||
|
||||
#### Returns
|
||||
|
||||
@@ -85,6 +104,13 @@ Add new columns with defined values.
|
||||
A promise that resolves to an object
|
||||
containing the new version number of the table after adding the columns.
|
||||
|
||||
#### Example
|
||||
|
||||
```ts
|
||||
await table.addColumns({ computed: [{ name: "doubled", valueSql: "x * 2" }] });
|
||||
const { rowsFilled } = await table.refreshColumn("doubled");
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### alterColumns()
|
||||
@@ -718,6 +744,32 @@ for await (const batch of table.query()) {
|
||||
|
||||
***
|
||||
|
||||
### refreshColumn()
|
||||
|
||||
```ts
|
||||
abstract refreshColumn(column): Promise<RefreshColumnResult>
|
||||
```
|
||||
|
||||
Fill the rows of a computed column that hold no value yet.
|
||||
|
||||
Rows appended since the last refresh are filled by the next one; rows
|
||||
already filled are left as they are, so the call is idempotent and does
|
||||
not observe a mutated input. Local tables only.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* **column**: `string`
|
||||
The name of the computed column to fill.
|
||||
|
||||
#### Returns
|
||||
|
||||
`Promise`<[`RefreshColumnResult`](../interfaces/RefreshColumnResult.md)>
|
||||
|
||||
A promise that resolves to the
|
||||
number of rows filled and the new version number of the table.
|
||||
|
||||
***
|
||||
|
||||
### restore()
|
||||
|
||||
```ts
|
||||
|
||||
@@ -105,6 +105,7 @@
|
||||
- [OptimizeOptions](interfaces/OptimizeOptions.md)
|
||||
- [OptimizeStats](interfaces/OptimizeStats.md)
|
||||
- [QueryExecutionOptions](interfaces/QueryExecutionOptions.md)
|
||||
- [RefreshColumnResult](interfaces/RefreshColumnResult.md)
|
||||
- [RemovalStats](interfaces/RemovalStats.md)
|
||||
- [RenameTableOptions](interfaces/RenameTableOptions.md)
|
||||
- [RestNamespaceConfig](interfaces/RestNamespaceConfig.md)
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
[**@lancedb/lancedb**](../README.md) • **Docs**
|
||||
|
||||
***
|
||||
|
||||
[@lancedb/lancedb](../globals.md) / RefreshColumnResult
|
||||
|
||||
# Interface: RefreshColumnResult
|
||||
|
||||
## Properties
|
||||
|
||||
### rowsFilled
|
||||
|
||||
```ts
|
||||
rowsFilled: number;
|
||||
```
|
||||
|
||||
***
|
||||
|
||||
### version
|
||||
|
||||
```ts
|
||||
version: number;
|
||||
```
|
||||
+1
-1
@@ -28,7 +28,7 @@
|
||||
<properties>
|
||||
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
|
||||
<arrow.version>15.0.0</arrow.version>
|
||||
<lance-core.version>11.0.0-beta.3</lance-core.version>
|
||||
<lance-core.version>11.0.0-beta.6</lance-core.version>
|
||||
<spotless.skip>false</spotless.skip>
|
||||
<spotless.version>2.30.0</spotless.version>
|
||||
<spotless.java.googlejavaformat.version>1.7</spotless.java.googlejavaformat.version>
|
||||
|
||||
@@ -3340,3 +3340,45 @@ describe("LSM merge insert", () => {
|
||||
await expect(table.query().useLsm(true).toArray()).rejects.toThrow();
|
||||
});
|
||||
});
|
||||
|
||||
describe("computed columns", () => {
|
||||
let tmpDir: tmp.DirResult;
|
||||
beforeEach(() => {
|
||||
tmpDir = tmp.dirSync({ unsafeCleanup: true });
|
||||
});
|
||||
afterEach(() => tmpDir.removeCallback());
|
||||
|
||||
it("declares a column and fills it on refresh", async () => {
|
||||
const db = await connect(tmpDir.name);
|
||||
const table = await db.createTable("computed", [{ x: 1 }, { x: 2 }]);
|
||||
|
||||
await table.addColumns({
|
||||
computed: [{ name: "doubled", valueSql: "x * 2" }],
|
||||
});
|
||||
let rows = await table.query().toArray();
|
||||
expect(rows.map((r) => r.doubled)).toEqual([null, null]);
|
||||
|
||||
const result = await table.refreshColumn("doubled");
|
||||
expect(result.rowsFilled).toBe(2);
|
||||
|
||||
rows = await table.query().toArray();
|
||||
expect(rows.map((r) => r.doubled).sort()).toEqual([2, 4]);
|
||||
});
|
||||
|
||||
it("fills rows added since the last refresh", async () => {
|
||||
const db = await connect(tmpDir.name);
|
||||
const table = await db.createTable("computed_append", [{ x: 1 }]);
|
||||
|
||||
await table.addColumns({
|
||||
computed: [{ name: "doubled", valueSql: "x * 2" }],
|
||||
});
|
||||
await table.refreshColumn("doubled");
|
||||
await table.add([{ x: 5 }]);
|
||||
|
||||
const result = await table.refreshColumn("doubled");
|
||||
expect(result.rowsFilled).toBe(1);
|
||||
|
||||
const rows = await table.query().toArray();
|
||||
expect(rows.map((r) => r.doubled).sort()).toEqual([10, 2]);
|
||||
});
|
||||
});
|
||||
|
||||
@@ -50,6 +50,7 @@ export {
|
||||
MergeResult,
|
||||
AddResult,
|
||||
AddColumnsResult,
|
||||
RefreshColumnResult,
|
||||
AlterColumnsResult,
|
||||
UpdateFieldMetadataResult,
|
||||
DeleteResult,
|
||||
|
||||
+57
-2
@@ -33,6 +33,7 @@ import {
|
||||
Job,
|
||||
Branches as NativeBranches,
|
||||
OptimizeStats,
|
||||
RefreshColumnResult,
|
||||
TableStatistics,
|
||||
Tags,
|
||||
UpdateFieldMetadataResult,
|
||||
@@ -525,18 +526,54 @@ export abstract class Table {
|
||||
abstract vectorSearch(vector: IntoVector | MultiVector): VectorQuery;
|
||||
/**
|
||||
* Add new columns with defined values.
|
||||
*
|
||||
* The `{ computed }` form stores the expression rather than evaluating it
|
||||
* now: the column is committed with no values, and rows get them from
|
||||
* {@link Table#refreshColumn}. Declaring one therefore costs the same on a
|
||||
* large table as on an empty one.
|
||||
*
|
||||
* A refresh does not revisit rows it has already filled, so mutating an
|
||||
* input leaves the value computed at fill time; recomputing means dropping
|
||||
* the column and declaring it again. While a declaration reads a column,
|
||||
* that column cannot be renamed, retyped or dropped.
|
||||
*
|
||||
* Computed columns are local-only: LanceDB Cloud and Enterprise reject a
|
||||
* declaration.
|
||||
* @param {AddColumnsSql[] | Field | Field[] | Schema} newColumnTransforms Either:
|
||||
* - An array of objects with column names and SQL expressions to calculate values
|
||||
* - A single Arrow Field defining one column with its data type (column will be initialized with null values)
|
||||
* - An array of Arrow Fields defining columns with their data types (columns will be initialized with null values)
|
||||
* - An Arrow Schema defining columns with their data types (columns will be initialized with null values)
|
||||
* - `{ computed }`, declaring columns defined by a SQL expression whose type and inputs are derived from it
|
||||
* @returns {Promise<AddColumnsResult>} A promise that resolves to an object
|
||||
* containing the new version number of the table after adding the columns.
|
||||
* @example
|
||||
* ```ts
|
||||
* await table.addColumns({ computed: [{ name: "doubled", valueSql: "x * 2" }] });
|
||||
* const { rowsFilled } = await table.refreshColumn("doubled");
|
||||
* ```
|
||||
*/
|
||||
abstract addColumns(
|
||||
newColumnTransforms: AddColumnsSql[] | Field | Field[] | Schema,
|
||||
newColumnTransforms:
|
||||
| AddColumnsSql[]
|
||||
| Field
|
||||
| Field[]
|
||||
| Schema
|
||||
| { computed: AddColumnsSql[] },
|
||||
): Promise<AddColumnsResult>;
|
||||
|
||||
/**
|
||||
* Fill the rows of a computed column that hold no value yet.
|
||||
*
|
||||
* Rows appended since the last refresh are filled by the next one; rows
|
||||
* already filled are left as they are, so the call is idempotent and does
|
||||
* not observe a mutated input. Local tables only.
|
||||
* @param {string} column The name of the computed column to fill.
|
||||
* @returns {Promise<RefreshColumnResult>} A promise that resolves to the
|
||||
* number of rows filled and the new version number of the table.
|
||||
*/
|
||||
abstract refreshColumn(column: string): Promise<RefreshColumnResult>;
|
||||
|
||||
/**
|
||||
* Alter the name or nullability of columns.
|
||||
* @param {ColumnAlteration[]} columnAlterations One or more alterations to
|
||||
@@ -1088,8 +1125,22 @@ export class LocalTable extends Table {
|
||||
// TODO: Support BatchUDF
|
||||
|
||||
async addColumns(
|
||||
newColumnTransforms: AddColumnsSql[] | Field | Field[] | Schema,
|
||||
newColumnTransforms:
|
||||
| AddColumnsSql[]
|
||||
| Field
|
||||
| Field[]
|
||||
| Schema
|
||||
| { computed: AddColumnsSql[] },
|
||||
): Promise<AddColumnsResult> {
|
||||
// Columns defined by an expression are declared, not materialized here.
|
||||
if (
|
||||
typeof newColumnTransforms === "object" &&
|
||||
!Array.isArray(newColumnTransforms) &&
|
||||
"computed" in newColumnTransforms
|
||||
) {
|
||||
return await this.inner.addComputedColumns(newColumnTransforms.computed);
|
||||
}
|
||||
|
||||
// Handle single Field -> convert to array of Fields
|
||||
if (newColumnTransforms instanceof Field) {
|
||||
newColumnTransforms = [newColumnTransforms];
|
||||
@@ -1124,6 +1175,10 @@ export class LocalTable extends Table {
|
||||
throw new Error("Invalid input type for addColumns");
|
||||
}
|
||||
|
||||
async refreshColumn(column: string): Promise<RefreshColumnResult> {
|
||||
return await this.inner.refreshColumn(column);
|
||||
}
|
||||
|
||||
async alterColumns(
|
||||
columnAlterations: ColumnAlteration[],
|
||||
): Promise<AlterColumnsResult> {
|
||||
|
||||
@@ -347,6 +347,30 @@ impl Table {
|
||||
Ok(res.into())
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn add_computed_columns(
|
||||
&self,
|
||||
columns: Vec<AddColumnsSql>,
|
||||
) -> napi::Result<AddColumnsResult> {
|
||||
let table = self.inner_ref()?;
|
||||
let mut builder = table.add_columns();
|
||||
for column in columns {
|
||||
builder = builder.computed(column.name, column.value_sql);
|
||||
}
|
||||
let res = builder.execute().await.default_error()?;
|
||||
Ok(res.into())
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn refresh_column(&self, column: String) -> napi::Result<RefreshColumnResult> {
|
||||
let res = self
|
||||
.inner_ref()?
|
||||
.refresh_column(column)
|
||||
.await
|
||||
.default_error()?;
|
||||
Ok(res.into())
|
||||
}
|
||||
|
||||
#[napi(catch_unwind)]
|
||||
pub async fn add_columns_with_schema(
|
||||
&self,
|
||||
@@ -1196,6 +1220,21 @@ pub struct AddColumnsResult {
|
||||
pub version: i64,
|
||||
}
|
||||
|
||||
#[napi(object)]
|
||||
pub struct RefreshColumnResult {
|
||||
pub rows_filled: i64,
|
||||
pub version: i64,
|
||||
}
|
||||
|
||||
impl From<lancedb::table::RefreshColumnResult> for RefreshColumnResult {
|
||||
fn from(value: lancedb::table::RefreshColumnResult) -> Self {
|
||||
Self {
|
||||
rows_filled: value.rows_filled as i64,
|
||||
version: value.version as i64,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<lancedb::table::AddColumnsResult> for AddColumnsResult {
|
||||
fn from(value: lancedb::table::AddColumnsResult) -> Self {
|
||||
Self {
|
||||
|
||||
@@ -335,6 +335,10 @@ class Table:
|
||||
) -> list[FtsToken]: ...
|
||||
async def delete(self, filter: Union[str, PyExpr]) -> DeleteResult: ...
|
||||
async def add_columns(self, columns: list[tuple[str, str]]) -> AddColumnsResult: ...
|
||||
async def add_computed_columns(
|
||||
self, columns: list[tuple[str, str]]
|
||||
) -> AddColumnsResult: ...
|
||||
async def refresh_column(self, column: str) -> RefreshColumnResult: ...
|
||||
async def add_columns_with_schema(self, schema: pa.Schema) -> AddColumnsResult: ...
|
||||
async def alter_columns(
|
||||
self, columns: list[dict[str, Any]]
|
||||
@@ -680,6 +684,10 @@ class LsmWriteSpec:
|
||||
class AddColumnsResult:
|
||||
version: int
|
||||
|
||||
class RefreshColumnResult:
|
||||
rows_filled: int
|
||||
version: int
|
||||
|
||||
class AlterColumnsResult:
|
||||
version: int
|
||||
|
||||
|
||||
@@ -958,9 +958,21 @@ class RemoteTable(Table):
|
||||
def count_rows(self, filter: Optional[str] = None) -> int:
|
||||
return LOOP.run(self._table.count_rows(filter))
|
||||
|
||||
def add_columns(self, transforms: Dict[str, str]) -> AddColumnsResult:
|
||||
def add_columns(
|
||||
self,
|
||||
transforms: Dict[str, str] | None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
) -> AddColumnsResult:
|
||||
if computed:
|
||||
raise NotImplementedError(
|
||||
"computed columns are supported only on local tables"
|
||||
)
|
||||
return LOOP.run(self._table.add_columns(transforms))
|
||||
|
||||
def refresh_column(self, column: str):
|
||||
raise NotImplementedError("computed columns are supported only on local tables")
|
||||
|
||||
def alter_columns(
|
||||
self, *alterations: Iterable[Dict[str, str]]
|
||||
) -> AlterColumnsResult:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -176,6 +176,7 @@ if TYPE_CHECKING:
|
||||
CompactionStats,
|
||||
Tag,
|
||||
AddColumnsResult,
|
||||
RefreshColumnResult,
|
||||
AddResult,
|
||||
AlterColumnsResult,
|
||||
UpdateFieldMetadataResult,
|
||||
@@ -1916,7 +1917,14 @@ class Table(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def add_columns(
|
||||
self, transforms: Dict[str, str] | pa.Field | List[pa.Field] | pa.Schema
|
||||
self,
|
||||
transforms: Dict[str, str]
|
||||
| pa.Field
|
||||
| List[pa.Field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
):
|
||||
"""
|
||||
Add new columns with defined values.
|
||||
@@ -1930,11 +1938,68 @@ class Table(ABC):
|
||||
Alternatively, a pyarrow Field or Schema can be provided to add
|
||||
new columns with the specified data types. The new columns will
|
||||
be initialized with null values.
|
||||
computed: Dict[str, str], optional
|
||||
A map of column name to a SQL expression defining the column. The
|
||||
column's type and inputs are derived from the expression, so no
|
||||
data type is supplied.
|
||||
|
||||
Unlike ``transforms``, the expression is stored rather than
|
||||
evaluated now: the column is committed with no values, and rows get
|
||||
them from [`refresh_column`][lancedb.table.Table.refresh_column].
|
||||
Declaring one therefore costs the same on a large table as on an
|
||||
empty one.
|
||||
|
||||
A refresh does not revisit rows it has already filled, so mutating
|
||||
an input leaves the value computed at fill time; recomputing means
|
||||
dropping the column and declaring it again. While a declaration
|
||||
reads a column, that column cannot be renamed, retyped or dropped.
|
||||
|
||||
Local tables only; LanceDB Cloud and Enterprise raise
|
||||
``NotImplementedError``. Cannot be combined with ``transforms``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
AddColumnsResult
|
||||
version: the new version number of the table after adding columns.
|
||||
|
||||
Examples
|
||||
--------
|
||||
>>> import lancedb
|
||||
>>> db = lancedb.connect("./.lancedb")
|
||||
>>> table = db.create_table("computed_demo", [{"x": 1}, {"x": 2}])
|
||||
>>> table.add_columns(computed={"doubled": "x * 2"})
|
||||
AddColumnsResult(version=2)
|
||||
>>> table.refresh_column("doubled")
|
||||
RefreshColumnResult(rows_filled=2, version=3)
|
||||
>>> table.to_arrow().sort_by("x").to_pandas()
|
||||
x doubled
|
||||
0 1 2
|
||||
1 2 4
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def refresh_column(self, column: str) -> "RefreshColumnResult":
|
||||
"""
|
||||
Fill the rows of a computed column that hold no value yet.
|
||||
|
||||
Declared with ``add_columns(computed=...)``, a column starts empty and
|
||||
gets its values here. Rows appended since the last refresh are filled
|
||||
by the next one; rows already filled are left as they are, so the call
|
||||
is idempotent and does not observe a mutated input.
|
||||
|
||||
Local tables only; LanceDB Cloud and Enterprise raise
|
||||
``NotImplementedError``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
column: str
|
||||
The name of the computed column to fill.
|
||||
|
||||
Returns
|
||||
-------
|
||||
RefreshColumnResult
|
||||
rows_filled: the number of rows given a value.
|
||||
version: the new version number of the table.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
@@ -3939,9 +4004,21 @@ class LanceTable(Table):
|
||||
return LOOP.run(self._table.index_stats(index_name))
|
||||
|
||||
def add_columns(
|
||||
self, transforms: Dict[str, str] | pa.field | List[pa.field] | pa.Schema
|
||||
self,
|
||||
transforms: Dict[str, str]
|
||||
| pa.field
|
||||
| List[pa.field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: Dict[str, str] | None = None,
|
||||
) -> AddColumnsResult:
|
||||
return LOOP.run(self._table.add_columns(transforms))
|
||||
return LOOP.run(self._table.add_columns(transforms, computed=computed))
|
||||
|
||||
def refresh_column(self, column: str) -> "RefreshColumnResult":
|
||||
"""Fill a computed column's unfilled rows. See
|
||||
[`AsyncTable.refresh_column`][lancedb.AsyncTable.refresh_column]."""
|
||||
return LOOP.run(self._table.refresh_column(column))
|
||||
|
||||
def alter_columns(
|
||||
self, *alterations: Iterable[Dict[str, str]]
|
||||
@@ -5856,7 +5933,14 @@ class AsyncTable:
|
||||
return await self._inner.update(updates_sql, where)
|
||||
|
||||
async def add_columns(
|
||||
self, transforms: dict[str, str] | pa.field | List[pa.field] | pa.Schema
|
||||
self,
|
||||
transforms: dict[str, str]
|
||||
| pa.field
|
||||
| List[pa.field]
|
||||
| pa.Schema
|
||||
| None = None,
|
||||
*,
|
||||
computed: dict[str, str] | None = None,
|
||||
) -> AddColumnsResult:
|
||||
"""
|
||||
Add new columns with defined values.
|
||||
@@ -5869,6 +5953,21 @@ class AsyncTable:
|
||||
each row in the table, and can reference existing columns.
|
||||
Alternatively, you can pass a pyarrow field or schema to add
|
||||
new columns with NULLs.
|
||||
computed: Dict[str, str], optional
|
||||
A map of column name to a SQL expression defining the column. The
|
||||
column's type and inputs are derived from the expression.
|
||||
|
||||
Unlike ``transforms``, the expression is stored rather than
|
||||
evaluated now: the column is committed with no values, and rows get
|
||||
them from
|
||||
[`refresh_column`][lancedb.table.AsyncTable.refresh_column].
|
||||
|
||||
A refresh does not revisit rows it has already filled, so mutating
|
||||
an input leaves the value computed at fill time. While a
|
||||
declaration reads a column, that column cannot be renamed, retyped
|
||||
or dropped.
|
||||
|
||||
Local tables only. Cannot be combined with ``transforms``.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -5882,11 +5981,43 @@ class AsyncTable:
|
||||
{isinstance(f, pa.Field) for f in transforms}
|
||||
):
|
||||
transforms = pa.schema(transforms)
|
||||
if computed:
|
||||
if transforms:
|
||||
raise ValueError(
|
||||
"add_columns cannot take both transforms and computed columns"
|
||||
)
|
||||
return await self._inner.add_computed_columns(list(computed.items()))
|
||||
if transforms is None:
|
||||
raise ValueError("add_columns requires transforms or computed columns")
|
||||
if isinstance(transforms, pa.Schema):
|
||||
return await self._inner.add_columns_with_schema(transforms)
|
||||
else:
|
||||
return await self._inner.add_columns(list(transforms.items()))
|
||||
|
||||
async def refresh_column(self, column: str) -> RefreshColumnResult:
|
||||
"""
|
||||
Fill the rows of a computed column that hold no value yet.
|
||||
|
||||
Declared with ``add_columns(computed=...)``, a column starts empty and
|
||||
gets its values here. Rows appended since the last refresh are filled
|
||||
by the next one; rows already filled are left as they are, so the call
|
||||
is idempotent and does not observe a mutated input.
|
||||
|
||||
Local tables only; LanceDB Cloud and Enterprise raise
|
||||
``NotImplementedError``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
column: str
|
||||
The name of the computed column to fill.
|
||||
|
||||
Returns
|
||||
-------
|
||||
RefreshColumnResult
|
||||
The number of rows filled and the new version of the table.
|
||||
"""
|
||||
return await self._inner.refresh_column(column)
|
||||
|
||||
async def alter_columns(
|
||||
self, *alterations: Iterable[dict[str, Any]]
|
||||
) -> AlterColumnsResult:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -3854,3 +3854,37 @@ async def test_async_search_runs_embedding_on_dedicated_executor(
|
||||
assert all(name.startswith("lancedb-embedding") for name in captured_threads), (
|
||||
f"embedding ran off the dedicated executor: {captured_threads}"
|
||||
)
|
||||
|
||||
|
||||
def test_computed_column_declare_and_refresh(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("computed", [{"x": 1}, {"x": 2}])
|
||||
|
||||
table.add_columns(computed={"doubled": "x * 2"})
|
||||
assert table.to_arrow()["doubled"].to_pylist() == [None, None]
|
||||
|
||||
result = table.refresh_column("doubled")
|
||||
assert result.rows_filled == 2
|
||||
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4]
|
||||
|
||||
table.add([{"x": 5}])
|
||||
assert table.refresh_column("doubled").rows_filled == 1
|
||||
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4, 10]
|
||||
|
||||
|
||||
def test_computed_column_rejects_transforms_and_computed_together(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
table = db.create_table("computed_mixed", [{"x": 1}])
|
||||
with pytest.raises(ValueError):
|
||||
table.add_columns({"a": "x + 1"}, computed={"b": "x * 2"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_computed_column_async(tmp_path):
|
||||
db = await lancedb.connect_async(tmp_path)
|
||||
table = await db.create_table("computed_async", [{"x": 3}])
|
||||
|
||||
await table.add_columns(computed={"tripled": "x * 3"})
|
||||
await table.refresh_column("tripled")
|
||||
|
||||
assert (await table.to_arrow())["tripled"].to_pylist() == [9]
|
||||
|
||||
+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
|
||||
|
||||
+3
-1
@@ -16,7 +16,8 @@ use query::{FTSQuery, HybridQuery, Query, VectorQuery};
|
||||
use session::Session;
|
||||
use table::{
|
||||
AddColumnsResult, AddResult, AlterColumnsResult, DeleteResult, DropColumnsResult, FtsToken,
|
||||
LsmWriteSpec, MergeResult, PyBlobFile, Table, UpdateFieldMetadataResult, UpdateResult,
|
||||
LsmWriteSpec, MergeResult, PyBlobFile, RefreshColumnResult, Table, UpdateFieldMetadataResult,
|
||||
UpdateResult,
|
||||
};
|
||||
|
||||
pub mod arrow;
|
||||
@@ -57,6 +58,7 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
m.add_class::<VectorQuery>()?;
|
||||
m.add_class::<RecordBatchStream>()?;
|
||||
m.add_class::<AddColumnsResult>()?;
|
||||
m.add_class::<RefreshColumnResult>()?;
|
||||
m.add_class::<AlterColumnsResult>()?;
|
||||
m.add_class::<UpdateFieldMetadataResult>()?;
|
||||
m.add_class::<AddResult>()?;
|
||||
|
||||
@@ -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>,
|
||||
|
||||
+50
-1
@@ -415,6 +415,32 @@ pub struct AddColumnsResult {
|
||||
pub version: u64,
|
||||
}
|
||||
|
||||
#[pyclass(get_all, from_py_object)]
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct RefreshColumnResult {
|
||||
pub rows_filled: u64,
|
||||
pub version: u64,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl RefreshColumnResult {
|
||||
pub fn __repr__(&self) -> String {
|
||||
format!(
|
||||
"RefreshColumnResult(rows_filled={}, version={})",
|
||||
self.rows_filled, self.version
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<lancedb::table::RefreshColumnResult> for RefreshColumnResult {
|
||||
fn from(result: lancedb::table::RefreshColumnResult) -> Self {
|
||||
Self {
|
||||
rows_filled: result.rows_filled,
|
||||
version: result.version,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl AddColumnsResult {
|
||||
pub fn __repr__(&self) -> String {
|
||||
@@ -579,7 +605,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,
|
||||
@@ -1510,6 +1536,29 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
pub fn add_computed_columns(
|
||||
self_: PyRef<'_, Self>,
|
||||
columns: Vec<(String, String)>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let mut builder = inner.add_columns();
|
||||
for (name, expression) in columns {
|
||||
builder = builder.computed(name, expression);
|
||||
}
|
||||
let result = builder.execute().await.infer_error()?;
|
||||
Ok(AddColumnsResult::from(result))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn refresh_column(self_: PyRef<'_, Self>, column: String) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
future_into_py(self_.py(), async move {
|
||||
let result = inner.refresh_column(column).await.infer_error()?;
|
||||
Ok(RefreshColumnResult::from(result))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn add_columns_with_schema(
|
||||
self_: PyRef<'_, Self>,
|
||||
schema: PyArrowType<Schema>,
|
||||
|
||||
@@ -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,7 +333,10 @@ 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 {
|
||||
if matches!(
|
||||
resolved,
|
||||
ConcreteFileVersion::V1 | ConcreteFileVersion::V2_0 | ConcreteFileVersion::V2_1
|
||||
) {
|
||||
params.data_storage_version = Some(LanceFileVersion::V2_2);
|
||||
}
|
||||
}
|
||||
@@ -499,7 +502,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_2
|
||||
ConcreteFileVersion::V2_2
|
||||
);
|
||||
}
|
||||
|
||||
@@ -512,7 +515,7 @@ mod tests {
|
||||
ensure_blob_storage_version(&blob_schema(), &mut params);
|
||||
assert_eq!(
|
||||
params.data_storage_version.unwrap().resolve(),
|
||||
LanceFileVersion::V2_2
|
||||
ConcreteFileVersion::V2_2
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -438,10 +438,9 @@ mod tests {
|
||||
.await
|
||||
.unwrap()
|
||||
.data_storage_format
|
||||
.lance_file_version()
|
||||
.unwrap();
|
||||
.lance_file_format();
|
||||
// Compare resolved versions since Stable/Next are aliases that resolve at storage time
|
||||
assert_eq!(storage_format.resolve(), data_storage_version.resolve());
|
||||
assert_eq!(storage_format, data_storage_version.resolve());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -71,6 +71,14 @@ pub enum Error {
|
||||
IndexNotFound { name: String },
|
||||
#[snafu(display("Embedding function '{name}' was not found. : {reason}"))]
|
||||
EmbeddingFunctionNotFound { name: String, reason: String },
|
||||
#[snafu(display("Column '{name}' was not found"))]
|
||||
ColumnNotFound { name: String },
|
||||
#[snafu(display("Column '{name}' already exists"))]
|
||||
ColumnAlreadyExists { name: String },
|
||||
#[snafu(display("Column '{name}' is not a computed column"))]
|
||||
NotAComputedColumn { name: String },
|
||||
#[snafu(display("Invalid expression for column '{column}': {message}"))]
|
||||
InvalidExpression { column: String, message: String },
|
||||
|
||||
#[snafu(display("Table '{name}' already exists"))]
|
||||
TableAlreadyExists { name: String },
|
||||
|
||||
@@ -2706,6 +2706,13 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
// A declaration reaches here as AllNulls, which the remote protocol
|
||||
// has no representation for.
|
||||
NewColumnTransform::AllNulls(_) => {
|
||||
return Err(Error::NotSupported {
|
||||
message: "computed columns are supported only on local tables".into(),
|
||||
});
|
||||
}
|
||||
_ => {
|
||||
return Err(Error::NotSupported {
|
||||
message: "Only SQL expressions are supported for adding columns".into(),
|
||||
@@ -6455,6 +6462,37 @@ mod tests {
|
||||
assert_eq!(result.version, if old_server { 0 } else { 43 });
|
||||
}
|
||||
|
||||
/// Computed columns are local-only. Both halves say so here rather than
|
||||
/// reaching the wire and failing somewhere less legible.
|
||||
#[tokio::test]
|
||||
async fn test_computed_columns_are_refused() {
|
||||
let table = Table::new_with_handler("my_table", |request| -> http::Response<String> {
|
||||
panic!("unexpected request: {}", request.url().path())
|
||||
});
|
||||
|
||||
let declared = Arc::new(Schema::new(vec![Field::new(
|
||||
"doubled",
|
||||
DataType::Int32,
|
||||
true,
|
||||
)]));
|
||||
let err = table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::AllNulls(declared))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(
|
||||
matches!(&err, Error::NotSupported { message } if message.contains("local tables")),
|
||||
"{err:?}"
|
||||
);
|
||||
|
||||
let err = table.refresh_column("doubled").await.unwrap_err();
|
||||
assert!(
|
||||
matches!(&err, Error::NotSupported { message } if message.contains("local tables")),
|
||||
"{err:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_prewarm_index() {
|
||||
let table = Table::new_with_handler("my_table", |request| {
|
||||
|
||||
@@ -69,6 +69,7 @@ pub mod add_columns;
|
||||
mod add_data;
|
||||
pub mod branch_merge;
|
||||
pub mod checkpoint;
|
||||
pub mod computed_columns;
|
||||
mod create_index;
|
||||
pub mod datafusion;
|
||||
pub(crate) mod dataset;
|
||||
@@ -78,6 +79,7 @@ pub mod merge;
|
||||
pub mod optimize;
|
||||
mod primary_key;
|
||||
pub mod query;
|
||||
pub mod refresh;
|
||||
pub mod schema_evolution;
|
||||
pub mod update;
|
||||
pub mod write_progress;
|
||||
@@ -91,6 +93,9 @@ pub use branch_merge::{
|
||||
MergeBranchResult, MergeBranchStatus, MergePreview, RowCountSummary,
|
||||
};
|
||||
pub use chrono::Duration;
|
||||
pub use computed_columns::{
|
||||
ComputedColumn, ComputedColumnKind, computed_column_from_field, computed_columns,
|
||||
};
|
||||
pub use delete::DeleteResult;
|
||||
use futures::future::join_all;
|
||||
pub use lance::dataset::refs::{BranchContents, Ref, TagContents, Tags as LanceTags};
|
||||
@@ -98,6 +103,7 @@ pub use lance::dataset::scanner::DatasetRecordBatchStream;
|
||||
pub use lance_index::optimize::OptimizeOptions;
|
||||
pub use lsm_stats::{BucketStats, GenerationStats, LsmStats, MemtableStats};
|
||||
pub use optimize::{CompactionOptions, OptimizeAction, OptimizeStats};
|
||||
pub use refresh::RefreshColumnResult;
|
||||
pub use schema_evolution::{
|
||||
AddColumnsResult, AlterColumnsResult, DropColumnsResult, FieldMetadataUpdate,
|
||||
UpdateFieldMetadataResult,
|
||||
@@ -782,6 +788,14 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync {
|
||||
transforms: NewColumnTransform,
|
||||
read_columns: Option<Vec<String>>,
|
||||
) -> Result<AddColumnsResult>;
|
||||
/// Fill a computed column's unfilled rows.
|
||||
///
|
||||
/// The default returns `NotSupported`; Lance-backed tables override it.
|
||||
async fn refresh_column(&self, _column: &str) -> Result<RefreshColumnResult> {
|
||||
Err(Error::NotSupported {
|
||||
message: "computed columns are supported only on local tables".into(),
|
||||
})
|
||||
}
|
||||
/// Alter columns in the table.
|
||||
async fn alter_columns(&self, alterations: &[ColumnAlteration]) -> Result<AlterColumnsResult>;
|
||||
/// Drop columns from the table.
|
||||
@@ -1674,6 +1688,29 @@ impl Table {
|
||||
AddColumnsBuilder::new(self.inner.clone())
|
||||
}
|
||||
|
||||
/// Fill the fragments of a computed column that hold no values yet.
|
||||
///
|
||||
/// Declared with
|
||||
/// [`AddColumnsBuilder::computed`](add_columns::AddColumnsBuilder::computed),
|
||||
/// a column starts empty and gets its values here. Fragments appended
|
||||
/// since the last refresh are filled by the next one; fragments already
|
||||
/// filled are left as they are, so the call is idempotent and does not
|
||||
/// observe a mutated input.
|
||||
///
|
||||
/// Local tables only.
|
||||
///
|
||||
/// ```
|
||||
/// # use lancedb::Table;
|
||||
/// # async fn refresh(table: &Table) -> Result<(), Box<dyn std::error::Error>> {
|
||||
/// let result = table.refresh_column("doubled").await?;
|
||||
/// println!("filled {} rows at version {}", result.rows_filled, result.version);
|
||||
/// # Ok(())
|
||||
/// # }
|
||||
/// ```
|
||||
pub async fn refresh_column(&self, column: impl AsRef<str>) -> Result<RefreshColumnResult> {
|
||||
self.inner.refresh_column(column.as_ref()).await
|
||||
}
|
||||
|
||||
/// Change a column's name or nullability.
|
||||
pub async fn alter_columns(
|
||||
&self,
|
||||
@@ -3341,6 +3378,12 @@ impl BaseTable for NativeTable {
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
async fn refresh_column(&self, column: &str) -> Result<RefreshColumnResult> {
|
||||
let result = refresh::execute_refresh_column(self, column).await?;
|
||||
self.bump_freshness();
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
async fn alter_columns(&self, alterations: &[ColumnAlteration]) -> Result<AlterColumnsResult> {
|
||||
let result = schema_evolution::execute_alter_columns(self, alterations).await?;
|
||||
self.bump_freshness();
|
||||
@@ -5339,7 +5382,7 @@ mod tests {
|
||||
pub async fn test_stats_includes_index_and_overlay_files() {
|
||||
use lance::dataset::WriteDestination;
|
||||
use lance::dataset::transaction::{DataOverlayGroup, Operation};
|
||||
use lance_file::version::{ConcreteFileVersion, LanceFileVersion};
|
||||
use lance_file::version::stable_file_version;
|
||||
use lance_file::writer::FileWriterOptions;
|
||||
use lance_io::utils::CachedFileSize;
|
||||
use lance_table::format::DataFile;
|
||||
@@ -5405,7 +5448,7 @@ mod tests {
|
||||
let fragment_id = dataset.get_fragments()[0].id() as u64;
|
||||
let foo_field_id = dataset.schema().field("foo").unwrap().id;
|
||||
let overlay_schema = dataset.schema().project_by_ids(&[foo_field_id], true);
|
||||
let file_version = ConcreteFileVersion::from(LanceFileVersion::Stable);
|
||||
let file_version = stable_file_version();
|
||||
|
||||
let filename = "overlay.lance".to_string();
|
||||
let store = dataset.object_store(None).await.unwrap();
|
||||
|
||||
@@ -8,6 +8,7 @@ use std::sync::Arc;
|
||||
use lance::dataset::NewColumnTransform;
|
||||
|
||||
use super::BaseTable;
|
||||
use super::computed_columns;
|
||||
use super::schema_evolution::AddColumnsResult;
|
||||
use crate::{Error, Result};
|
||||
|
||||
@@ -15,6 +16,7 @@ use crate::{Error, Result};
|
||||
pub struct AddColumnsBuilder {
|
||||
parent: Arc<dyn BaseTable>,
|
||||
transform: Option<NewColumnTransform>,
|
||||
computed: Vec<(String, String)>,
|
||||
read_columns: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
@@ -23,6 +25,7 @@ impl std::fmt::Debug for AddColumnsBuilder {
|
||||
f.debug_struct("AddColumnsBuilder")
|
||||
.field("parent", &self.parent)
|
||||
.field("has_transform", &self.transform.is_some())
|
||||
.field("computed", &self.computed)
|
||||
.field("read_columns", &self.read_columns)
|
||||
.finish()
|
||||
}
|
||||
@@ -33,19 +36,57 @@ impl AddColumnsBuilder {
|
||||
Self {
|
||||
parent,
|
||||
transform: None,
|
||||
computed: Vec::new(),
|
||||
read_columns: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Set how the new columns' values are produced. Required.
|
||||
/// Set how the new columns' values are produced.
|
||||
pub fn transform(mut self, transform: NewColumnTransform) -> Self {
|
||||
self.transform = Some(transform);
|
||||
self
|
||||
}
|
||||
|
||||
/// Add a column defined by `expression`, evaluated by a later refresh
|
||||
/// rather than by this commit. Its type and inputs are derived from the
|
||||
/// expression.
|
||||
///
|
||||
/// The column is committed with no values, so declaring one costs the same
|
||||
/// on an empty table as on a large one. Rows get values from
|
||||
/// [`Table::refresh_column`](super::Table::refresh_column), which fills
|
||||
/// every fragment that has none -- including fragments appended since the
|
||||
/// last refresh.
|
||||
///
|
||||
/// Refresh does not revisit a fragment it has filled, so mutating an input
|
||||
/// leaves the value computed at fill time; recomputing means dropping the
|
||||
/// column and declaring it again. An input cannot be renamed, retyped or
|
||||
/// dropped while a declaration reads it, since the expression names it.
|
||||
///
|
||||
/// Local tables only: LanceDB Cloud and Enterprise reject a declaration
|
||||
/// with `NotSupported`.
|
||||
///
|
||||
/// ```
|
||||
/// # use lancedb::Table;
|
||||
/// # async fn declare(table: &Table) -> Result<(), Box<dyn std::error::Error>> {
|
||||
/// table
|
||||
/// .add_columns()
|
||||
/// .computed("doubled", "x * 2")
|
||||
/// .execute()
|
||||
/// .await?;
|
||||
/// let filled = table.refresh_column("doubled").await?;
|
||||
/// println!("filled {} rows", filled.rows_filled);
|
||||
/// # Ok(())
|
||||
/// # }
|
||||
/// ```
|
||||
pub fn computed(mut self, name: impl Into<String>, expression: impl Into<String>) -> Self {
|
||||
self.computed.push((name.into(), expression.into()));
|
||||
self
|
||||
}
|
||||
|
||||
/// Limit which existing columns a [`NewColumnTransform::BatchUDF`] mapper
|
||||
/// receives. Every other transform determines what it reads, so setting
|
||||
/// this alongside one is an error rather than a silent no-op.
|
||||
/// receives. Every other transform, and a computed column, determines what
|
||||
/// it reads, so setting this alongside one is an error rather than a silent
|
||||
/// no-op.
|
||||
pub fn read_columns(mut self, columns: impl IntoIterator<Item = impl Into<String>>) -> Self {
|
||||
self.read_columns = Some(columns.into_iter().map(Into::into).collect());
|
||||
self
|
||||
@@ -56,24 +97,43 @@ impl AddColumnsBuilder {
|
||||
let Self {
|
||||
parent,
|
||||
transform,
|
||||
computed,
|
||||
read_columns,
|
||||
} = self;
|
||||
|
||||
let Some(transform) = transform else {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "add_columns requires a transform".into(),
|
||||
});
|
||||
};
|
||||
|
||||
if read_columns.is_some() && !matches!(transform, NewColumnTransform::BatchUDF(_)) {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "read_columns applies only to a BatchUDF transform; \
|
||||
every other transform determines what it reads"
|
||||
match (transform, computed.is_empty()) {
|
||||
(None, true) => Err(Error::InvalidInput {
|
||||
message: "add_columns requires a transform or a computed column".into(),
|
||||
}),
|
||||
// The two commit through different transforms, so one call covering
|
||||
// both would be two commits and could half-apply.
|
||||
(Some(_), false) => Err(Error::InvalidInput {
|
||||
message: "add_columns cannot mix a transform with computed columns; \
|
||||
they cannot be added atomically in one call"
|
||||
.into(),
|
||||
});
|
||||
}),
|
||||
(Some(transform), true) => {
|
||||
if read_columns.is_some() && !matches!(transform, NewColumnTransform::BatchUDF(_)) {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "read_columns applies only to a BatchUDF transform; \
|
||||
every other transform determines what it reads"
|
||||
.into(),
|
||||
});
|
||||
}
|
||||
parent.add_columns(transform, read_columns).await
|
||||
}
|
||||
(None, false) => {
|
||||
if read_columns.is_some() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "read_columns applies only to a BatchUDF transform; \
|
||||
a computed column's inputs come from its expression"
|
||||
.into(),
|
||||
});
|
||||
}
|
||||
let transform = computed_columns::declare(parent.schema().await?, &computed)?;
|
||||
parent.add_columns(transform, None).await
|
||||
}
|
||||
}
|
||||
|
||||
parent.add_columns(transform, read_columns).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -85,8 +145,8 @@ mod tests {
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use lance::dataset::{BatchUDF, NewColumnTransform};
|
||||
|
||||
use crate::Table;
|
||||
use crate::connect;
|
||||
use crate::{Error, Table};
|
||||
|
||||
async fn table_with_two_columns(name: &str) -> Table {
|
||||
let conn = connect("memory://").execute().await.unwrap();
|
||||
@@ -98,10 +158,7 @@ mod tests {
|
||||
async fn test_requires_a_transform() {
|
||||
let table = table_with_two_columns("no_transform").await;
|
||||
let err = table.add_columns().execute().await.unwrap_err();
|
||||
assert!(
|
||||
err.to_string().contains("requires a transform"),
|
||||
"got: {err}"
|
||||
);
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -117,7 +174,7 @@ mod tests {
|
||||
.execute()
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(err.to_string().contains("BatchUDF"), "got: {err}");
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
assert!(
|
||||
@@ -126,6 +183,47 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_mixing_transform_and_computed_is_rejected() {
|
||||
let table = table_with_two_columns("mixed_add").await;
|
||||
let err = table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::SqlExpressions(vec![(
|
||||
"eager".into(),
|
||||
"x * 2".into(),
|
||||
)]))
|
||||
.computed("lazy", "x * 3")
|
||||
.execute()
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
assert!(schema.field_with_name("eager").is_err());
|
||||
assert!(schema.field_with_name("lazy").is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_read_columns_with_computed_is_rejected() {
|
||||
let table = table_with_two_columns("read_cols_computed").await;
|
||||
let err = table
|
||||
.add_columns()
|
||||
.computed("doubled", "x * 2")
|
||||
.read_columns(["x"])
|
||||
.execute()
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidInput { .. }));
|
||||
assert!(
|
||||
table
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name("doubled")
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_read_columns_limits_what_a_batch_udf_sees() {
|
||||
let table = table_with_two_columns("read_cols_udf").await;
|
||||
|
||||
@@ -0,0 +1,705 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Computed columns.
|
||||
//!
|
||||
//! A computed column is defined by a rule rather than by values supplied at
|
||||
//! write time. Declaring one commits the column carrying that rule in field
|
||||
//! metadata but no data, so the cost does not scale with the table; a later
|
||||
//! refresh fills the rows.
|
||||
//!
|
||||
//! The rule is tagged by kind ([`ComputedColumnKind`]) because kinds differ in
|
||||
//! where the column's type and inputs come from. A SQL expression is
|
||||
//! self-describing -- both are derived from the expression, so a caller writes
|
||||
//! neither -- while a kind resolved through a registry cannot be typed without
|
||||
//! consulting it. Only SQL exists today; the tag is what lets another kind be
|
||||
//! added without a second reading of the same key.
|
||||
//!
|
||||
//! [`computed_columns`] and [`computed_column_from_field`] read declarations
|
||||
//! back off a schema.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_schema::{Field as ArrowField, Schema as ArrowSchema, SchemaRef};
|
||||
use lance::dataset::NewColumnTransform;
|
||||
use lance_datafusion::planner::Planner;
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
/// Field metadata key marking a column as computed. The value is `"true"`.
|
||||
pub const COMPUTED_COLUMN_META_KEY: &str = "computed_column";
|
||||
|
||||
/// Field metadata key naming the kind of rule that defines the column.
|
||||
pub const KIND_META_KEY: &str = "computed_column.kind";
|
||||
|
||||
/// Field metadata key holding the SQL expression that defines the column.
|
||||
pub const EXPRESSION_META_KEY: &str = "computed_column.expression";
|
||||
|
||||
/// Field metadata key holding the column's inputs, as a JSON array of names.
|
||||
pub const INPUTS_META_KEY: &str = "computed_column.inputs";
|
||||
|
||||
/// Value of [`KIND_META_KEY`] for a column defined by a SQL expression.
|
||||
pub const SQL_KIND: &str = "sql";
|
||||
|
||||
/// The rule that defines a computed column's values.
|
||||
///
|
||||
/// Non-exhaustive: a kind added later is an additive change, and a caller that
|
||||
/// only handles the kinds it knows keeps compiling.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[non_exhaustive]
|
||||
pub enum ComputedColumnKind {
|
||||
/// A SQL expression evaluated by DataFusion. It is the whole definition:
|
||||
/// the column's type and its inputs are both derived from it.
|
||||
Sql {
|
||||
/// The expression.
|
||||
expression: String,
|
||||
},
|
||||
/// A kind this version does not understand, written by a newer one.
|
||||
///
|
||||
/// Reported rather than hidden so a caller can tell a column it cannot
|
||||
/// refresh apart from one that was never computed. Nothing produces this.
|
||||
Unrecognized {
|
||||
/// The kind as it was found in the metadata.
|
||||
kind: String,
|
||||
},
|
||||
}
|
||||
|
||||
/// A computed column's declaration, as read back from field metadata.
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct ComputedColumn {
|
||||
/// Name of the computed column.
|
||||
pub name: String,
|
||||
/// The rule that defines it.
|
||||
pub kind: ComputedColumnKind,
|
||||
/// Columns the rule reads, recorded at declaration time.
|
||||
///
|
||||
/// Outside the kind because every kind has inputs and the consumers that
|
||||
/// use them -- refresh planning, dependency ordering -- do not care which
|
||||
/// kind produced them. Where they come from does differ, and that is
|
||||
/// settled at declaration: derived from a SQL expression, supplied by the
|
||||
/// caller for a kind that cannot be parsed.
|
||||
pub inputs: Vec<String>,
|
||||
}
|
||||
|
||||
/// Build the field metadata recording a SQL binding.
|
||||
fn computed_column_metadata(expression: &str, inputs: &[String]) -> HashMap<String, String> {
|
||||
HashMap::from([
|
||||
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
|
||||
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
|
||||
(EXPRESSION_META_KEY.to_string(), expression.to_string()),
|
||||
(
|
||||
INPUTS_META_KEY.to_string(),
|
||||
serde_json::to_string(inputs).unwrap_or_else(|_| "[]".to_string()),
|
||||
),
|
||||
])
|
||||
}
|
||||
|
||||
/// Read a field's computed-column declaration, if it carries one.
|
||||
///
|
||||
/// A field flagged computed but carrying no kind, or a SQL one missing its
|
||||
/// expression, is not a computed column here: without the rule there is
|
||||
/// nothing to refresh from, so it is reported as absent rather than as a
|
||||
/// half-formed declaration. An unrecognized kind is different -- the rule is
|
||||
/// there and intact, this version just cannot act on it -- and comes back as
|
||||
/// [`ComputedColumnKind::Unrecognized`].
|
||||
pub fn computed_column_from_field(field: &ArrowField) -> Option<ComputedColumn> {
|
||||
let metadata = field.metadata();
|
||||
if metadata.get(COMPUTED_COLUMN_META_KEY).map(String::as_str) != Some("true") {
|
||||
return None;
|
||||
}
|
||||
let kind = match metadata.get(KIND_META_KEY)?.as_str() {
|
||||
SQL_KIND => ComputedColumnKind::Sql {
|
||||
expression: metadata.get(EXPRESSION_META_KEY)?.clone(),
|
||||
},
|
||||
other => ComputedColumnKind::Unrecognized {
|
||||
kind: other.to_string(),
|
||||
},
|
||||
};
|
||||
let inputs = metadata
|
||||
.get(INPUTS_META_KEY)
|
||||
.and_then(|raw| serde_json::from_str::<Vec<String>>(raw).ok())
|
||||
.unwrap_or_default();
|
||||
Some(ComputedColumn {
|
||||
name: field.name().clone(),
|
||||
kind,
|
||||
inputs,
|
||||
})
|
||||
}
|
||||
|
||||
/// Read every computed-column declaration carried by `schema`, in field order.
|
||||
///
|
||||
/// Introspection is a pure read of the schema the caller already holds, the
|
||||
/// way a SQL catalog reports a generation expression as another column of
|
||||
/// `information_schema.columns`.
|
||||
pub fn computed_columns(schema: &ArrowSchema) -> Vec<ComputedColumn> {
|
||||
schema
|
||||
.fields()
|
||||
.iter()
|
||||
.filter_map(|field| computed_column_from_field(field))
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Reject a schema change to a column some declaration reads.
|
||||
///
|
||||
/// A binding is SQL text naming its inputs, so renaming, retyping or dropping
|
||||
/// one leaves an expression that no longer resolves. Refusing the change keeps
|
||||
/// a declaration that survived [`plan`] evaluable for as long as it exists.
|
||||
///
|
||||
/// Paths are compared at their root: a declaration reading `metadata` is
|
||||
/// invalidated by a change to `metadata.age` just as surely.
|
||||
pub(crate) fn ensure_not_an_input(schema: &ArrowSchema, paths: &[&str]) -> Result<()> {
|
||||
let root = |path: &str| path.split('.').next().unwrap_or(path).to_string();
|
||||
for declaration in computed_columns(schema) {
|
||||
for path in paths {
|
||||
// A declaration does not read itself, so it is free to be dropped
|
||||
// or renamed along with its binding.
|
||||
if declaration.name == root(path) {
|
||||
continue;
|
||||
}
|
||||
if declaration
|
||||
.inputs
|
||||
.iter()
|
||||
.any(|input| root(input) == root(path))
|
||||
{
|
||||
return Err(Error::InvalidInput {
|
||||
message: format!(
|
||||
"column '{}' is read by computed column '{}'; drop that column first",
|
||||
path, declaration.name
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Resolve `(name, expression)` pairs against `schema` into fields carrying
|
||||
/// their bindings.
|
||||
///
|
||||
/// Everything that can be known statically is checked here rather than at
|
||||
/// refresh time: that the expression parses, that every column it reads
|
||||
/// exists, and that the target name is free. A declaration that survives this
|
||||
/// is one a refresh can always act on.
|
||||
pub(crate) fn plan(schema: SchemaRef, columns: &[(String, String)]) -> Result<Vec<ArrowField>> {
|
||||
if columns.is_empty() {
|
||||
return Err(Error::InvalidInput {
|
||||
message: "at least one computed column is required".into(),
|
||||
});
|
||||
}
|
||||
|
||||
let planner = Planner::new(schema.clone());
|
||||
let mut fields = Vec::with_capacity(columns.len());
|
||||
let mut declared: Vec<&str> = Vec::with_capacity(columns.len());
|
||||
|
||||
for (name, expression) in columns {
|
||||
if schema.field_with_name(name).is_ok() || declared.contains(&name.as_str()) {
|
||||
return Err(Error::ColumnAlreadyExists { name: name.clone() });
|
||||
}
|
||||
|
||||
let expr = planner
|
||||
.parse_expr(expression)
|
||||
.and_then(|expr| planner.optimize_expr(expr))
|
||||
.map_err(|e| Error::InvalidExpression {
|
||||
column: name.clone(),
|
||||
message: e.to_string(),
|
||||
})?;
|
||||
|
||||
let mut inputs = Planner::column_names_in_expr(&expr);
|
||||
inputs.sort();
|
||||
inputs.dedup();
|
||||
|
||||
// Resolved here rather than left to the planner so an unknown column
|
||||
// names itself in the error instead of surfacing as a plan failure.
|
||||
let mut indices = Vec::with_capacity(inputs.len());
|
||||
for input in &inputs {
|
||||
let index = schema
|
||||
.index_of(input)
|
||||
.map_err(|_| Error::InvalidExpression {
|
||||
column: name.clone(),
|
||||
message: format!("unknown column '{input}'"),
|
||||
})?;
|
||||
indices.push(index);
|
||||
}
|
||||
|
||||
// Physical expressions address columns by position, so the planner
|
||||
// that types the expression has to be built on the projected schema
|
||||
// the refresh will actually read.
|
||||
let read_schema =
|
||||
Arc::new(
|
||||
schema
|
||||
.project(&indices)
|
||||
.map_err(|e| Error::InvalidExpression {
|
||||
column: name.clone(),
|
||||
message: e.to_string(),
|
||||
})?,
|
||||
);
|
||||
let physical = Planner::new(read_schema.clone())
|
||||
.create_physical_expr(&expr)
|
||||
.map_err(|e| Error::InvalidExpression {
|
||||
column: name.clone(),
|
||||
message: e.to_string(),
|
||||
})?;
|
||||
let data_type =
|
||||
physical
|
||||
.data_type(read_schema.as_ref())
|
||||
.map_err(|e| Error::InvalidExpression {
|
||||
column: name.clone(),
|
||||
message: e.to_string(),
|
||||
})?;
|
||||
|
||||
// Declared columns start entirely null, so nullability is a property
|
||||
// of the declaration rather than of what the expression yields.
|
||||
fields.push(
|
||||
ArrowField::new(name, data_type, true)
|
||||
.with_metadata(computed_column_metadata(expression, &inputs)),
|
||||
);
|
||||
declared.push(name);
|
||||
}
|
||||
|
||||
Ok(fields)
|
||||
}
|
||||
|
||||
/// Build the transform that declares `columns` against `schema`.
|
||||
///
|
||||
/// An all-null column is how a binding with no values yet is carried into a
|
||||
/// commit; that it is spelled `AllNulls` is a detail of the commit, not of the
|
||||
/// column, which is why this is internal and
|
||||
/// [`AddColumnsBuilder::computed`](super::AddColumnsBuilder::computed) is the
|
||||
/// public way in.
|
||||
pub(crate) fn declare(
|
||||
schema: SchemaRef,
|
||||
columns: &[(String, String)],
|
||||
) -> Result<NewColumnTransform> {
|
||||
let fields = plan(schema, columns)?;
|
||||
Ok(NewColumnTransform::AllNulls(Arc::new(ArrowSchema::new(
|
||||
fields,
|
||||
))))
|
||||
}
|
||||
|
||||
/// Commit a declaration of a kind this version does not produce, the way a
|
||||
/// newer lancedb would leave one behind. Shared with the refresh tests, which
|
||||
/// need the same column to check that refresh refuses it.
|
||||
#[cfg(test)]
|
||||
pub(super) async fn add_foreign_kind(table: &crate::Table, name: &str, kind: &str) {
|
||||
use arrow_schema::DataType;
|
||||
|
||||
let field = ArrowField::new(name, DataType::Int32, true).with_metadata(HashMap::from([
|
||||
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
|
||||
(KIND_META_KEY.to_string(), kind.to_string()),
|
||||
(INPUTS_META_KEY.to_string(), r#"["x"]"#.to_string()),
|
||||
]));
|
||||
table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::AllNulls(Arc::new(ArrowSchema::new(
|
||||
vec![field],
|
||||
))))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use arrow_array::record_batch;
|
||||
use arrow_schema::DataType;
|
||||
use futures::TryStreamExt;
|
||||
use lance::dataset::ColumnAlteration;
|
||||
|
||||
use super::*;
|
||||
use crate::connect;
|
||||
use crate::query::{ExecutableQuery, QueryBase, Select};
|
||||
use crate::{Error, Table};
|
||||
|
||||
async fn table_with_ints(name: &str) -> Table {
|
||||
let conn = connect("memory://").execute().await.unwrap();
|
||||
let batch = record_batch!(("x", Int32, [1, 2, 3])).unwrap();
|
||||
conn.create_table(name, batch).execute().await.unwrap()
|
||||
}
|
||||
|
||||
/// Declare `columns` the way a caller would: plan the expressions, then
|
||||
/// add them through the ordinary column API.
|
||||
async fn add_computed(table: &Table, columns: &[(String, String)]) -> Result<u64> {
|
||||
let mut builder = table.add_columns();
|
||||
for (name, expression) in columns {
|
||||
builder = builder.computed(name, expression);
|
||||
}
|
||||
Ok(builder.execute().await?.version)
|
||||
}
|
||||
|
||||
async fn declared(table: &Table) -> Vec<ComputedColumn> {
|
||||
computed_columns(table.schema().await.unwrap().as_ref())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_declare_infers_type_and_inputs() {
|
||||
let table = table_with_ints("declare_infers").await;
|
||||
let initial = table.version().await.unwrap();
|
||||
|
||||
let version = add_computed(&table, &[("doubled".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(version > initial);
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
let field = schema.field_with_name("doubled").unwrap();
|
||||
assert_eq!(field.data_type(), &DataType::Int32);
|
||||
assert!(field.is_nullable());
|
||||
|
||||
assert_eq!(
|
||||
declared(&table).await,
|
||||
vec![ComputedColumn {
|
||||
name: "doubled".into(),
|
||||
kind: ComputedColumnKind::Sql {
|
||||
expression: "x * 2".into()
|
||||
},
|
||||
inputs: vec!["x".into()],
|
||||
}]
|
||||
);
|
||||
}
|
||||
|
||||
/// The binding reaches the schema only if `AllNulls` carries per-field
|
||||
/// metadata through the commit. The whole representation rests on it.
|
||||
#[tokio::test]
|
||||
async fn test_all_nulls_preserves_field_metadata() {
|
||||
let table = table_with_ints("metadata_survives").await;
|
||||
add_computed(&table, &[("doubled".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
let metadata = schema.field_with_name("doubled").unwrap().metadata();
|
||||
assert_eq!(
|
||||
metadata.get(COMPUTED_COLUMN_META_KEY).map(String::as_str),
|
||||
Some("true")
|
||||
);
|
||||
assert_eq!(metadata.get(KIND_META_KEY).map(String::as_str), Some("sql"));
|
||||
assert_eq!(
|
||||
metadata.get(EXPRESSION_META_KEY).map(String::as_str),
|
||||
Some("x * 2")
|
||||
);
|
||||
assert_eq!(
|
||||
metadata.get(INPUTS_META_KEY).map(String::as_str),
|
||||
Some(r#"["x"]"#)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_declared_column_is_all_null() {
|
||||
let table = table_with_ints("declare_is_null").await;
|
||||
add_computed(&table, &[("doubled".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let batches = table
|
||||
.query()
|
||||
.select(Select::columns(&["doubled"]))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let total: usize = batches.iter().map(|b| b.num_rows()).sum();
|
||||
assert_eq!(total, 3);
|
||||
for batch in &batches {
|
||||
assert_eq!(batch["doubled"].null_count(), batch.num_rows());
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_unknown_column_fails_at_declare_time() {
|
||||
let table = table_with_ints("unknown_input").await;
|
||||
let err = add_computed(&table, &[("bad".into(), "missing + 1".into())])
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "bad"));
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
assert!(schema.field_with_name("bad").is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_unparsable_expression_fails_at_declare_time() {
|
||||
let table = table_with_ints("bad_syntax").await;
|
||||
let err = add_computed(&table, &[("bad".into(), "x *".into())])
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "bad"));
|
||||
assert!(
|
||||
table
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name("bad")
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
/// A user-defined function is an expression like any other; only its
|
||||
/// resolution is missing. When a registry-aware planner exists this
|
||||
/// becomes a supported declaration rather than a new API.
|
||||
#[tokio::test]
|
||||
async fn test_unregistered_function_is_rejected_for_now() {
|
||||
let table = table_with_ints("udf_not_yet").await;
|
||||
let err = add_computed(&table, &[("vec".into(), "embed(x)".into())])
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "vec"));
|
||||
assert!(
|
||||
table
|
||||
.schema()
|
||||
.await
|
||||
.unwrap()
|
||||
.field_with_name("vec")
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_existing_column_name_is_rejected() {
|
||||
let table = table_with_ints("name_taken").await;
|
||||
let err = add_computed(&table, &[("x".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::ColumnAlreadyExists { name } if name == "x"));
|
||||
assert!(declared(&table).await.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_constant_expression_needs_no_inputs() {
|
||||
let table = table_with_ints("constant").await;
|
||||
add_computed(&table, &[("answer".into(), "42".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let declared = declared(&table).await;
|
||||
assert_eq!(declared.len(), 1);
|
||||
assert!(declared[0].inputs.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_multiple_columns_in_one_commit() {
|
||||
let table = table_with_ints("multi").await;
|
||||
let initial = table.version().await.unwrap();
|
||||
|
||||
add_computed(
|
||||
&table,
|
||||
&[
|
||||
("plus".into(), "x + 1".into()),
|
||||
("squared".into(), "x * x".into()),
|
||||
],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(table.version().await.unwrap(), initial + 1);
|
||||
let declared = declared(&table).await;
|
||||
assert_eq!(declared.len(), 2);
|
||||
assert_eq!(declared[0].name, "plus");
|
||||
assert_eq!(declared[1].name, "squared");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_duplicate_declaration_in_one_call_is_rejected() {
|
||||
let table = table_with_ints("dupe").await;
|
||||
let err = add_computed(
|
||||
&table,
|
||||
&[
|
||||
("dup".into(), "x + 1".into()),
|
||||
("dup".into(), "x + 2".into()),
|
||||
],
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::ColumnAlreadyExists { name } if name == "dup"));
|
||||
assert!(declared(&table).await.is_empty());
|
||||
}
|
||||
|
||||
/// A column added by an ordinary transform is materialized, not bound, so
|
||||
/// it carries no declaration to report.
|
||||
#[tokio::test]
|
||||
async fn test_ordinary_columns_are_not_reported_as_computed() {
|
||||
let table = table_with_ints("plain").await;
|
||||
assert!(declared(&table).await.is_empty());
|
||||
|
||||
table
|
||||
.add_columns()
|
||||
.transform(NewColumnTransform::SqlExpressions(vec![(
|
||||
"eager".into(),
|
||||
"x * 2".into(),
|
||||
)]))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(declared(&table).await.is_empty());
|
||||
}
|
||||
|
||||
/// Built-in functions type the column the same way an operator does.
|
||||
#[tokio::test]
|
||||
async fn test_builtin_function_inference() {
|
||||
let conn = connect("memory://").execute().await.unwrap();
|
||||
let batch = record_batch!(("name", Utf8, ["ada", "grace"]), ("n", Int32, [-1, 2])).unwrap();
|
||||
let table = conn
|
||||
.create_table("builtins", batch)
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
add_computed(
|
||||
&table,
|
||||
&[
|
||||
("shout".into(), "upper(name)".into()),
|
||||
("width".into(), "length(name)".into()),
|
||||
("magnitude".into(), "abs(n)".into()),
|
||||
],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let schema = table.schema().await.unwrap();
|
||||
assert_eq!(
|
||||
schema.field_with_name("shout").unwrap().data_type(),
|
||||
&DataType::Utf8
|
||||
);
|
||||
assert_eq!(
|
||||
schema.field_with_name("magnitude").unwrap().data_type(),
|
||||
&DataType::Int32
|
||||
);
|
||||
// length() returns a width-dependent integer type; assert it is one
|
||||
// rather than pinning which.
|
||||
assert!(
|
||||
schema
|
||||
.field_with_name("width")
|
||||
.unwrap()
|
||||
.data_type()
|
||||
.is_integer()
|
||||
);
|
||||
|
||||
let declared = declared(&table).await;
|
||||
assert_eq!(declared.len(), 3);
|
||||
assert_eq!(declared[0].inputs, vec!["name".to_string()]);
|
||||
assert_eq!(declared[2].inputs, vec!["n".to_string()]);
|
||||
}
|
||||
|
||||
/// The reason the kind is tagged: a declaration written by a newer version
|
||||
/// has to read back as a computed column this one cannot evaluate, not as
|
||||
/// an ordinary column. Reported as absent it would be refreshable by
|
||||
/// nothing and redeclarable over, silently.
|
||||
#[tokio::test]
|
||||
async fn test_unrecognized_kind_is_reported_rather_than_hidden() {
|
||||
let table = table_with_ints("foreign_kind").await;
|
||||
super::add_foreign_kind(&table, "embedding", "udf").await;
|
||||
|
||||
assert_eq!(
|
||||
declared(&table).await,
|
||||
vec![ComputedColumn {
|
||||
name: "embedding".into(),
|
||||
kind: ComputedColumnKind::Unrecognized { kind: "udf".into() },
|
||||
inputs: vec!["x".into()],
|
||||
}]
|
||||
);
|
||||
|
||||
let err = add_computed(&table, &[("embedding".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(err, Error::ColumnAlreadyExists { name } if name == "embedding"));
|
||||
}
|
||||
|
||||
/// A kind is what makes a declaration readable at all, so the flag alone
|
||||
/// is half-formed in the same way a missing expression is.
|
||||
#[test]
|
||||
fn test_flag_without_a_kind_is_not_a_declaration() {
|
||||
let field =
|
||||
ArrowField::new("half", DataType::Int32, true).with_metadata(HashMap::from([(
|
||||
COMPUTED_COLUMN_META_KEY.to_string(),
|
||||
"true".to_string(),
|
||||
)]));
|
||||
assert_eq!(computed_column_from_field(&field), None);
|
||||
}
|
||||
|
||||
/// A SQL declaration is its expression; without one there is nothing to
|
||||
/// refresh from.
|
||||
#[test]
|
||||
fn test_sql_kind_without_an_expression_is_not_a_declaration() {
|
||||
let field = ArrowField::new("half", DataType::Int32, true).with_metadata(HashMap::from([
|
||||
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
|
||||
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
|
||||
]));
|
||||
assert_eq!(computed_column_from_field(&field), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_inputs_are_deduplicated_and_sorted() {
|
||||
let conn = connect("memory://").execute().await.unwrap();
|
||||
let batch = record_batch!(("b", Int32, [1, 2]), ("a", Int32, [3, 4])).unwrap();
|
||||
let table = conn.create_table("dedupe", batch).execute().await.unwrap();
|
||||
|
||||
add_computed(&table, &[("total".into(), "b + a + b".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
declared(&table).await[0].inputs,
|
||||
vec!["a".to_string(), "b".to_string()]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_dropping_an_input_is_refused() {
|
||||
let table = table_with_ints("drop_input").await;
|
||||
add_computed(&table, &[("doubled".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let err = table.drop_columns(&["x"]).await.unwrap_err();
|
||||
assert!(
|
||||
matches!(&err, Error::InvalidInput { message } if message.contains("doubled")),
|
||||
"{err:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_renaming_an_input_is_refused() {
|
||||
let table = table_with_ints("rename_input").await;
|
||||
add_computed(&table, &[("doubled".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let err = table
|
||||
.alter_columns(&[ColumnAlteration::new("x".into()).rename("y".into())])
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(
|
||||
matches!(&err, Error::InvalidInput { message } if message.contains("doubled")),
|
||||
"{err:?}"
|
||||
);
|
||||
}
|
||||
|
||||
/// Nothing resolves against nullability, so it is not a rebinding.
|
||||
#[tokio::test]
|
||||
async fn test_altering_an_input_nullability_is_allowed() {
|
||||
let table = table_with_ints("nullable_input").await;
|
||||
add_computed(&table, &[("doubled".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
table
|
||||
.alter_columns(&[ColumnAlteration::new("x".into()).set_nullable(true)])
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
/// A declaration does not read itself, so it travels with its binding.
|
||||
#[tokio::test]
|
||||
async fn test_dropping_the_computed_column_is_allowed() {
|
||||
let table = table_with_ints("drop_computed").await;
|
||||
add_computed(&table, &[("doubled".into(), "x * 2".into())])
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
table.drop_columns(&["doubled"]).await.unwrap();
|
||||
assert!(declared(&table).await.is_empty());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,523 @@
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
//! Filling computed columns.
|
||||
//!
|
||||
//! A row without a value gets one; a row that has one keeps it. Refresh is
|
||||
//! therefore idempotent and does not observe input mutation -- once a row is
|
||||
//! filled, changing what the expression reads leaves the stored result alone.
|
||||
//!
|
||||
//! Convergence comes from staging nothing when nothing would change, so an
|
||||
//! expression yielding null settles after one pass rather than re-selecting
|
||||
//! the same rows forever. Fragments that already cover the column and hold no
|
||||
//! nulls are skipped without evaluating it at all.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::RecordBatch;
|
||||
use arrow_schema::Schema as ArrowSchema;
|
||||
use futures::{TryStreamExt, stream};
|
||||
use lance::Dataset;
|
||||
use lance::dataset::WriteDestination;
|
||||
use lance::dataset::fragment::FileFragment;
|
||||
use lance::dataset::transaction::Operation;
|
||||
use lance_core::ROW_ID;
|
||||
use lance_core::datatypes::Schema as LanceSchema;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::NativeTable;
|
||||
use super::computed_columns::{ComputedColumnKind, computed_column_from_field};
|
||||
use crate::{Error, Result};
|
||||
|
||||
/// Alias the expression is projected under, so its result and the column's
|
||||
/// current values can be read side by side.
|
||||
const COMPUTED_ALIAS: &str = "__lancedb_computed";
|
||||
|
||||
/// The result of refreshing a computed column.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
|
||||
pub struct RefreshColumnResult {
|
||||
/// Rows that had a value computed.
|
||||
#[serde(default)]
|
||||
pub rows_filled: u64,
|
||||
/// The commit version associated with the operation.
|
||||
#[serde(default)]
|
||||
pub version: u64,
|
||||
}
|
||||
|
||||
/// Internal implementation of the refresh logic.
|
||||
pub(crate) async fn execute_refresh_column(
|
||||
table: &NativeTable,
|
||||
column: &str,
|
||||
) -> Result<RefreshColumnResult> {
|
||||
table.dataset.ensure_mutable()?;
|
||||
let dataset = table.dataset.get().await?;
|
||||
|
||||
let expression = declared_expression(&dataset, column)?;
|
||||
let field = dataset
|
||||
.schema()
|
||||
.field(column)
|
||||
.ok_or_else(|| Error::ColumnNotFound {
|
||||
name: column.to_string(),
|
||||
})?;
|
||||
// The dataset's own field, so the identity write_column checks against the
|
||||
// manifest holds by construction.
|
||||
let column_schema = LanceSchema {
|
||||
fields: vec![field.clone()],
|
||||
metadata: Default::default(),
|
||||
};
|
||||
|
||||
let mut rows_filled = 0u64;
|
||||
let mut replacements = Vec::new();
|
||||
for fragment in fragments_to_consider(&dataset, column, field.id).await? {
|
||||
let Some((filled, values)) =
|
||||
fill_fragment(&dataset, &fragment, column, &expression).await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
rows_filled += filled;
|
||||
replacements.push(
|
||||
fragment
|
||||
.write_column(stream::iter(values.into_iter().map(Ok)), &column_schema)
|
||||
.await?,
|
||||
);
|
||||
}
|
||||
|
||||
if replacements.is_empty() {
|
||||
return Ok(RefreshColumnResult {
|
||||
rows_filled: 0,
|
||||
version: dataset.version().version,
|
||||
});
|
||||
}
|
||||
|
||||
let read_version = dataset.version().version;
|
||||
let new_dataset = Dataset::commit(
|
||||
WriteDestination::Dataset(dataset.clone()),
|
||||
Operation::DataReplacement { replacements },
|
||||
Some(read_version),
|
||||
None,
|
||||
None,
|
||||
Arc::new(Default::default()),
|
||||
false,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let version = new_dataset.version().version;
|
||||
table.dataset.update(new_dataset);
|
||||
Ok(RefreshColumnResult {
|
||||
rows_filled,
|
||||
version,
|
||||
})
|
||||
}
|
||||
|
||||
/// The SQL expression `column` is declared with.
|
||||
fn declared_expression(dataset: &Dataset, column: &str) -> Result<String> {
|
||||
let schema = ArrowSchema::from(dataset.schema());
|
||||
let field = schema
|
||||
.field_with_name(column)
|
||||
.map_err(|_| Error::ColumnNotFound {
|
||||
name: column.to_string(),
|
||||
})?;
|
||||
let declaration =
|
||||
computed_column_from_field(field).ok_or_else(|| Error::NotAComputedColumn {
|
||||
name: column.to_string(),
|
||||
})?;
|
||||
match declaration.kind {
|
||||
ComputedColumnKind::Sql { expression } => Ok(expression),
|
||||
ComputedColumnKind::Unrecognized { kind } => Err(Error::NotSupported {
|
||||
message: format!(
|
||||
"computed column '{column}' is defined by '{kind}', which this version of \
|
||||
lancedb cannot evaluate"
|
||||
),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
/// Quote `name` as a lance SQL identifier.
|
||||
///
|
||||
/// Lance's dialect delimits with backticks, so a double-quoted name would
|
||||
/// parse as a string literal rather than a column.
|
||||
fn quote_identifier(name: &str) -> String {
|
||||
format!("`{}`", name.replace('`', "``"))
|
||||
}
|
||||
|
||||
/// Fragments that could hold a row needing a value.
|
||||
///
|
||||
/// A fragment whose data files do not carry the field cannot hold one that
|
||||
/// does. One that carries it is asked, since a row rewrite -- an update, or a
|
||||
/// compaction folding an unfilled fragment into a filled one -- can leave
|
||||
/// nulls behind a covering file.
|
||||
async fn fragments_to_consider(
|
||||
dataset: &Dataset,
|
||||
column: &str,
|
||||
field_id: i32,
|
||||
) -> Result<Vec<FileFragment>> {
|
||||
let unfilled = format!("{} IS NULL", quote_identifier(column));
|
||||
let mut considered = Vec::new();
|
||||
for fragment in dataset.get_fragments() {
|
||||
let covered = fragment
|
||||
.metadata()
|
||||
.files
|
||||
.iter()
|
||||
.any(|file| file.fields.contains(&field_id));
|
||||
if !covered || fragment.count_rows(Some(unfilled.clone())).await? > 0 {
|
||||
considered.push(fragment);
|
||||
}
|
||||
}
|
||||
Ok(considered)
|
||||
}
|
||||
|
||||
/// Compute one fragment's column, keeping every value it already holds.
|
||||
///
|
||||
/// `Ok(None)` when no live row gained a value, which is what keeps a refresh
|
||||
/// from restaging a fragment whose expression yields null. Deleted rows are
|
||||
/// carried through so the values line up positionally with the fragment's data
|
||||
/// files; they are never read back, but the column file has to cover them.
|
||||
async fn fill_fragment(
|
||||
dataset: &Dataset,
|
||||
fragment: &FileFragment,
|
||||
column: &str,
|
||||
expression: &str,
|
||||
) -> Result<Option<(u64, Vec<RecordBatch>)>> {
|
||||
let mut scanner = dataset.scan();
|
||||
scanner
|
||||
.with_fragments(vec![fragment.metadata().clone()])
|
||||
.with_row_id()
|
||||
.include_deleted_rows()
|
||||
.project_with_transform(&[
|
||||
(column, quote_identifier(column).as_str()),
|
||||
(COMPUTED_ALIAS, expression),
|
||||
])?;
|
||||
|
||||
let projected = Arc::new(ArrowSchema::new(vec![
|
||||
ArrowSchema::from(dataset.schema())
|
||||
.field_with_name(column)
|
||||
.map_err(|_| Error::ColumnNotFound {
|
||||
name: column.to_string(),
|
||||
})?
|
||||
.clone(),
|
||||
]));
|
||||
|
||||
let missing = |name: &str| Error::Runtime {
|
||||
message: format!("refreshing {column} produced no {name} column"),
|
||||
};
|
||||
|
||||
let mut filled = 0u64;
|
||||
let mut values = Vec::new();
|
||||
let mut batches = scanner.try_into_stream().await?;
|
||||
while let Some(batch) = batches.try_next().await? {
|
||||
let existing = batch
|
||||
.column_by_name(column)
|
||||
.ok_or_else(|| missing(column))?;
|
||||
let computed = batch
|
||||
.column_by_name(COMPUTED_ALIAS)
|
||||
.ok_or_else(|| missing("expression"))?;
|
||||
let row_ids = batch
|
||||
.column_by_name(ROW_ID)
|
||||
.ok_or_else(|| missing(ROW_ID))?;
|
||||
|
||||
// A row is filled only if it gains a value: an expression yielding null
|
||||
// leaves it as unfilled as it was, which is what lets a refresh settle.
|
||||
// A deleted row has a null row id; its value is written but not counted.
|
||||
let unfilled = arrow::compute::is_null(existing.as_ref())?;
|
||||
filled += (0..unfilled.len())
|
||||
.filter(|i| unfilled.value(*i) && row_ids.is_valid(*i) && computed.is_valid(*i))
|
||||
.count() as u64;
|
||||
|
||||
let merged = arrow_select::zip::zip(&unfilled, computed, existing)?;
|
||||
values.push(RecordBatch::try_new(projected.clone(), vec![merged])?);
|
||||
}
|
||||
|
||||
Ok((filled > 0).then_some((filled, values)))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use arrow_array::{Int32Array, record_batch};
|
||||
use futures::TryStreamExt;
|
||||
|
||||
use crate::connect;
|
||||
use crate::query::{ExecutableQuery, QueryBase, Select};
|
||||
use crate::{Error, Result, Table};
|
||||
|
||||
async fn table_with(name: &str, values: Vec<i32>) -> Table {
|
||||
let conn = connect("memory://").execute().await.unwrap();
|
||||
let batch = record_batch!(("x", Int32, values)).unwrap();
|
||||
conn.create_table(name, batch).execute().await.unwrap()
|
||||
}
|
||||
|
||||
async fn declare_doubled(table: &Table) -> Result<u64> {
|
||||
Ok(table
|
||||
.add_columns()
|
||||
.computed("doubled", "x * 2")
|
||||
.execute()
|
||||
.await?
|
||||
.version)
|
||||
}
|
||||
|
||||
async fn read(table: &Table, column: &str) -> Vec<Option<i32>> {
|
||||
let batches = table
|
||||
.query()
|
||||
.select(Select::columns(&[column]))
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
let mut values: Vec<Option<i32>> = batches
|
||||
.iter()
|
||||
.flat_map(|batch| {
|
||||
batch[column]
|
||||
.as_any()
|
||||
.downcast_ref::<Int32Array>()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.collect();
|
||||
values.sort();
|
||||
values
|
||||
}
|
||||
|
||||
async fn append(table: &Table, values: Vec<i32>) {
|
||||
let batch = record_batch!(("x", Int32, values)).unwrap();
|
||||
table.add(batch).execute().await.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_fills_a_declared_column() {
|
||||
let table = table_with("refresh_fills", vec![1, 2, 3]).await;
|
||||
let declared = declare_doubled(&table).await.unwrap();
|
||||
assert_eq!(read(&table, "doubled").await, vec![None, None, None]);
|
||||
|
||||
let result = table.refresh_column("doubled").await.unwrap();
|
||||
assert!(result.version > declared);
|
||||
assert_eq!(result.rows_filled, 3);
|
||||
assert_eq!(
|
||||
read(&table, "doubled").await,
|
||||
vec![Some(2), Some(4), Some(6)]
|
||||
);
|
||||
}
|
||||
|
||||
/// Values written after the last refresh must be reachable by another one.
|
||||
#[tokio::test]
|
||||
async fn test_refresh_fills_rows_appended_since_the_last_refresh() {
|
||||
let table = table_with("refresh_appended", vec![1, 2]).await;
|
||||
declare_doubled(&table).await.unwrap();
|
||||
table.refresh_column("doubled").await.unwrap();
|
||||
|
||||
append(&table, vec![5, 6]).await;
|
||||
assert_eq!(
|
||||
read(&table, "doubled").await,
|
||||
vec![None, None, Some(2), Some(4)]
|
||||
);
|
||||
|
||||
let result = table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(result.rows_filled, 2);
|
||||
assert_eq!(
|
||||
read(&table, "doubled").await,
|
||||
vec![Some(2), Some(4), Some(10), Some(12)]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_with_nothing_to_fill() {
|
||||
let table = table_with("refresh_noop", vec![1, 2, 3]).await;
|
||||
declare_doubled(&table).await.unwrap();
|
||||
table.refresh_column("doubled").await.unwrap();
|
||||
|
||||
let again = table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(again.rows_filled, 0);
|
||||
assert_eq!(
|
||||
read(&table, "doubled").await,
|
||||
vec![Some(2), Some(4), Some(6)]
|
||||
);
|
||||
}
|
||||
|
||||
/// A row is filled only by gaining a value, so an expression yielding null
|
||||
/// settles at once instead of re-selecting the same rows forever. Nothing
|
||||
/// is staged, so the version does not move either.
|
||||
#[tokio::test]
|
||||
async fn test_refresh_converges_on_a_null_result() {
|
||||
let table = table_with("refresh_null_result", vec![1, 2, 3]).await;
|
||||
let declared = table
|
||||
.add_columns()
|
||||
.computed("maybe", "nullif(x, x)")
|
||||
.execute()
|
||||
.await
|
||||
.unwrap()
|
||||
.version;
|
||||
|
||||
let first = table.refresh_column("maybe").await.unwrap();
|
||||
assert_eq!(first.rows_filled, 0);
|
||||
assert_eq!(first.version, declared);
|
||||
assert_eq!(read(&table, "maybe").await, vec![None, None, None]);
|
||||
|
||||
let again = table.refresh_column("maybe").await.unwrap();
|
||||
assert_eq!(again.rows_filled, 0);
|
||||
assert_eq!(again.version, declared);
|
||||
}
|
||||
|
||||
/// The contract's boundary: a filled fragment is not revisited, so
|
||||
/// mutating an input leaves the value computed at fill time.
|
||||
#[tokio::test]
|
||||
async fn test_refresh_does_not_observe_input_mutation() {
|
||||
let table = table_with("refresh_mutation", vec![1]).await;
|
||||
declare_doubled(&table).await.unwrap();
|
||||
table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(read(&table, "doubled").await, vec![Some(2)]);
|
||||
|
||||
table.update().column("x", "3").execute().await.unwrap();
|
||||
|
||||
let again = table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(again.rows_filled, 0);
|
||||
assert_eq!(read(&table, "doubled").await, vec![Some(2)]);
|
||||
}
|
||||
|
||||
/// A row rewrite before the first refresh materializes the declared
|
||||
/// column as null behind a covering data file. Those rows are still
|
||||
/// unfilled and a later refresh has to reach them.
|
||||
#[tokio::test]
|
||||
async fn test_update_before_the_first_refresh() {
|
||||
let table = table_with("refresh_update_first", vec![1]).await;
|
||||
declare_doubled(&table).await.unwrap();
|
||||
|
||||
table.update().column("x", "3").execute().await.unwrap();
|
||||
|
||||
let result = table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(result.rows_filled, 1);
|
||||
assert_eq!(read(&table, "doubled").await, vec![Some(6)]);
|
||||
}
|
||||
|
||||
/// The contract holds row by row, not fragment by fragment: revisiting a
|
||||
/// fragment to fill one row must not recompute a filled row sitting beside
|
||||
/// it, even where the input behind it has since changed.
|
||||
#[tokio::test]
|
||||
async fn test_refresh_does_not_recompute_a_filled_row_beside_an_unfilled_one() {
|
||||
let table = table_with("refresh_mixed", vec![1, 2]).await;
|
||||
declare_doubled(&table).await.unwrap();
|
||||
table.refresh_column("doubled").await.unwrap();
|
||||
|
||||
append(&table, vec![5]).await;
|
||||
table
|
||||
.update()
|
||||
.column("x", "100")
|
||||
.only_if("x = 1")
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
table
|
||||
.optimize(crate::table::OptimizeAction::Compact {
|
||||
options: crate::table::CompactionOptions::default(),
|
||||
remap_options: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let result = table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(result.rows_filled, 1);
|
||||
// 2 is the mutated row keeping the value it was filled with, not 200.
|
||||
assert_eq!(
|
||||
read(&table, "doubled").await,
|
||||
vec![Some(2), Some(4), Some(10)]
|
||||
);
|
||||
}
|
||||
|
||||
/// Filling a fragment must not disturb the values it already holds, which
|
||||
/// is what makes a compaction-mixed fragment safe to revisit.
|
||||
#[tokio::test]
|
||||
async fn test_refresh_preserves_already_filled_rows() {
|
||||
let table = table_with("refresh_preserves", vec![1, 2]).await;
|
||||
declare_doubled(&table).await.unwrap();
|
||||
table.refresh_column("doubled").await.unwrap();
|
||||
|
||||
append(&table, vec![5]).await;
|
||||
table
|
||||
.optimize(crate::table::OptimizeAction::Compact {
|
||||
options: crate::table::CompactionOptions::default(),
|
||||
remap_options: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let result = table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(result.rows_filled, 1);
|
||||
assert_eq!(
|
||||
read(&table, "doubled").await,
|
||||
vec![Some(2), Some(4), Some(10)]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_leaves_deleted_rows_alone() {
|
||||
let table = table_with("refresh_deleted", vec![1, 2, 3, 4]).await;
|
||||
declare_doubled(&table).await.unwrap();
|
||||
table.delete("x = 2").await.unwrap();
|
||||
|
||||
let result = table.refresh_column("doubled").await.unwrap();
|
||||
assert_eq!(result.rows_filled, 3);
|
||||
assert_eq!(
|
||||
read(&table, "doubled").await,
|
||||
vec![Some(2), Some(6), Some(8)]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_a_constant_expression() {
|
||||
let table = table_with("refresh_constant", vec![1, 2, 3]).await;
|
||||
table
|
||||
.add_columns()
|
||||
.computed("answer", "42")
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let result = table.refresh_column("answer").await.unwrap();
|
||||
assert_eq!(result.rows_filled, 3);
|
||||
}
|
||||
|
||||
/// A name needing quotes reaches the evaluator intact: it is carried as a
|
||||
/// projection alias, never spliced into SQL text.
|
||||
#[tokio::test]
|
||||
async fn test_refresh_a_column_whose_name_needs_quoting() {
|
||||
let table = table_with("refresh_quoted", vec![1, 2, 3]).await;
|
||||
table
|
||||
.add_columns()
|
||||
.computed("double value", "x * 2")
|
||||
.execute()
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let result = table.refresh_column("double value").await.unwrap();
|
||||
assert_eq!(result.rows_filled, 3);
|
||||
assert_eq!(
|
||||
read(&table, "double value").await,
|
||||
vec![Some(2), Some(4), Some(6)]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_rejects_a_plain_column() {
|
||||
let table = table_with("refresh_plain", vec![1, 2, 3]).await;
|
||||
let err = table.refresh_column("x").await.unwrap_err();
|
||||
assert!(matches!(err, Error::NotAComputedColumn { name } if name == "x"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_refresh_rejects_an_unknown_column() {
|
||||
let table = table_with("refresh_missing", vec![1, 2, 3]).await;
|
||||
let err = table.refresh_column("nope").await.unwrap_err();
|
||||
assert!(matches!(err, Error::ColumnNotFound { name } if name == "nope"));
|
||||
}
|
||||
|
||||
/// A declaration of a kind this version cannot evaluate is refused by
|
||||
/// name, rather than mistaken for a plain column or fed to the SQL path.
|
||||
#[tokio::test]
|
||||
async fn test_refresh_rejects_a_kind_it_cannot_evaluate() {
|
||||
let table = table_with("refresh_foreign", vec![1, 2, 3]).await;
|
||||
super::super::computed_columns::add_foreign_kind(&table, "embedding", "udf").await;
|
||||
|
||||
let err = table.refresh_column("embedding").await.unwrap_err();
|
||||
assert!(matches!(err, Error::NotSupported { message } if message.contains("udf")));
|
||||
}
|
||||
}
|
||||
@@ -8,11 +8,13 @@
|
||||
//! - [`alter_columns`](execute_alter_columns): Rename columns, change types, or modify nullability
|
||||
//! - [`drop_columns`](execute_drop_columns): Remove columns from the table
|
||||
|
||||
use arrow_schema::Schema as ArrowSchema;
|
||||
use lance::dataset::{ColumnAlteration, NewColumnTransform};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
use super::NativeTable;
|
||||
use super::computed_columns;
|
||||
use crate::Result;
|
||||
|
||||
/// The result of an add columns operation.
|
||||
@@ -116,6 +118,14 @@ pub(crate) async fn execute_alter_columns(
|
||||
) -> Result<AlterColumnsResult> {
|
||||
table.dataset.ensure_mutable()?;
|
||||
let mut dataset = (*table.dataset.get().await?).clone();
|
||||
// Nullability is not part of what an expression resolves against, so only
|
||||
// a rename or a retype can invalidate a binding.
|
||||
let rebinding = alterations
|
||||
.iter()
|
||||
.filter(|alteration| alteration.rename.is_some() || alteration.data_type.is_some())
|
||||
.map(|alteration| alteration.path.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
computed_columns::ensure_not_an_input(&ArrowSchema::from(dataset.schema()), &rebinding)?;
|
||||
dataset.alter_columns(alterations).await?;
|
||||
let version = dataset.version().version;
|
||||
table.dataset.update(dataset);
|
||||
@@ -131,6 +141,7 @@ pub(crate) async fn execute_drop_columns(
|
||||
) -> Result<DropColumnsResult> {
|
||||
table.dataset.ensure_mutable()?;
|
||||
let mut dataset = (*table.dataset.get().await?).clone();
|
||||
computed_columns::ensure_not_an_input(&ArrowSchema::from(dataset.schema()), columns)?;
|
||||
dataset.drop_columns(columns).await?;
|
||||
let version = dataset.version().version;
|
||||
table.dataset.update(dataset);
|
||||
|
||||
@@ -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,14 @@ async fn storage_format_version(table: &Table) -> LanceFileVersion {
|
||||
.await
|
||||
.unwrap()
|
||||
.data_storage_format
|
||||
.lance_file_version()
|
||||
.unwrap()
|
||||
.resolve()
|
||||
.lance_file_format()
|
||||
}
|
||||
|
||||
fn supports_blob_v2(version: ConcreteFileVersion) -> bool {
|
||||
matches!(
|
||||
version,
|
||||
ConcreteFileVersion::V2_2 | ConcreteFileVersion::V2_3
|
||||
)
|
||||
}
|
||||
|
||||
async fn uses_stable_row_ids(table: &Table) -> bool {
|
||||
@@ -112,7 +117,7 @@ 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_format_version(&table).await));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -127,7 +132,7 @@ 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_format_version(&table).await));
|
||||
assert!(!uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -139,7 +144,7 @@ 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_format_version(&table).await));
|
||||
assert!(!uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -171,7 +176,7 @@ 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_format_version(&table).await));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
assert_eq!(table.count_rows(None).await?, 1);
|
||||
Ok(())
|
||||
@@ -281,7 +286,7 @@ 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_format_version(&table).await));
|
||||
assert!(!uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -297,7 +302,7 @@ async fn namespace_create_applies_blob_defaults() -> Result<()> {
|
||||
.execute()
|
||||
.await?;
|
||||
|
||||
assert!(storage_format_version(&table).await >= LanceFileVersion::V2_2);
|
||||
assert!(supports_blob_v2(storage_format_version(&table).await));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
Ok(())
|
||||
}
|
||||
@@ -474,7 +479,7 @@ 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_format_version(&table).await));
|
||||
assert!(uses_stable_row_ids(&table).await);
|
||||
|
||||
let ids = collect_row_ids(&table).await?;
|
||||
@@ -1305,7 +1310,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_format_version(&table).await),
|
||||
"blob v2 columns require storage >= 2.2"
|
||||
);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user