Compare commits

..

4 Commits

Author SHA1 Message Date
Will Jones 011def461c docs(python): fix cross-references that resolved to the wrong page
`mkdocs build --strict` only catches references it cannot resolve. A bare
anchor such as `[limit][]` or `[vector search][search]` is matched by
autorefs against any heading on the site, so six of them silently linked
into the JavaScript reference instead. The relative links in
`permutation.py` and `remote/errors.py` pointed at in-page anchors and
paths that do not exist.

Targets that still exist here or in an imported inventory now use
mkdocstrings references; the guide pages deleted in #2770 use their
lancedb.com URLs.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-30 16:19:38 -07:00
Will Jones ed6be12ad6 docs: clear the mkdocs warning backlog so --strict passes
`mkdocs build` emitted 61 warnings on main, and rendering the previously
undocumented classes in this PR pushed that to 158. That backlog is what
blocks turning on strict mode (#3707), so clear it here rather than leave
it worse than we found it.

Most of it was one systematic false positive: griffe cannot see the
generated `__init__` of a pydantic dataclass, so every documented
parameter looked unknown. `warn_unknown_params` turns that check off.

The rest were real docstring bugs, in 15 docstrings:

* Prose trailing a `Parameters` section is read as parameter names, which
  invented parameters called `The`, `you` and `To`. Moved into `Notes` or
  the summary.
* numpydoc only reads a type when the colon has spaces around it. Where
  the documented name is a pydantic attribute rather than a signature
  parameter, griffe has no signature to fall back on and the type was
  dropped. Affects nine embedding classes.
* `num_partitions, default sqrt(num_rows)` and friends parse as a list of
  names, rendering a bogus `default` parameter.
* One parameter indented five spaces instead of four.

`nodejs/CONTRIBUTING.md` links to the repo-root CONTRIBUTING.md, which
does not resolve once typedoc copies the file into `docs/src/js/_media/`;
an absolute URL works from both places.

`mkdocs build --strict` now exits 0.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-29 14:15:20 -07:00
Will Jones ac2b689cdb docs(python): render index/embeddings/remote/rerankers from __all__
Four packages are now rendered by a single mkdocstrings directive each,
driven by the module's `__all__`, instead of a hand-maintained list of
symbols. These were where most of the drift was: 7 of 12 rerankers and
14 of 17 embedding functions had never been listed.

`lancedb.embeddings` had no `__all__`; without one mkdocstrings renders
no members at all for a re-export package, so one is added.

AGENTS.md gains a section describing how the reference page is wired up
and how to check a docs build locally, plus a step in the "adding a new
method on Table" checklist.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-29 13:59:23 -07:00
Will Jones 4fc8114871 docs(python): add missing public APIs to the Python reference
The Python API reference page had drifted from the public API. Branch
management (`Branches` / `AsyncBranches`, which own `diff` and `merge`),
structured full-text query classes, take queries, blob helpers,
namespace connections, most rerankers and embedding functions, the
PyTorch dataloader, and several other public symbols were never listed,
so they did not appear in the rendered docs.

Also fixes docstring cross-references that pointed at guide pages which
have since moved off this site, and at unresolvable relative targets
(`[Table](Table)`, `[PyArrow Table](pyarrow.Table)`).

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-29 13:46:45 -07:00
180 changed files with 2313 additions and 55773 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
[tool.bumpversion]
current_version = "0.37.1-beta.1"
current_version = "0.37.1-beta.0"
parse = """(?x)
(?P<major>0|[1-9]\\d*)\\.
(?P<minor>0|[1-9]\\d*)\\.
-243
View File
@@ -1,243 +0,0 @@
name: Check doc links
# Checking external links is inherently noisy: third-party sites rate-limit
# automated clients, reject non-browser user agents, and go down temporarily.
# Blocking pull requests on that trades a lot of false failures for very little
# signal, so this runs on a schedule and reports findings in a single tracking
# issue instead of failing anyone's build.
on:
schedule:
- cron: "0 7 * * *"
workflow_dispatch:
# The report lives in one repository-global issue, so runs must not overlap: a
# lookup racing a create produces duplicate issues, and a healthy run closing
# the issue while a failing run only rewrites its body would leave a broken
# report closed. The group is deliberately ref-independent so that a manual
# dispatch serializes against the scheduled run.
concurrency:
group: docs-link-check
cancel-in-progress: false
permissions: {}
env:
REPORT_TITLE: "Docs link checker report"
jobs:
scan:
name: Scan links
runs-on: ubuntu-24.04
# lychee-action is pinned by SHA, but its wrapper downloads the lychee
# release tarball at run time without verifying a digest, and hands the
# resulting binary a GitHub token. Release assets remain replaceable, so
# that binary is confined to a job whose token can only read public
# content; everything that writes runs in the report job below.
permissions:
contents: read
outputs:
checker_outcome: ${{ steps.lychee.outcome }}
exit_code: ${{ steps.lychee.outputs.exit_code }}
status: ${{ steps.validate.outputs.status }}
steps:
- name: Checkout
uses: actions/checkout@v6
with:
# workflow_dispatch can run from any ref, but the report is
# repository-global. Always measure the default branch so a manual
# run from a topic branch cannot close a report that main warrants,
# or overwrite it with branch-only findings.
ref: ${{ github.event.repository.default_branch }}
persist-credentials: false
- name: Check links
id: lychee
continue-on-error: true
uses: lycheeverse/lychee-action@e7477775783ea5526144ba13e8db5eec57747ce8 # v2.9.0
with:
# Restricted to http(s) on purpose. Much of docs/src is generated
# API reference (the js/ tree comes from `npm run docs` in nodejs)
# and the hand-written pages use mkdocstrings cross-references and
# nav-relative paths that only resolve in the site mkdocs builds,
# not in this checkout, so relative links would be reported as
# broken on every run.
args: >-
--scheme https
--scheme http
--no-progress
--max-retries 3
--timeout 20
'docs/src/**/*.md'
format: json
output: ./lychee/out.json
jobSummary: false
# The report issue, not a red workflow run, is the signal for link
# findings and checker failures alike.
fail: false
- name: Validate report
id: validate
# lychee does not reserve exit code 2 for broken links: its CLI
# parser also exits 2 on an invalid option, before any link was
# checked or any report written. Only a parseable report whose
# counts agree with a completed exit code (0 or 2) counts as a link
# verdict. Everything else becomes a checker-error report instead of
# failing the workflow. Exit 2 covers timeouts as well as errors, and a
# timed-out host is exactly the transient unavailability this report
# exists to surface, so both count as findings. Requiring total > 0
# also catches a glob that silently stopped matching any file.
if: always()
env:
CHECKER_OUTCOME: ${{ steps.lychee.outcome }}
EXIT_CODE: ${{ steps.lychee.outputs.exit_code }}
run: |
status=checker-error
if [[ "$CHECKER_OUTCOME" == success ]] &&
[[ "$EXIT_CODE" == 0 || "$EXIT_CODE" == 2 ]] &&
jq -e --argjson code "$EXIT_CODE" '
(.total > 0) and
(if $code == 0
then .errors == 0 and .timeouts == 0
and (.error_map | length == 0) and (.timeout_map | length == 0)
else (.errors + .timeouts) > 0
and ((.error_map | length) + (.timeout_map | length)) > 0
end)
' ./lychee/out.json
then
if [[ "$EXIT_CODE" == 0 ]]; then
status=healthy
else
status=findings
fi
fi
echo "status=$status" >> "$GITHUB_OUTPUT"
echo "Validated link check as $status"
- name: Upload report
if: steps.validate.outputs.status == 'findings'
uses: actions/upload-artifact@v7
with:
name: link-report
path: ./lychee/out.json
retention-days: 7
report:
name: Update report issue
needs: scan
runs-on: ubuntu-24.04
# Deliberately no checkout: this job needs the report artifact and the
# issues API, not the repository contents.
permissions:
issues: write
env:
CHECKER_OUTCOME: ${{ needs.scan.outputs.checker_outcome }}
EXIT_CODE: ${{ needs.scan.outputs.exit_code }}
STATUS: ${{ needs.scan.outputs.status }}
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
steps:
- name: Find existing report issue
id: report
# Matched on title alone, and through search rather than a listing:
# the issue action applies labels in a separate call after creating the
# issue, so a label filter misses a half-created report, and this
# repository has far more open issues than one listing page holds.
# Closed issues are included because a healthy run closes the report:
# an open-only lookup would forget that identity and the next failing
# run would open a duplicate. The oldest match stays the canonical
# report and is reopened below when a problem recurs.
run: |
match=$(gh issue list --repo "$GITHUB_REPOSITORY" --state all \
--search "in:title \"$REPORT_TITLE\" author:app/github-actions" \
--limit 50 --json number,title,state \
--jq "[.[] | select(.title == \"$REPORT_TITLE\")] | sort_by(.number) | first // empty")
echo "number=$(jq -r '.number // empty' <<<"$match")" >> "$GITHUB_OUTPUT"
echo "state=$(jq -r '.state // empty' <<<"$match")" >> "$GITHUB_OUTPUT"
- name: Download report
if: env.STATUS == 'findings'
uses: actions/download-artifact@v8
with:
name: link-report
path: ./lychee
- name: Compose report
if: env.STATUS == 'findings'
run: |
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
{
echo "Broken documentation links found by [\`$GITHUB_WORKFLOW\`]($run_url)."
echo
echo "This issue is rewritten by every scheduled run and closed automatically once all links resolve."
echo
echo "Entries can be false positives: some sites rate-limit or block automated clients while working fine in a browser. Confirm before editing the docs, and add persistent offenders to \`--exclude\` in \`.github/workflows/docs-link-check.yml\`."
echo
# Timeouts are reported alongside errors: entries land in
# timeout_map with a status text instead of an HTTP code.
jq -r '
"\(.errors) of \(.total) links failed, \(.timeouts) timed out.",
"",
([(.error_map | to_entries[]), (.timeout_map | to_entries[])]
| group_by(.key)[] |
"### Errors in \(.[0].key)",
"",
(map(.value[])[] | "* [\(.status.code // .status.text // "ERR")] <\(.url)> — \(.status.details // .status.text // "unknown error")"),
"")
' ./lychee/out.json
} > ./lychee/issue.md
- name: Compose checker error report
if: env.STATUS == 'checker-error'
run: |
mkdir -p ./lychee
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
{
echo "The documentation link check did not complete in [the latest run]($run_url)."
echo
echo "This issue is rewritten by every scheduled run and closed automatically once a trustworthy run finds that all links resolve."
echo
echo "The checker did not produce a trustworthy link verdict. Treat the previous result, if any, as stale until a later run completes."
echo
echo "* Action outcome: \`$CHECKER_OUTCOME\`"
echo "* Exit code: \`${EXIT_CODE:-not reported}\`"
echo "* Verdict validation: \`failed\`"
} > ./lychee/issue.md
- name: Reopen report issue
# A healthy run closes the report, and the issue action below only
# rewrites the body of whatever number it is given. Without an
# explicit reopen, a later finding or checker error would rewrite a
# closed issue. A CLOSED state implies the lookup found a canonical
# issue, so no separate emptiness check.
if: >-
env.STATUS != 'healthy' &&
steps.report.outputs.state == 'CLOSED'
env:
ISSUE_NUMBER: ${{ steps.report.outputs.number }}
run: |
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
gh issue reopen "$ISSUE_NUMBER" --repo "$GITHUB_REPOSITORY" \
--comment "The documentation link checker reported a problem again in [the latest run]($run_url)."
- name: Report link-check problem
if: env.STATUS != 'healthy'
uses: peter-evans/create-issue-from-file@fca9117c27cdc29c6c4db3b86c48e4115a786710 # v6.0.0
with:
# Empty on the first failing run, which creates the issue; afterwards
# the same issue is updated in place.
issue-number: ${{ steps.report.outputs.number }}
title: ${{ env.REPORT_TITLE }}
content-filepath: ./lychee/issue.md
labels: documentation
- name: Close report issue once links are healthy
# An OPEN state implies the lookup found a canonical issue; a report
# that is already closed needs nothing.
if: >-
env.STATUS == 'healthy' &&
steps.report.outputs.state == 'OPEN'
env:
ISSUE_NUMBER: ${{ steps.report.outputs.number }}
run: |
run_url="$GITHUB_SERVER_URL/$GITHUB_REPOSITORY/actions/runs/$GITHUB_RUN_ID"
gh issue close "$ISSUE_NUMBER" --repo "$GITHUB_REPOSITORY" \
--comment "All documentation links resolved in [the latest run]($run_url)."
+6 -8
View File
@@ -296,18 +296,16 @@ jobs:
cargo update -p aws-types --precise 1.3.9
cargo update -p aws-sigv4 --precise 1.3.5
cargo update -p aws-credential-types --precise 1.2.8
# aws-smithy-checksums must stay at or above 0.63.13: OpenDAL's S3
# service needs crc-fast ~1.9, and older releases pin it to ~1.3.
cargo update -p aws-smithy-checksums --precise 0.63.13
cargo update -p aws-smithy-checksums --precise 0.63.9
cargo update -p aws-smithy-runtime --precise 1.9.3
cargo update -p aws-smithy-http --precise 0.62.6
cargo update -p aws-smithy-eventstream --precise 0.60.14
cargo update -p aws-smithy-http --precise 0.62.4
cargo update -p aws-smithy-eventstream --precise 0.60.12
cargo update -p aws-smithy-http-client --precise 1.1.3
cargo update -p aws-smithy-observability --precise 0.1.4
cargo update -p aws-smithy-query --precise 0.60.8
cargo update -p aws-smithy-runtime-api --precise 1.9.3
cargo update -p aws-smithy-async --precise 1.2.7
cargo update -p aws-smithy-types --precise 1.3.6
cargo update -p aws-smithy-runtime-api --precise 1.9.1
cargo update -p aws-smithy-async --precise 1.2.6
cargo update -p aws-smithy-types --precise 1.3.5
cargo update -p aws-smithy-xml --precise 0.60.11
cargo update -p home --precise 0.5.9
- name: cargo +${{ matrix.msrv }} check
Generated
+306 -327
View File
File diff suppressed because it is too large Load Diff
+15 -15
View File
@@ -13,20 +13,20 @@ categories = ["database-implementations"]
rust-version = "1.91.0"
[workspace.dependencies]
lance = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-core = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-datagen = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-file = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-io = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-index = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-linalg = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-namespace = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-namespace-impls = { version = "=11.0.0-beta.4", default-features = false, rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-table = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-testing = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-datafusion = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-encoding = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance-arrow = { version = "=11.0.0-beta.4", rev = "11007a4c36fe3a8ecf236511b4ed034c6a980d34", git = "https://github.com/lance-format/lance.git" }
lance = { "version" = "=10.0.0-beta.5", default-features = false, "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-core = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-datagen = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-file = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-io = { "version" = "=10.0.0-beta.5", default-features = false, "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-index = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-linalg = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace-impls = { "version" = "=10.0.0-beta.5", default-features = false, "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-table = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-testing = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-datafusion = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-encoding = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
lance-arrow = { "version" = "=10.0.0-beta.5", "tag" = "v10.0.0-beta.5", "git" = "https://github.com/lance-format/lance.git" }
ahash = "0.8"
# Note that this one does not include pyarrow
arrow = { version = "58.0.0", optional = false }
@@ -52,7 +52,7 @@ env_logger = "0.11"
half = { "version" = "2.7.1", default-features = false, features = [
"num-traits",
] }
futures = "0.3"
futures = "0"
log = "0.4"
metrics = "0.24"
metrics-util = "0.19"
+1 -1
View File
@@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`:
<dependency>
<groupId>com.lancedb</groupId>
<artifactId>lancedb-core</artifactId>
<version>0.37.1-beta.1</version>
<version>0.37.1-beta.0</version>
</dependency>
```
-97
View File
@@ -25,27 +25,6 @@ the underlying connection has been closed.
## Methods
### cancelJob()
```ts
abstract cancelJob(jobId): Promise<boolean>
```
Request cancellation of a server-side job by id.
Resolves to true if the server accepted the cancellation, false if no
such job exists. Cancelling an already-terminal job is a no-op success.
#### Parameters
* **jobId**: `string`
#### Returns
`Promise`&lt;`boolean`&gt;
***
### cloneTable()
```ts
@@ -386,26 +365,6 @@ Drop an existing table.
***
### getJob()
```ts
abstract getJob(jobId): Promise<null | JobDescription>
```
Describe a single server-side job by id.
Resolves to `null` when the server has no such job.
#### Parameters
* **jobId**: `string`
#### Returns
`Promise`&lt;`null` \| [`JobDescription`](../interfaces/JobDescription.md)&gt;
***
### isOpen()
```ts
@@ -420,62 +379,6 @@ Return true if the connection has not been closed
***
### job()
```ts
abstract job(jobId): Job
```
A [Job](Job.md) handle for a server-side job by id.
The handle is constructed without a server round trip; an unknown id
surfaces when the handle is used. Dropping the handle has no effect on
the job itself.
#### Parameters
* **jobId**: `string`
#### Returns
[`Job`](Job.md)
***
### jobHistory()
```ts
abstract jobHistory(jobId?): Promise<Table<any>>
```
The lifecycle event history of a server-side job, as an Arrow table.
Lists history across all jobs when `jobId` is omitted.
#### Parameters
* **jobId?**: `string`
#### Returns
`Promise`&lt;`Table`&lt;`any`&gt;&gt;
***
### listJobs()
```ts
abstract listJobs(): Promise<JobInfo[]>
```
List server-side jobs across the database's tables.
#### Returns
`Promise`&lt;[`JobInfo`](../interfaces/JobInfo.md)[]&gt;
***
### listNamespaces()
```ts
-83
View File
@@ -1,83 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / Job
# Class: Job
A handle to an operation that may still be running.
## Constructors
### new Job()
```ts
new Job(): Job
```
#### Returns
[`Job`](Job.md)
## Accessors
### id
```ts
get id(): null | string
```
Identifies the operation on the server that is running it. Operations
that run in this process have no server id. The value is opaque.
#### Returns
`null` \| `string`
## Methods
### cancel()
```ts
cancel(): Promise<void>
```
Request cancellation. Cancelling a finished operation is a no-op.
#### Returns
`Promise`&lt;`void`&gt;
***
### status()
```ts
status(): Promise<string>
```
The operation's current lifecycle state: "running", "finished",
"failed", or "cancelled".
A point snapshot; unlike [Job.wait](Job.md#wait) it does not block or reject
on a terminal failure state. States a newer server reports that this
client version does not know pass through as-is.
#### Returns
`Promise`&lt;`string`&gt;
***
### wait()
```ts
wait(): Promise<void>
```
Wait until the operation reaches a terminal state.
#### Returns
`Promise`&lt;`void`&gt;
+3 -32
View File
@@ -295,29 +295,6 @@ await table.createIndex("my_float_col");
***
### createIndexAsync()
```ts
abstract createIndexAsync(column, options?): Promise<Job>
```
Create an index, returning a handle to the indexing job.
The job may already be complete when returned; callers must not assume
the index exists until [Job.wait](Job.md#wait) resolves.
#### Parameters
* **column**: `string`
* **options?**: `Partial`&lt;[`IndexOptions`](../interfaces/IndexOptions.md)&gt;
#### Returns
`Promise`&lt;[`Job`](Job.md)&gt;
***
### currentBranch()
```ts
@@ -431,10 +408,9 @@ Read the [LsmWriteSpec](../interfaces/LsmWriteSpec.md) currently installed on th
Resolves to `undefined` when the MemWAL LSM write path is not enabled (no
spec has been set, or it was removed with [Table#unsetLsmWriteSpec](Table.md#unsetlsmwritespec)).
The returned spec mirrors what was passed to
[Table#setLsmWriteSpec](Table.md#setlsmwritespec), except that `maintainedIndexes` always
reports the concrete list resolved when the spec was set — `undefined`
never round-trips.
The returned spec — including its `maintainedIndexes` and
`writerConfigDefaults` — mirrors what was passed to
[Table#setLsmWriteSpec](Table.md#setlsmwritespec).
#### Returns
@@ -807,11 +783,6 @@ All variants require the table to have an unenforced primary key
([Table#setUnenforcedPrimaryKey](Table.md#setunenforcedprimarykey)); bucket sharding additionally
requires it to be the single column being bucketed.
Omitting `maintainedIndexes` maintains every index on the table, resolved
here, failing if one cannot be maintained — name them to install anyway.
Naming them pins an exact set, and a still-building index is rejected
rather than quietly omitted.
#### Parameters
* **spec**: [`LsmWriteSpec`](../interfaces/LsmWriteSpec.md)
-4
View File
@@ -25,7 +25,6 @@
- [Connection](classes/Connection.md)
- [HeaderProvider](classes/HeaderProvider.md)
- [Index](classes/Index.md)
- [Job](classes/Job.md)
- [MakeArrowTableOptions](classes/MakeArrowTableOptions.md)
- [MatchQuery](classes/MatchQuery.md)
- [MergeInsertBuilder](classes/MergeInsertBuilder.md)
@@ -89,9 +88,6 @@
- [IvfFlatOptions](interfaces/IvfFlatOptions.md)
- [IvfPqOptions](interfaces/IvfPqOptions.md)
- [IvfRqOptions](interfaces/IvfRqOptions.md)
- [JobDescription](interfaces/JobDescription.md)
- [JobFailureInfo](interfaces/JobFailureInfo.md)
- [JobInfo](interfaces/JobInfo.md)
- [ListNamespacesOptions](interfaces/ListNamespacesOptions.md)
- [ListNamespacesResponse](interfaces/ListNamespacesResponse.md)
- [LsmWriteSpec](interfaces/LsmWriteSpec.md)
-66
View File
@@ -1,66 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / JobDescription
# Interface: JobDescription
A described job from `Connection.getJob`.
## Properties
### creationMs
```ts
creationMs: number;
```
When the job was created, in milliseconds since the epoch.
***
### failure?
```ts
optional failure: JobFailureInfo;
```
Why the job failed, when the job is failed and the server reports a
reason.
***
### jobId
```ts
jobId: string;
```
***
### jobType
```ts
jobType: string;
```
***
### specJson?
```ts
optional specJson: string;
```
The job-type-specific specification as a JSON string, when present.
***
### state
```ts
state: string;
```
Lifecycle state: "running", "finished", "failed", or "cancelled".
-33
View File
@@ -1,33 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / JobFailureInfo
# Interface: JobFailureInfo
The server's account of why a job failed.
## Properties
### message?
```ts
optional message: string;
```
***
### phase?
```ts
optional phase: string;
```
***
### retryable?
```ts
optional retryable: boolean;
```
-58
View File
@@ -1,58 +0,0 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / JobInfo
# Interface: JobInfo
A row from `Connection.listJobs`: one server-side job.
## Properties
### createdAtMillis
```ts
createdAtMillis: number;
```
When the job was created, in milliseconds since the epoch.
***
### jobId
```ts
jobId: string;
```
The job id -- what `Connection.getJob` and `Connection.cancelJob`
accept.
***
### jobType
```ts
jobType: string;
```
***
### state
```ts
state: string;
```
Lifecycle state: "running", "finished", "failed", or "cancelled".
***
### table
```ts
table: string;
```
The table the job runs against, without URI or namespace.
+1 -3
View File
@@ -34,9 +34,7 @@ Bucket and identity variants: the sharding column.
optional maintainedIndexes: string[];
```
Indexes the MemWAL keeps up to date. Omit to maintain every supported
index, resolved on install — a snapshot, so indexes created later are not
maintained. Pass `[]` for none.
Names of indexes the MemWAL should keep up to date during writes.
***
+1 -4
View File
@@ -44,7 +44,4 @@ The number of rows in the table
totalBytes: number;
```
The total size, in bytes, of the table's data files, index files, and
overlay files
Read from the manifest, so this excludes deletion files and manifests.
The total number of bytes in the table
+1 -1
View File
@@ -31,7 +31,7 @@ is also an [asynchronous API client](#connections-asynchronous).
## Namespaces (Synchronous)
A namespace-backed connection resolves tables through a
[Lance namespace](https://lance-format.github.io/lance-namespace/) service instead of
[Lance namespace](https://lancedb.github.io/lance-namespace/) service instead of
listing a storage directory.
::: lancedb.connect_namespace
+1 -1
View File
@@ -8,7 +8,7 @@
<parent>
<groupId>com.lancedb</groupId>
<artifactId>lancedb-parent</artifactId>
<version>0.37.1-beta.1</version>
<version>0.37.1-beta.0</version>
<relativePath>../pom.xml</relativePath>
</parent>
+2 -2
View File
@@ -6,7 +6,7 @@
<groupId>com.lancedb</groupId>
<artifactId>lancedb-parent</artifactId>
<version>0.37.1-beta.1</version>
<version>0.37.1-beta.0</version>
<packaging>pom</packaging>
<name>${project.artifactId}</name>
<description>LanceDB Java SDK Parent POM</description>
@@ -28,7 +28,7 @@
<properties>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
<arrow.version>15.0.0</arrow.version>
<lance-core.version>11.0.0-beta.3</lance-core.version>
<lance-core.version>10.0.0-beta.5</lance-core.version>
<spotless.skip>false</spotless.skip>
<spotless.version>2.30.0</spotless.version>
<spotless.java.googlejavaformat.version>1.7</spotless.java.googlejavaformat.version>
+1 -1
View File
@@ -1,7 +1,7 @@
[package]
name = "lancedb-nodejs"
edition.workspace = true
version = "0.37.1-beta.1"
version = "0.37.1-beta.0"
publish = false
license.workspace = true
description.workspace = true
-144
View File
@@ -6,9 +6,7 @@ import * as arrow17 from "apache-arrow-17";
import * as arrow18 from "apache-arrow-18";
import {
Vector as CurrentVector,
convertToTable,
tableFromIPC as currentTableFromIPC,
fromBufferToRecordBatch,
fromDataToBuffer,
fromRecordBatchToBuffer,
@@ -21,7 +19,6 @@ import {
FunctionOptions,
} from "../lancedb/embedding/embedding_function";
import { EmbeddingFunctionConfig } from "../lancedb/embedding/registry";
import { sanitizeTable } from "../lancedb/sanitize";
// biome-ignore lint/suspicious/noExplicitAny: skip
function sampleRecords(): Array<Record<string, any>> {
@@ -67,11 +64,7 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
tableFromIPC,
DataType,
Dictionary,
RecordBatch: ArrowRecordBatch,
Table: ArrowTable,
Uint8: ArrowUint8,
makeData: arrowMakeData,
vectorFromArray,
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
} = <any>arrow;
type Schema = ApacheArrow["Schema"];
@@ -204,35 +197,6 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
expect(table.getChild("d")?.toJSON()).toEqual([9n, 10n, null]);
});
it("will use a provided FixedSizeList schema with typed array values", function () {
const schema = new Schema([
new Field("text", new Utf8(), false),
new Field(
"vector",
new FixedSizeList(3, new Field("item", new Float32(), false)),
false,
),
]);
const table = makeArrowTable(
[
{
text: "foo",
vector: new Float32Array([1, 2, 3]),
},
],
{ schema },
);
expect(table.getChild("text")?.toJSON()).toEqual(["foo"]);
expect(
table
.getChild("vector")
?.toJSON()
.map((value) => value.toJSON()),
).toEqual([[1, 2, 3]]);
});
it("will assume the column `vector` is FixedSizeList<Float32> by default", async function () {
const schema = new Schema([
new Field("a", new Float(Precision.DOUBLE), true),
@@ -1061,114 +1025,6 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
});
describe("when using two versions of arrow", function () {
it("preserves a dictionary shared by multiple fields", async function () {
const values = ["alpha", "beta", "alpha"];
const dictionaryVector = vectorFromArray(values);
const batch = new ArrowRecordBatch({
first: dictionaryVector.data[0],
second: dictionaryVector.data[0],
});
const table = new ArrowTable([batch]);
const sanitized = sanitizeTable(table);
expect([...sanitized.getChild("first")!]).toEqual(values);
expect([...sanitized.getChild("second")!]).toEqual(values);
const firstType = sanitized.schema.fields[0].type as {
dictionary: unknown;
};
const secondType = sanitized.schema.fields[1].type as {
dictionary: unknown;
};
expect(secondType.dictionary).toBe(firstType.dictionary);
expect(sanitized.batches[0].data.children[1].dictionary).toBe(
sanitized.batches[0].data.children[0].dictionary,
);
const buf = await fromDataToBuffer(table);
const actual = currentTableFromIPC(buf);
expect([...actual.getChild("first")!]).toEqual(values);
expect([...actual.getChild("second")!]).toEqual(values);
});
it("preserves shared dictionary data from another Arrow version", async function () {
const values = ["alpha", "beta", "alpha"];
const dictionaryVector = vectorFromArray(values);
const firstBatch = new ArrowRecordBatch({
label: dictionaryVector.slice(0, 2).data[0],
});
const secondBatch = new ArrowRecordBatch({
label: dictionaryVector.slice(2).data[0],
});
const table = new ArrowTable([firstBatch, secondBatch]);
const sanitized = sanitizeTable(table);
expect([...sanitized.getChild("label")!]).toEqual(values);
const dictionaries = sanitized.batches.map(
(batch) => batch.data.children[0].dictionary,
);
expect(dictionaries[0]).toBeInstanceOf(CurrentVector);
expect(dictionaries[1]).toBe(dictionaries[0]);
const buf = await fromDataToBuffer(table);
const actual = currentTableFromIPC(buf);
expect([...actual.getChild("label")!]).toEqual(values);
});
it("preserves shared chunks in growing dictionaries", async function () {
const type = new Dictionary(new Utf8(), new Int32(), 42, false);
const firstDictionary = vectorFromArray(["alpha", "beta"], new Utf8());
const secondDictionary = firstDictionary.concat(
vectorFromArray(["gamma"], new Utf8()),
);
const firstData = arrowMakeData({
type,
data: Int32Array.from([0, 1]),
dictionary: firstDictionary,
});
const secondData = arrowMakeData({
type,
data: Int32Array.from([2]),
dictionary: secondDictionary,
});
const table = new ArrowTable([
new ArrowRecordBatch({ label: firstData }),
new ArrowRecordBatch({ label: secondData }),
]);
const sanitized = sanitizeTable(table);
const expected = ["alpha", "beta", "gamma"];
expect([...sanitized.getChild("label")!]).toEqual(expected);
const firstLocalDictionary =
sanitized.batches[0].data.children[0].dictionary!;
const secondLocalDictionary =
sanitized.batches[1].data.children[0].dictionary!;
expect(secondLocalDictionary.data[0]).toBe(
firstLocalDictionary.data[0],
);
const buf = await fromTableToBuffer(sanitized);
const actual = currentTableFromIPC(buf);
expect([...actual.getChild("label")!]).toEqual(expected);
});
it("can serialize list data from another Arrow version", async function () {
const values = [["anime", "action"], [], null];
const vector = vectorFromArray(
values,
new List(new Field("item", new Utf8(), true)),
);
const table = new ArrowTable({ tags: vector });
const buf = await fromDataToBuffer(table);
const actual = currentTableFromIPC(buf);
const actualTags = actual.getChild("tags");
expect(actualTags?.get(0)?.toJSON()).toEqual(values[0]);
expect(actualTags?.get(1)?.toJSON()).toEqual(values[1]);
expect(actualTags?.get(2)).toBeNull();
});
it("can still import data", async function () {
const schema = new arrow15.Schema([
new arrow15.Field("id", new arrow15.Int32()),
-60
View File
@@ -11,11 +11,8 @@ import {
Float16,
Float32,
Float64,
Int32,
Schema,
Utf8,
fromDataToBuffer,
tableFromIPC,
} from "../lancedb/arrow";
import { EmbeddingFunction, LanceSchema } from "../lancedb/embedding";
import { getRegistry, register } from "../lancedb/embedding/registry";
@@ -187,63 +184,6 @@ describe("embedding functions", () => {
const vector0 = JSON.parse(JSON.stringify(arr[0].vector));
expect(vector0).toEqual([1, 2, 3]);
});
it("should append generated vectors to a non-nullable schema", async () => {
@register("non_nullable_schema_test")
class MockEmbeddingFunction extends EmbeddingFunction<string> {
ndims() {
return 3;
}
embeddingDataType(): Float {
return new Float64();
}
async computeSourceEmbeddings(data: string[]) {
return data.map(() => [1, 2, 3]);
}
}
const schema = new Schema([
new Field("id", new Int32()),
new Field("text", new Utf8()),
new Field("type", new Utf8()),
new Field(
"vector",
new FixedSizeList(3, new Field("item", new Float64())),
),
]);
const func = new MockEmbeddingFunction();
const db = await connect(tmpDir.name);
const table = await db.createEmptyTable("test_non_nullable", schema, {
embeddingFunction: {
function: func,
sourceColumn: "text",
},
});
const data = [
{ id: 1, text: "Carrot", type: "vegetable" },
{ id: 2, text: "Apple", type: "fruit" },
];
const buffer = await fromDataToBuffer(
data,
undefined,
await table.schema(),
);
const generatedTable = tableFromIPC(buffer);
const vectorField = generatedTable.schema.fields.find(
(field) => field.name === "vector",
);
expect(vectorField?.nullable).toBe(false);
await table.add(data);
const rows = await table.query().toArray();
expect(rows).toHaveLength(2);
for (const row of rows) {
expect([...row.vector]).toEqual([1, 2, 3]);
}
});
it("should error when appending to a table with an unregistered embedding function", async () => {
@register("mock")
class MockEmbeddingFunction extends EmbeddingFunction<string> {
-14
View File
@@ -1,14 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
import packageJson = require("../package.json");
describe("package metadata", () => {
it("requires Node.js type declarations compatible with the runtime", () => {
expect(packageJson.engines.node).toBe(">= 18");
expect(packageJson.peerDependencies["@types/node"]).toBe(">=18");
expect(packageJson.peerDependenciesMeta["@types/node"]).toEqual({
optional: true,
});
});
});
-75
View File
@@ -110,81 +110,6 @@ describe("Query outputSchema", () => {
});
});
describe("Search pagination", () => {
let tmpDir: tmp.DirResult;
let table: Table;
beforeEach(async () => {
tmpDir = tmp.dirSync({ unsafeCleanup: true });
const db = await connect(tmpDir.name);
const schema = new Schema([
new Field("id", new Int64(), false),
new Field("text", new Utf8(), false),
new Field(
"vector",
new FixedSizeList(2, new Field("item", new Float32())),
false,
),
]);
const data = makeArrowTable(
[
{ id: 1n, text: "common", vector: [0, 0] },
{ id: 2n, text: "common common", vector: [1, 1] },
{ id: 3n, text: "common common common", vector: [2, 2] },
{ id: 4n, text: "common common common common", vector: [3, 3] },
],
{ schema },
);
table = await db.createTable("test", data);
});
afterEach(() => {
tmpDir.removeCallback();
});
it("applies offset after the vector search limit", async () => {
const allResults = await table
.vectorSearch([0, 0])
.select(["id"])
.limit(4)
.toArray();
const secondPage = await table
.vectorSearch([0, 0])
.select(["id"])
.limit(2)
.offset(2)
.toArray();
expect(allResults).toHaveLength(4);
expect(secondPage).toHaveLength(2);
expect(secondPage.map((row) => row.id)).toEqual(
allResults.slice(2, 4).map((row) => row.id),
);
});
it("applies offset after the full-text search limit", async () => {
await table.createIndex("text", { config: Index.fts() });
const allResults = await table
.search("common", "fts")
.select(["id"])
.limit(4)
.toArray();
const secondPage = await table
.search("common", "fts")
.select(["id"])
.limit(2)
.offset(2)
.toArray();
expect(allResults).toHaveLength(4);
expect(secondPage).toHaveLength(2);
expect(secondPage.map((row) => row.id)).toEqual(
allResults.slice(2, 4).map((row) => row.id),
);
});
});
describe("Query orderBy", () => {
let tmpDir: tmp.DirResult;
let table: Table;
-125
View File
@@ -170,38 +170,6 @@ describe("remote connection", () => {
);
});
it("surfaces JSON server errors from remote table operations", async () => {
await withMockDatabase(
(req, res) => {
const path = req.url ?? "";
if (path.endsWith("/describe/")) {
res.writeHead(200, { "Content-Type": "application/json" }).end(
JSON.stringify({
name: "broken_table",
version: 1,
schema: { fields: [] },
}),
);
return;
}
if (path.endsWith("/count_rows/")) {
res
.writeHead(400, { "Content-Type": "application/json" })
.end(JSON.stringify({ error: "count rows failed" }));
return;
}
res.writeHead(404).end();
},
async (db) => {
const table = await db.openTable("broken_table");
await expect(table.countRows()).rejects.toThrow("count rows failed");
},
);
});
it("should pass on requested extra headers", async () => {
await withMockDatabase(
(req, res) => {
@@ -909,96 +877,3 @@ describe("remote connection", () => {
});
});
});
describe("remote connection jobs surface", () => {
it("lists, describes, cancels, and reads history", async () => {
const { tableFromArrays, tableToIPC } = await import("apache-arrow");
const eventsTable = tableFromArrays({ state: ["created", "succeeded"] });
const eventsBody = Buffer.from(tableToIPC(eventsTable, "stream"));
await withMockDatabase(
(req, res) => {
let body = "";
req.on("data", (chunk) => {
body += chunk;
});
req.on("end", () => {
const payload = body.length > 0 ? JSON.parse(body) : {};
if (req.url === "/v1/jobs/list") {
if (payload["page_token"] === undefined) {
res
.writeHead(200, { "Content-Type": "application/json" })
.end(
'{"jobs": [{"job_id": "job-1", "table": "t1", ' +
'"job_type": "create_index", "state": "in_progress", ' +
'"created_at_millis": 1000}], "page_token": "next"}',
);
} else {
res
.writeHead(200, { "Content-Type": "application/json" })
.end(
'{"jobs": [{"job_id": "job-2", "table": "t2", ' +
'"job_type": "create_index", "state": "succeeded", ' +
'"created_at_millis": 2000}]}',
);
}
} else if (req.url === "/v1/jobs/describe") {
if (payload["job_id"] !== "job-1") {
res.writeHead(404).end("no such job");
return;
}
res
.writeHead(200, { "Content-Type": "application/json" })
.end(
'{"job_id": "job-1", "job_type": "create_index", ' +
'"job_state": "FAILED", "creation_ms": 1000, ' +
'"spec": {"column": "vec"}, "failure": {"phase": "execute", ' +
'"message": "worker died", "retryable": true}}',
);
} else if (req.url === "/v1/jobs/cancel") {
if (payload["job_id"] !== "job-1") {
res.writeHead(404).end("no such job");
return;
}
res
.writeHead(200, { "Content-Type": "application/json" })
.end('{"job_id": "job-1"}');
} else if (req.url === "/v1/jobs/query_events") {
res
.writeHead(200, {
"Content-Type": "application/vnd.apache.arrow.stream",
})
.end(eventsBody);
} else {
res.writeHead(404).end();
}
});
},
async (db) => {
const jobs = await db.listJobs();
expect(jobs.map((job) => job.jobId)).toEqual(["job-1", "job-2"]);
expect(jobs[0].state).toEqual("running");
expect(jobs[1].state).toEqual("finished");
const description = await db.getJob("job-1");
expect(description?.state).toEqual("failed");
expect(JSON.parse(description?.specJson ?? "")).toEqual({
column: "vec",
});
expect(description?.failure?.message).toEqual("worker died");
expect(await db.getJob("missing")).toBeNull();
expect(await db.cancelJob("job-1")).toBe(true);
expect(await db.cancelJob("missing")).toBe(false);
const history = await db.jobHistory("job-1");
expect(history.numRows).toEqual(2);
const job = db.job("job-1");
expect(job.id).toEqual("job-1");
expect(await job.status()).toEqual("failed");
await expect(job.wait()).rejects.toThrow("worker died");
},
);
});
});
+2 -52
View File
@@ -86,44 +86,6 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
await expect(table.countRows()).resolves.toBe(3);
});
it("should support a foreign Float64 vector schema end to end", async () => {
const conn = await connect(tmpDir.name);
const schema = new arrow.Schema([
new arrow.Field("resource_id", new arrow.Int32(), false),
new arrow.Field(
"vector",
new arrow.FixedSizeList(
3,
new arrow.Field("value", new arrow.Float64(), true),
),
false,
),
]);
const data = [
{
// biome-ignore lint/style/useNamingConvention: matches the reported schema
resource_id: 0,
vector: [0.1, 0.1, 0.1],
},
];
const resources = await conn.createTable("resources", data, { schema });
const existing = await resources
.query()
.where("resource_id = 0")
.limit(1)
.toArray();
expect(existing).toHaveLength(1);
const matched = await resources
.search(Float64Array.from(data[0].vector))
.limit(1)
.toArray();
expect(matched).toHaveLength(1);
expect(matched[0]["resource_id"]).toBe(0);
});
it("should support branches", async () => {
await table.add([{ id: 1 }]);
expect(await table.countRows()).toBe(1);
@@ -277,16 +239,8 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
},
numIndices: 0,
numRows: 3,
// Full on-disk size of the two data files, footers and metadata included.
totalBytes: 684,
totalBytes: 44,
});
// Index files count toward totalBytes too (only deletion files and
// manifests are excluded).
await table.createIndex("id", { config: Index.btree() });
const statsWithIndex = await table.stats();
expect(statsWithIndex.numIndices).toBe(1);
expect(statsWithIndex.totalBytes).toBeGreaterThan(684);
});
it("should overwrite data if asked", async () => {
@@ -897,11 +851,7 @@ describe("When creating an index", () => {
afterEach(() => tmpDir.removeCallback());
it("should create a vector index on vector columns", async () => {
const job = await tbl.createIndexAsync("vec");
expect(job.id).toBeNull();
await job.wait();
// Cancelling a job that already finished succeeds and does nothing.
await job.cancel();
await tbl.createIndex("vec");
// check index directory
const indexDir = path.join(tmpDir.name, "test.lance", "_indices");
-62
View File
@@ -1,7 +1,6 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
import { tableFromIPC } from "apache-arrow";
import {
Data,
SchemaLike,
@@ -21,9 +20,6 @@ import type {
CreateNamespaceResponse,
DescribeNamespaceResponse,
DropNamespaceResponse,
Job,
JobDescription,
JobInfo,
ListNamespacesResponse,
} from "./native";
export type {
@@ -440,40 +436,6 @@ export abstract class Connection {
newName: string,
options?: RenameTableOptions,
): Promise<void>;
/**
* A {@link Job} handle for a server-side job by id.
*
* The handle is constructed without a server round trip; an unknown id
* surfaces when the handle is used. Dropping the handle has no effect on
* the job itself.
*/
abstract job(jobId: string): Job;
/** List server-side jobs across the database's tables. */
abstract listJobs(): Promise<JobInfo[]>;
/**
* Describe a single server-side job by id.
*
* Resolves to `null` when the server has no such job.
*/
abstract getJob(jobId: string): Promise<JobDescription | null>;
/**
* Request cancellation of a server-side job by id.
*
* Resolves to true if the server accepted the cancellation, false if no
* such job exists. Cancelling an already-terminal job is a no-op success.
*/
abstract cancelJob(jobId: string): Promise<boolean>;
/**
* The lifecycle event history of a server-side job, as an Arrow table.
*
* Lists history across all jobs when `jobId` is omitted.
*/
abstract jobHistory(jobId?: string): Promise<ArrowTable>;
}
/** @hideconstructor */
@@ -760,30 +722,6 @@ export class LocalConnection extends Connection {
options?.newNamespacePath,
);
}
job(jobId: string): Job {
return this.inner.job(jobId);
}
async listJobs(): Promise<JobInfo[]> {
return this.inner.listJobs();
}
async getJob(jobId: string): Promise<JobDescription | null> {
return this.inner.getJob(jobId);
}
async cancelJob(jobId: string): Promise<boolean> {
return this.inner.cancelJob(jobId);
}
async jobHistory(jobId?: string): Promise<ArrowTable> {
const buf = await this.inner.jobHistory(jobId);
if (buf.length === 0) {
return new ArrowTable();
}
return tableFromIPC(buf);
}
}
/**
+1 -7
View File
@@ -85,13 +85,7 @@ export {
RenameTableOptions,
} from "./connection";
export {
Job,
JobDescription,
JobFailureInfo,
JobInfo,
Session,
} from "./native.js";
export { Session } from "./native.js";
export {
ExecutableQuery,
+29 -174
View File
@@ -9,7 +9,7 @@
// comes from the exact same library instance. This is not always the case
// and so we must sanitize the input to ensure that it is compatible.
import { BufferType, Data, Vector } from "apache-arrow";
import { BufferType, Data } from "apache-arrow";
import type { IntBitWidth, TKeys, TimeBitWidth } from "apache-arrow/type";
import {
Binary,
@@ -74,20 +74,6 @@ import {
Utf8,
} from "./arrow";
type SanitizationContext = {
types: WeakMap<object, DataType>;
vectors: WeakMap<object, Vector>;
data: WeakMap<object, Data<DataType>>;
};
function createSanitizationContext(): SanitizationContext {
return {
types: new WeakMap(),
vectors: new WeakMap(),
data: new WeakMap(),
};
}
export function sanitizeMetadata(
metadataLike?: unknown,
): Map<string, string> | undefined {
@@ -200,13 +186,6 @@ export function sanitizeInterval(typeLike: object) {
}
export function sanitizeList(typeLike: object) {
return sanitizeListWithContext(typeLike, createSanitizationContext());
}
function sanitizeListWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
throw Error(
"Expected a List type to have an array-like `children` property",
@@ -215,35 +194,19 @@ function sanitizeListWithContext(
if (typeLike.children.length !== 1) {
throw Error("Expected a List type to have exactly one child");
}
return new List(sanitizeFieldWithContext(typeLike.children[0], context));
return new List(sanitizeField(typeLike.children[0]));
}
export function sanitizeStruct(typeLike: object) {
return sanitizeStructWithContext(typeLike, createSanitizationContext());
}
function sanitizeStructWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
throw Error(
"Expected a Struct type to have an array-like `children` property",
);
}
return new Struct(
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
);
return new Struct(typeLike.children.map((child) => sanitizeField(child)));
}
export function sanitizeUnion(typeLike: object) {
return sanitizeUnionWithContext(typeLike, createSanitizationContext());
}
function sanitizeUnionWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (
!("typeIds" in typeLike) ||
!("mode" in typeLike) ||
@@ -263,7 +226,7 @@ function sanitizeUnionWithContext(
typeLike.mode,
// biome-ignore lint/suspicious/noExplicitAny: skip
typeLike.typeIds as any,
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
typeLike.children.map((child) => sanitizeField(child)),
);
}
@@ -271,19 +234,6 @@ export function sanitizeTypedUnion(
typeLike: object,
// eslint-disable-next-line @typescript-eslint/naming-convention
UnionType: typeof DenseUnion | typeof SparseUnion,
) {
return sanitizeTypedUnionWithContext(
typeLike,
UnionType,
createSanitizationContext(),
);
}
function sanitizeTypedUnionWithContext(
typeLike: object,
// eslint-disable-next-line @typescript-eslint/naming-convention
UnionType: typeof DenseUnion | typeof SparseUnion,
context: SanitizationContext,
) {
if (!("typeIds" in typeLike)) {
throw Error(
@@ -298,7 +248,7 @@ function sanitizeTypedUnionWithContext(
return new UnionType(
typeLike.typeIds as Int32Array | number[],
typeLike.children.map((child) => sanitizeFieldWithContext(child, context)),
typeLike.children.map((child) => sanitizeField(child)),
);
}
@@ -312,16 +262,6 @@ export function sanitizeFixedSizeBinary(typeLike: object) {
}
export function sanitizeFixedSizeList(typeLike: object) {
return sanitizeFixedSizeListWithContext(
typeLike,
createSanitizationContext(),
);
}
function sanitizeFixedSizeListWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("listSize" in typeLike) || typeof typeLike.listSize !== "number") {
throw Error("Expected a FixedSizeList type to have a `listSize` property");
}
@@ -335,18 +275,11 @@ function sanitizeFixedSizeListWithContext(
}
return new FixedSizeList(
typeLike.listSize,
sanitizeFieldWithContext(typeLike.children[0], context),
sanitizeField(typeLike.children[0]),
);
}
export function sanitizeMap(typeLike: object) {
return sanitizeMapWithContext(typeLike, createSanitizationContext());
}
function sanitizeMapWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("children" in typeLike) || !Array.isArray(typeLike.children)) {
throw Error(
"Expected a Map type to have an array-like `children` property",
@@ -359,10 +292,7 @@ function sanitizeMapWithContext(
throw Error("Expected a Map type to have exactly one child");
}
return new Map_(
sanitizeFieldWithContext(typeLike.children[0], context),
typeLike.keysSorted,
);
return new Map_(sanitizeField(typeLike.children[0]), typeLike.keysSorted);
}
export function sanitizeDuration(typeLike: object) {
@@ -373,13 +303,6 @@ export function sanitizeDuration(typeLike: object) {
}
export function sanitizeDictionary(typeLike: object) {
return sanitizeDictionaryWithContext(typeLike, createSanitizationContext());
}
function sanitizeDictionaryWithContext(
typeLike: object,
context: SanitizationContext,
) {
if (!("id" in typeLike) || typeof typeLike.id !== "number") {
throw Error("Expected a Dictionary type to have an `id` property");
}
@@ -393,8 +316,8 @@ function sanitizeDictionaryWithContext(
throw Error("Expected a Dictionary type to have an `isOrdered` property");
}
return new Dictionary(
sanitizeTypeWithContext(typeLike.dictionary, context),
sanitizeTypeWithContext(typeLike.indices, context) as TKeys,
sanitizeType(typeLike.dictionary),
sanitizeType(typeLike.indices) as TKeys,
typeLike.id,
typeLike.isOrdered,
);
@@ -402,23 +325,12 @@ function sanitizeDictionaryWithContext(
// biome-ignore lint/suspicious/noExplicitAny: skip
export function sanitizeType(typeLike: unknown): DataType<any> {
return sanitizeTypeWithContext(typeLike, createSanitizationContext());
}
function sanitizeTypeWithContext(
typeLike: unknown,
context: SanitizationContext,
): DataType {
if (typeof typeLike === "string") {
return dataTypeFromName(typeLike);
}
if (typeof typeLike !== "object" || typeLike === null) {
throw Error("Expected a Type but object was null/undefined");
}
const cached = context.types.get(typeLike);
if (cached !== undefined) {
return cached;
}
if (
!("typeId" in typeLike) ||
!(
@@ -437,16 +349,6 @@ function sanitizeTypeWithContext(
throw Error("Type's typeId property was not a function or number");
}
const type = sanitizeTypeById(typeLike, typeId, context);
context.types.set(typeLike, type);
return type;
}
function sanitizeTypeById(
typeLike: object,
typeId: Type,
context: SanitizationContext,
): DataType {
switch (typeId) {
case Type.NONE:
throw Error("Received a Type with a typeId of NONE");
@@ -473,21 +375,21 @@ function sanitizeTypeById(
case Type.Interval:
return sanitizeInterval(typeLike);
case Type.List:
return sanitizeListWithContext(typeLike, context);
return sanitizeList(typeLike);
case Type.Struct:
return sanitizeStructWithContext(typeLike, context);
return sanitizeStruct(typeLike);
case Type.Union:
return sanitizeUnionWithContext(typeLike, context);
return sanitizeUnion(typeLike);
case Type.FixedSizeBinary:
return sanitizeFixedSizeBinary(typeLike);
case Type.FixedSizeList:
return sanitizeFixedSizeListWithContext(typeLike, context);
return sanitizeFixedSizeList(typeLike);
case Type.Map:
return sanitizeMapWithContext(typeLike, context);
return sanitizeMap(typeLike);
case Type.Duration:
return sanitizeDuration(typeLike);
case Type.Dictionary:
return sanitizeDictionaryWithContext(typeLike, context);
return sanitizeDictionary(typeLike);
case Type.Int8:
return new Int8();
case Type.Int16:
@@ -531,9 +433,9 @@ function sanitizeTypeById(
case Type.TimestampSecond:
return sanitizeTypedTimestamp(typeLike, TimestampSecond);
case Type.DenseUnion:
return sanitizeTypedUnionWithContext(typeLike, DenseUnion, context);
return sanitizeTypedUnion(typeLike, DenseUnion);
case Type.SparseUnion:
return sanitizeTypedUnionWithContext(typeLike, SparseUnion, context);
return sanitizeTypedUnion(typeLike, SparseUnion);
case Type.IntervalDayTime:
return new IntervalDayTime();
case Type.IntervalYearMonth:
@@ -552,13 +454,6 @@ function sanitizeTypeById(
}
export function sanitizeField(fieldLike: unknown): Field {
return sanitizeFieldWithContext(fieldLike, createSanitizationContext());
}
function sanitizeFieldWithContext(
fieldLike: unknown,
context: SanitizationContext,
): Field {
if (fieldLike instanceof Field) {
return fieldLike;
}
@@ -576,7 +471,7 @@ function sanitizeFieldWithContext(
}
let type: DataType;
try {
type = sanitizeTypeWithContext(fieldLike.type, context);
type = sanitizeType(fieldLike.type);
} catch (error: unknown) {
throw Error(
`Unable to sanitize type for field: ${fieldLike.name} due to error: ${error}`,
@@ -606,13 +501,6 @@ function sanitizeFieldWithContext(
* than lancedb is using.
*/
export function sanitizeSchema(schemaLike: SchemaLike): Schema {
return sanitizeSchemaWithContext(schemaLike, createSanitizationContext());
}
function sanitizeSchemaWithContext(
schemaLike: SchemaLike,
context: SanitizationContext,
): Schema {
if (schemaLike instanceof Schema) {
return schemaLike;
}
@@ -634,7 +522,7 @@ function sanitizeSchemaWithContext(
);
}
const sanitizedFields = schemaLike.fields.map((field) =>
sanitizeFieldWithContext(field, context),
sanitizeField(field),
);
return new Schema(sanitizedFields, metadata);
}
@@ -656,18 +544,13 @@ export function sanitizeTable(tableLike: TableLike): Table {
"The table passed in does not appear to be a table (no 'columns' property)",
);
}
const context = createSanitizationContext();
const schema = sanitizeSchemaWithContext(tableLike.schema, context);
const batches = tableLike.batches.map((batch) =>
sanitizeRecordBatch(batch, context),
);
const schema = sanitizeSchema(tableLike.schema);
const batches = tableLike.batches.map(sanitizeRecordBatch);
return new Table(schema, batches);
}
function sanitizeRecordBatch(
batchLike: RecordBatchLike,
context: SanitizationContext,
): RecordBatch {
function sanitizeRecordBatch(batchLike: RecordBatchLike): RecordBatch {
if (batchLike instanceof RecordBatch) {
return batchLike;
}
@@ -684,43 +567,19 @@ function sanitizeRecordBatch(
"The record batch passed in does not appear to be a record batch (no 'data' property)",
);
}
const schema = sanitizeSchemaWithContext(batchLike.schema, context);
const data = sanitizeData(batchLike.data, context) as Data<Struct>;
const schema = sanitizeSchema(batchLike.schema);
const data = sanitizeData(batchLike.data);
return new RecordBatch(schema, data);
}
type DictionaryVectorLike = {
data: readonly DataLike[];
};
type DictionaryDataLike = DataLike & {
dictionary?: DictionaryVectorLike;
};
function sanitizeData(
dataLike: DataLike,
context: SanitizationContext,
): Data<DataType> {
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
): import("apache-arrow").Data<Struct<any>> {
if (dataLike instanceof Data) {
return dataLike;
}
const cachedData = context.data.get(dataLike);
if (cachedData !== undefined) {
return cachedData;
}
const dictionaryLike = (dataLike as DictionaryDataLike).dictionary;
let dictionary: Vector | undefined;
if (dictionaryLike !== undefined) {
dictionary = context.vectors.get(dictionaryLike);
if (dictionary === undefined) {
dictionary = new Vector(
dictionaryLike.data.map((data) => sanitizeData(data, context)),
);
context.vectors.set(dictionaryLike, dictionary);
}
}
const data = new Data(
sanitizeTypeWithContext(dataLike.type, context),
return new Data(
dataLike.type,
dataLike.offset,
dataLike.length,
dataLike.nullCount,
@@ -730,11 +589,7 @@ function sanitizeData(
[BufferType.VALIDITY]: dataLike.nullBitmap,
[BufferType.TYPE]: dataLike.typeIds,
},
dataLike.children.map((child) => sanitizeData(child, context)),
dictionary,
);
context.data.set(dataLike, data);
return data;
}
const constructorsByTypeName = {
+4 -42
View File
@@ -30,7 +30,6 @@ import {
DropColumnsResult,
IndexConfig,
IndexStatistics,
Job,
Branches as NativeBranches,
OptimizeStats,
TableStatistics,
@@ -197,11 +196,7 @@ export interface LsmWriteSpec {
column?: string;
/** Bucket variant: the number of buckets, in `[1, 1024]`. */
numBuckets?: number;
/**
* Indexes the MemWAL keeps up to date. Omit to maintain every supported
* index, resolved on install — a snapshot, so indexes created later are not
* maintained. Pass `[]` for none.
*/
/** Names of indexes the MemWAL should keep up to date during writes. */
maintainedIndexes?: string[];
/** Default `ShardWriter` configuration recorded in the MemWAL index. */
writerConfigDefaults?: Record<string, string>;
@@ -363,17 +358,6 @@ export abstract class Table {
options?: Partial<IndexOptions>,
): Promise<void>;
/**
* Create an index, returning a handle to the indexing job.
*
* The job may already be complete when returned; callers must not assume
* the index exists until {@link Job.wait} resolves.
*/
abstract createIndexAsync(
column: string,
options?: Partial<IndexOptions>,
): Promise<Job>;
/**
* Drop an index from the table.
*
@@ -599,11 +583,6 @@ export abstract class Table {
* All variants require the table to have an unenforced primary key
* ({@link Table#setUnenforcedPrimaryKey}); bucket sharding additionally
* requires it to be the single column being bucketed.
*
* Omitting `maintainedIndexes` maintains every index on the table, resolved
* here, failing if one cannot be maintained — name them to install anyway.
* Naming them pins an exact set, and a still-building index is rejected
* rather than quietly omitted.
* @param {LsmWriteSpec} spec The sharding spec to install.
* @returns {Promise<void>}
* @example
@@ -631,10 +610,9 @@ export abstract class Table {
*
* Resolves to `undefined` when the MemWAL LSM write path is not enabled (no
* spec has been set, or it was removed with {@link Table#unsetLsmWriteSpec}).
* The returned spec mirrors what was passed to
* {@link Table#setLsmWriteSpec}, except that `maintainedIndexes` always
* reports the concrete list resolved when the spec was set — `undefined`
* never round-trips.
* The returned spec — including its `maintainedIndexes` and
* `writerConfigDefaults` — mirrors what was passed to
* {@link Table#setLsmWriteSpec}.
* @returns {Promise<LsmWriteSpec | undefined>}
*/
abstract getLsmWriteSpec(): Promise<LsmWriteSpec | undefined>;
@@ -962,22 +940,6 @@ export class LocalTable extends Table {
);
}
async createIndexAsync(
column: string,
options?: Partial<IndexOptions>,
): Promise<Job> {
// biome-ignore lint/suspicious/noExplicitAny: skip
const nativeIndex = (options?.config as any)?.inner;
return await this.inner.createIndexAsync(
nativeIndex,
column,
options?.replace,
options?.waitTimeoutSeconds,
options?.name,
options?.train,
);
}
async dropIndex(name: string): Promise<void> {
await this.inner.dropIndex(name);
}
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-darwin-arm64",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"os": ["darwin"],
"cpu": ["arm64"],
"main": "lancedb.darwin-arm64.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-arm64-gnu",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"os": ["linux"],
"cpu": ["arm64"],
"main": "lancedb.linux-arm64-gnu.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-arm64-musl",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"os": ["linux"],
"cpu": ["arm64"],
"main": "lancedb.linux-arm64-musl.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-x64-gnu",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"os": ["linux"],
"cpu": ["x64"],
"main": "lancedb.linux-x64-gnu.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-linux-x64-musl",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"os": ["linux"],
"cpu": ["x64"],
"main": "lancedb.linux-x64-musl.node",
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-win32-arm64-msvc",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"os": [
"win32"
],
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-win32-x64-msvc",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"os": ["win32"],
"cpu": ["x64"],
"main": "lancedb.win32-x64-msvc.node",
+2 -8
View File
@@ -1,12 +1,12 @@
{
"name": "@lancedb/lancedb",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "@lancedb/lancedb",
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"cpu": [
"x64",
"arm64"
@@ -55,13 +55,7 @@
"openai": "4.29.2"
},
"peerDependencies": {
"@types/node": ">=18",
"apache-arrow": ">=15.0.0 <=18.1.0"
},
"peerDependenciesMeta": {
"@types/node": {
"optional": true
}
}
},
"node_modules/@aws-crypto/crc32": {
+1 -7
View File
@@ -11,7 +11,7 @@
"ann"
],
"private": false,
"version": "0.37.1-beta.1",
"version": "0.37.1-beta.0",
"main": "dist/index.js",
"exports": {
".": "./dist/index.js",
@@ -101,12 +101,6 @@
"openai": "4.29.2"
},
"peerDependencies": {
"@types/node": ">=18",
"apache-arrow": ">=15.0.0 <=18.1.0"
},
"peerDependenciesMeta": {
"@types/node": {
"optional": true
}
}
}
-63
View File
@@ -340,69 +340,6 @@ impl Connection {
self.get_inner()?.drop_all_tables(&ns).await.default_error()
}
/// A `Job` handle for a server-side job by id.
///
/// The handle is constructed without a server round trip; an unknown id
/// surfaces when the handle is used.
#[napi]
pub fn job(&self, job_id: String) -> napi::Result<crate::job::Job> {
let job = self.get_inner()?.job(job_id).default_error()?;
Ok(crate::job::Job::new(job))
}
/// List server-side jobs across the database's tables.
#[napi(catch_unwind)]
pub async fn list_jobs(&self) -> napi::Result<Vec<crate::job::JobInfo>> {
let jobs = self.get_inner()?.list_jobs().await.default_error()?;
Ok(jobs.into_iter().map(Into::into).collect())
}
/// Describe a single server-side job by id. `null` when the server has
/// no such job.
#[napi(catch_unwind)]
pub async fn get_job(
&self,
job_id: String,
) -> napi::Result<Option<crate::job::JobDescription>> {
let description = self.get_inner()?.get_job(&job_id).await.default_error()?;
Ok(description.map(Into::into))
}
/// Request cancellation of a server-side job by id. Returns true if the
/// server accepted the cancellation, false if no such job exists.
#[napi(catch_unwind)]
pub async fn cancel_job(&self, job_id: String) -> napi::Result<bool> {
self.get_inner()?.cancel_job(&job_id).await.default_error()
}
/// The lifecycle event history of a server-side job (all jobs when
/// `job_id` is null), as an Arrow IPC stream buffer. Empty when there is
/// no history.
#[napi(catch_unwind)]
pub async fn job_history(&self, job_id: Option<String>) -> napi::Result<Buffer> {
let batches = self
.get_inner()?
.job_history(job_id.as_deref())
.await
.default_error()?;
let Some(first) = batches.first() else {
return Ok(Buffer::from(Vec::<u8>::new()));
};
let mut out = Vec::new();
let mut writer = arrow_ipc::writer::StreamWriter::try_new(&mut out, &first.schema())
.map_err(|e| napi::Error::from_reason(e.to_string()))?;
for batch in &batches {
writer
.write(batch)
.map_err(|e| napi::Error::from_reason(e.to_string()))?;
}
writer
.finish()
.map_err(|e| napi::Error::from_reason(e.to_string()))?;
drop(writer);
Ok(Buffer::from(out))
}
#[napi(catch_unwind)]
/// Describe a namespace and return its properties.
pub async fn describe_namespace(
-133
View File
@@ -1,133 +0,0 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
use std::sync::Arc;
use napi_derive::napi;
use crate::error::NapiErrorExt;
/// A handle to an operation that may still be running.
#[napi]
pub struct Job {
inner: Arc<lancedb::Job>,
}
impl Job {
pub(crate) fn new(inner: lancedb::Job) -> Self {
Self {
inner: Arc::new(inner),
}
}
}
#[napi]
impl Job {
/// Identifies the operation on the server that is running it. Operations
/// that run in this process have no server id. The value is opaque.
#[napi(getter)]
pub fn id(&self) -> Option<String> {
self.inner.id().map(str::to_string)
}
/// The operation's current lifecycle state: "running", "finished",
/// "failed", or "cancelled".
///
/// A point snapshot; unlike {@link Job.wait} it does not block or reject
/// on a terminal failure state. States a newer server reports that this
/// client version does not know pass through as-is.
#[napi(catch_unwind)]
pub async fn status(&self) -> napi::Result<String> {
self.inner.status().await.default_error()
}
/// Wait until the operation reaches a terminal state.
///
/// Jobs that complete without a resource result resolve successfully.
/// Resource results are not exposed on this binding yet; unsupported
/// success results reject with a generic error.
#[napi(catch_unwind)]
pub async fn wait(&self) -> napi::Result<()> {
match self.inner.wait().await.default_error()? {
lancedb::JobResult::None => Ok(()),
// JobResult is non_exhaustive; Function and future variants fail closed.
_ => Err(napi::Error::from_reason(
"unsupported job result".to_string(),
)),
}
}
/// Request cancellation. Cancelling a finished operation is a no-op.
#[napi(catch_unwind)]
pub async fn cancel(&self) -> napi::Result<()> {
self.inner.cancel().await.default_error()
}
}
/// A row from `Connection.listJobs`: one server-side job.
#[napi(object)]
pub struct JobInfo {
/// The job id -- what `Connection.getJob` and `Connection.cancelJob`
/// accept.
pub job_id: String,
/// The table the job runs against, without URI or namespace.
pub table: String,
pub job_type: String,
/// Lifecycle state: "running", "finished", "failed", or "cancelled".
pub state: String,
/// When the job was created, in milliseconds since the epoch.
pub created_at_millis: i64,
}
impl From<lancedb::database::JobInfo> for JobInfo {
fn from(info: lancedb::database::JobInfo) -> Self {
Self {
job_id: info.job_id,
table: info.table,
job_type: info.job_type,
state: info.state,
created_at_millis: info.created_at_millis,
}
}
}
/// The server's account of why a job failed.
#[napi(object)]
pub struct JobFailureInfo {
pub phase: Option<String>,
pub message: Option<String>,
pub retryable: Option<bool>,
}
/// A described job from `Connection.getJob`.
#[napi(object)]
pub struct JobDescription {
pub job_id: String,
pub job_type: String,
/// Lifecycle state: "running", "finished", "failed", or "cancelled".
pub state: String,
/// When the job was created, in milliseconds since the epoch.
pub creation_ms: i64,
/// The job-type-specific specification as a JSON string, when present.
pub spec_json: Option<String>,
/// Why the job failed, when the job is failed and the server reports a
/// reason.
pub failure: Option<JobFailureInfo>,
}
impl From<lancedb::database::JobDescription> for JobDescription {
fn from(description: lancedb::database::JobDescription) -> Self {
Self {
job_id: description.job_id,
job_type: description.job_type,
state: description.state,
creation_ms: description.creation_ms,
spec_json: (!description.spec.is_null()).then(|| description.spec.to_string()),
failure: description.failure.map(|failure| JobFailureInfo {
phase: failure.phase,
message: failure.message,
retryable: failure.retryable,
}),
}
}
}
-1
View File
@@ -11,7 +11,6 @@ mod error;
mod header;
mod index;
mod iterator;
mod job;
pub mod merge;
pub mod otel;
pub mod permutation;
+9 -49
View File
@@ -168,39 +168,6 @@ impl Table {
builder.execute().await.default_error()
}
#[napi(catch_unwind)]
pub async fn create_index_async(
&self,
index: Option<&Index>,
column: String,
replace: Option<bool>,
wait_timeout_s: Option<i64>,
name: Option<String>,
train: Option<bool>,
) -> napi::Result<crate::job::Job> {
let lancedb_index = if let Some(index) = index {
index.consume()?
} else {
lancedb::index::Index::Auto
};
let mut builder = self.inner_ref()?.create_index(&[column], lancedb_index);
if let Some(replace) = replace {
builder = builder.replace(replace);
}
if let Some(timeout) = wait_timeout_s {
builder =
builder.wait_timeout(std::time::Duration::from_secs(timeout.try_into().unwrap()));
}
if let Some(name) = name {
builder = builder.name(name);
}
if let Some(train) = train {
builder = builder.train(train);
}
let job = builder.execute_async().await.default_error()?;
Ok(crate::job::Job::new(job))
}
#[napi(catch_unwind)]
pub async fn drop_index(&self, index_name: String) -> napi::Result<()> {
self.inner_ref()?
@@ -339,9 +306,7 @@ impl Table {
let transforms = NewColumnTransform::SqlExpressions(transforms);
let res = self
.inner_ref()?
.add_columns()
.transform(transforms)
.execute()
.add_columns(transforms, None)
.await
.default_error()?;
Ok(res.into())
@@ -358,9 +323,7 @@ impl Table {
let transforms = NewColumnTransform::AllNulls(schema);
let res = self
.inner_ref()?
.add_columns()
.transform(transforms)
.execute()
.add_columns(transforms, None)
.await
.default_error()?;
Ok(res.into())
@@ -772,8 +735,7 @@ pub struct LsmWriteSpec {
pub column: Option<String>,
/// Bucket variant: the number of buckets, in `[1, 1024]`.
pub num_buckets: Option<u32>,
/// Indexes the MemWAL keeps up to date. Omitted resolves every
/// maintainable index on install; an empty array means none.
/// Names of indexes the MemWAL should keep up to date during writes.
pub maintained_indexes: Option<Vec<String>>,
/// Default `ShardWriter` configuration recorded in the MemWAL index.
pub writer_config_defaults: Option<HashMap<String, String>>,
@@ -783,6 +745,7 @@ impl TryFrom<LsmWriteSpec> for lancedb::table::LsmWriteSpec {
type Error = napi::Error;
fn try_from(value: LsmWriteSpec) -> napi::Result<Self> {
let maintained = value.maintained_indexes.unwrap_or_default();
let writer_config_defaults = value.writer_config_defaults.unwrap_or_default();
let spec = match value.spec_type.as_str() {
"bucket" => {
@@ -809,7 +772,7 @@ impl TryFrom<LsmWriteSpec> for lancedb::table::LsmWriteSpec {
}
};
Ok(spec
.with_maintained_indexes(value.maintained_indexes)
.with_maintained_indexes(maintained)
.with_writer_config_defaults(writer_config_defaults))
}
}
@@ -827,7 +790,7 @@ impl From<lancedb::table::LsmWriteSpec> for LsmWriteSpec {
spec_type: "bucket".to_string(),
column: Some(column),
num_buckets: Some(num_buckets),
maintained_indexes,
maintained_indexes: Some(maintained_indexes),
writer_config_defaults: Some(writer_config_defaults),
},
Native::Identity {
@@ -838,7 +801,7 @@ impl From<lancedb::table::LsmWriteSpec> for LsmWriteSpec {
spec_type: "identity".to_string(),
column: Some(column),
num_buckets: None,
maintained_indexes,
maintained_indexes: Some(maintained_indexes),
writer_config_defaults: Some(writer_config_defaults),
},
Native::Unsharded {
@@ -848,7 +811,7 @@ impl From<lancedb::table::LsmWriteSpec> for LsmWriteSpec {
spec_type: "unsharded".to_string(),
column: None,
num_buckets: None,
maintained_indexes,
maintained_indexes: Some(maintained_indexes),
writer_config_defaults: Some(writer_config_defaults),
},
}
@@ -1043,10 +1006,7 @@ impl From<lancedb::index::IndexStatistics> for IndexStatistics {
#[napi(object)]
pub struct TableStatistics {
/// The total size, in bytes, of the table's data files, index files, and
/// overlay files
///
/// Read from the manifest, so this excludes deletion files and manifests.
/// The total number of bytes in the table
pub total_bytes: i64,
/// The number of rows in the table
+3 -3
View File
@@ -1,6 +1,6 @@
[package]
name = "lancedb-python"
version = "0.37.1-beta.1"
version = "0.37.1-beta.0"
publish = false
edition.workspace = true
description = "Python bindings for LanceDB"
@@ -26,7 +26,7 @@ lance-namespace-impls.workspace = true
lance-io.workspace = true
env_logger.workspace = true
log.workspace = true
pyo3 = { version = "0.28", features = ["extension-module", "abi3-py310", "chrono"] }
pyo3 = { version = "0.28", features = ["extension-module", "abi3-py39", "chrono"] }
chrono = { version = "0.4", default-features = false, features = ["clock"] }
pyo3-async-runtimes = { version = "0.28", features = [
"attributes",
@@ -43,7 +43,7 @@ libc = "0.2"
[build-dependencies]
pyo3-build-config = { version = "0.28", features = [
"extension-module",
"abi3-py310",
"abi3-py39",
] }
[features]
+1 -2
View File
@@ -60,7 +60,7 @@ tests = [
"pytest-asyncio>=0.21",
"duckdb>=0.9.0",
"pytz>=2023.3",
"polars>=0.19, <=1.32.3",
"polars>=0.19, <=1.3.0",
"pyarrow<25",
"pyarrow-stubs>=16.0",
"pylance==9.0.0rc1",
@@ -140,7 +140,6 @@ include = [
"python/lancedb/remote/errors.py",
"python/lancedb/embeddings/__init__.py",
"python/lancedb/_lancedb.pyi",
"python/type_tests/connect.py",
]
exclude = ["python/tests/"]
pythonVersion = "3.13"
-8
View File
@@ -12,7 +12,6 @@ __version__ = importlib.metadata.version("lancedb")
from ._lancedb import connect as lancedb_connect
from ._lancedb import FtsToken
from ._lancedb import Function
from ._lancedb import tokenize as _tokenize
from .common import URI, sanitize_uri
from urllib.parse import urlparse
@@ -21,10 +20,8 @@ from .remote import ClientConfig
from .remote.db import RemoteDBConnection
from .expr import Expr, col, lit, func
from .schema import blob, vector, BlobType
from .job import AsyncJob, Job
from .table import AsyncTable, Table
from .types import BaseTokenizerType
from ._udf import FunctionCapability, udf
from ._lancedb import Session
from .namespace import (
connect_namespace,
@@ -503,14 +500,11 @@ __all__ = [
"connect_namespace",
"connect_namespace_async",
"AsyncConnection",
"AsyncJob",
"AsyncLanceNamespaceDBConnection",
"AsyncTable",
"FtsToken",
"col",
"Expr",
"Function",
"FunctionCapability",
"func",
"lit",
"URI",
@@ -519,12 +513,10 @@ __all__ = [
"BlobType",
"vector",
"DBConnection",
"Job",
"LanceDBConnection",
"LanceNamespaceDBConnection",
"RemoteDBConnection",
"Session",
"Table",
"udf",
"__version__",
]
+23 -6
View File
@@ -14,10 +14,14 @@ import pyarrow as pa
from .expr import Expr
from .schema import blob_v2_column_paths
from .types import BlobMode, QueryProjection, QueryProjectionSpec
from .util import get_uri_scheme
if TYPE_CHECKING:
from _typeshed import WriteableBuffer
from .remote.table import RemoteTable
from .table import AsyncTable, Table
BLOB_MODE_TO_HANDLING = {
"lazy": "blobs_descriptions",
"bytes": "all_binary",
@@ -100,6 +104,22 @@ def validate_blob_mode(blob_mode: BlobMode) -> None:
raise ValueError(f"blob_mode must be one of {modes}, got {blob_mode!r}")
def supports_blob_auto_row_id(table: Table | AsyncTable | RemoteTable) -> bool:
"""Blob auto row-id applies to native tables, not LanceDB Cloud."""
from .remote.table import RemoteTable
if isinstance(table, RemoteTable):
return False
inner = getattr(table, "_inner", None)
if inner is not None:
uri = inner.database().uri
if isinstance(uri, str) and get_uri_scheme(uri) == "db":
return False
return True
def projection_includes_blob_column(
projection: QueryProjection,
blob_columns: Iterable[str],
@@ -144,14 +164,16 @@ def v2_projection_needs_row_id(
def blob_auto_row_id_for_scan(
table: Table | AsyncTable | RemoteTable,
schema: pa.Schema,
projection: QueryProjection,
*,
with_row_id: bool | None,
) -> bool:
"""Auto row-id only applies when the caller said nothing about row ids."""
if with_row_id is not None:
return False
if not supports_blob_auto_row_id(table):
return False
return v2_projection_needs_row_id(schema, projection, with_row_id=False)
@@ -164,11 +186,6 @@ def finalize_blob_query_table(
) -> pa.Table:
if user_requested_row_id or not blob_auto_row_id:
return tbl
if "_rowid" not in tbl.column_names:
# A backend that ignores the row-id request leaves nothing to stash. Hand
# back the projection as-is so fetch_blobs raises the error that names the
# ways to supply row ids, rather than failing here about a hidden column.
return tbl
return stash_auto_row_ids(tbl, blob_paths)
-108
View File
@@ -1,108 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Private first-class Function namespace facades for database connections.
These helpers are internal submission and lookup surfaces. They are not durable
resources and are not part of the public top-level export surface.
"""
from __future__ import annotations
from collections.abc import Callable
from typing import TYPE_CHECKING
from . import _udf
from ._lancedb import Function
from .job import AsyncJob, Job
if TYPE_CHECKING:
from .db import AsyncConnection, DBConnection
class _SyncFunctions:
"""Synchronous `db.functions` facade."""
__slots__ = ("_connection",)
def __init__(self, connection: DBConnection) -> None:
self._connection = connection
def __repr__(self) -> str:
return "_SyncFunctions()"
def register(self, name: str, decorated_udf: Callable[..., object]) -> Job:
"""Register a decorated UDF and return a synchronous [Job][lancedb.job.Job]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = self._connection._submit_register_function(name, definition)
return Job(AsyncJob(native_job))
def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> Job:
"""Conditionally replace a Function; return sync [Job][lancedb.job.Job]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = self._connection._submit_replace_function(
name, current, definition
)
return Job(AsyncJob(native_job))
def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return self._connection._lookup_function_by_name(name)
def get_by_id(self, function_id: str) -> Function:
"""Return the immutable Function for an exact Function ID."""
return self._connection._lookup_function_by_id(function_id)
def remove(self, name: str, current: Function) -> None:
"""Conditionally remove a Function catalog name binding."""
return self._connection._remove_function_name(name, current)
def revoke(self, function: Function) -> None:
"""Revoke an exact immutable Function by administrator set-bit."""
return self._connection._revoke_function(function)
class _AsyncFunctions:
"""Asynchronous `async_db.functions` facade."""
__slots__ = ("_connection",)
def __init__(self, connection: AsyncConnection) -> None:
self._connection = connection
def __repr__(self) -> str:
return "_AsyncFunctions()"
async def register(
self, name: str, decorated_udf: Callable[..., object]
) -> AsyncJob:
"""Register a decorated UDF and return an [AsyncJob][lancedb.job.AsyncJob]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = await self._connection._register_function(name, definition)
return AsyncJob(native_job)
async def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> AsyncJob:
"""Conditionally replace a Function; return [AsyncJob][lancedb.job.AsyncJob]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = await self._connection._replace_function(name, current, definition)
return AsyncJob(native_job)
async def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return await self._connection._lookup_function_by_name(name)
async def get_by_id(self, function_id: str) -> Function:
"""Return the immutable Function for an exact Function ID."""
return await self._connection._lookup_function_by_id(function_id)
async def remove(self, name: str, current: Function) -> None:
"""Conditionally remove a Function catalog name binding."""
return await self._connection._remove_function_name(name, current)
async def revoke(self, function: Function) -> None:
"""Revoke an exact immutable Function by administrator set-bit."""
return await self._connection._revoke_function(function)
+4 -139
View File
@@ -146,23 +146,6 @@ class Connection(object):
start_after: Optional[str],
limit: Optional[int],
) -> list[str]: ... # Deprecated: Use list_tables instead
def job(self, job_id: str) -> Job: ...
async def list_jobs(self) -> List[JobInfo]: ...
async def get_job(self, job_id: str) -> Optional[JobDescription]: ...
async def cancel_job(self, job_id: str) -> bool: ...
async def job_history(
self, job_id: Optional[str] = None
) -> List[pa.RecordBatch]: ...
async def _register_function(
self, name: str, definition: "_FunctionDefinition"
) -> Job: ...
async def _replace_function(
self, name: str, current: Function, definition: "_FunctionDefinition"
) -> Job: ...
async def _lookup_function_by_name(self, name: str) -> Function: ...
async def _lookup_function_by_id(self, function_id: str) -> Function: ...
async def _remove_function_name(self, name: str, current: Function) -> None: ...
async def _revoke_function(self, function: Function) -> None: ...
async def create_table(
self,
name: str,
@@ -226,85 +209,6 @@ class BlobFile:
def read_range(self, offset: int, length: int) -> bytes: ...
def read_up_to(self, length: int) -> bytes: ...
class Function:
@property
def id(self) -> str: ...
@property
def parameters(self) -> tuple[tuple[str, pa.DataType], ...]: ...
@property
def output_type(self) -> pa.DataType: ...
@property
def output_nullable(self) -> bool: ...
def __call__(self, **kwargs: Any) -> "_FunctionCall": ...
class _FunctionCall:
"""Private unresolved Function call authoring value (FF-028)."""
...
class _FunctionDefinition:
"""Private owner of the Rust FunctionDefinition registration input."""
def _to_json(self) -> str: ...
def _new_function_definition(
*,
parameters: list[tuple[str, pa.DataType]],
output_type: pa.DataType,
output_nullable: bool,
module: str,
callable_name: str,
source: str,
python: str,
packages: list[str],
capabilities: list[tuple[str, str, Optional[str]]],
) -> _FunctionDefinition: ...
class Job:
@property
def id(self) -> Optional[str]: ...
async def status(self) -> str: ...
async def wait(self) -> Optional[Function]: ...
async def cancel(self) -> None: ...
class JobInfo:
@property
def job_id(self) -> str: ...
@property
def table(self) -> str: ...
@property
def job_type(self) -> str: ...
@property
def state(self) -> str: ...
@property
def created_at_millis(self) -> int: ...
class JobFailureInfo:
@property
def phase(self) -> Optional[str]: ...
@property
def message(self) -> Optional[str]: ...
@property
def retryable(self) -> Optional[bool]: ...
@property
def error_code(self) -> Optional[str]: ...
class JobDescription:
@property
def job_id(self) -> str: ...
@property
def job_type(self) -> str: ...
@property
def state(self) -> str: ...
@property
def creation_ms(self) -> int: ...
@property
def spec_json(self) -> Optional[str]: ...
@property
def failure(self) -> Optional[JobFailureInfo]: ...
@property
def result(self) -> Optional[Function]: ...
class Table:
def name(self) -> str: ...
def __repr__(self) -> str: ...
@@ -344,38 +248,6 @@ class Table:
name: Optional[str],
train: Optional[bool],
): ...
async def create_index_async(
self,
column: str,
index: Union[
IvfFlat,
IvfSq,
IvfPq,
HnswPq,
HnswSq,
HnswFlat,
BTree,
Bitmap,
LabelList,
Fm,
FTS,
],
replace: Optional[bool],
wait_timeout: Optional[object],
*,
name: Optional[str],
train: Optional[bool],
) -> Job: ...
async def _add_generated_column(
self, column_name: str, call: _FunctionCall
) -> Job: ...
async def _generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]: ...
async def _refresh_generated_column(self, column_name: str) -> Job: ...
async def _alter_generated_column(
self, column_name: str, new_call: _FunctionCall
) -> Job: ...
async def list_versions(self) -> List[Dict[str, Any]]: ...
async def version(self) -> int: ...
async def checkout(self, version: Union[int, str]): ...
@@ -413,10 +285,6 @@ class Table:
async def set_lsm_write_spec(self, spec: LsmWriteSpec) -> None: ...
async def unset_lsm_write_spec(self) -> None: ...
async def get_lsm_write_spec(self) -> Optional[LsmWriteSpec]: ...
async def checkpoint_lsm(self) -> None: ...
async def flush_lsm(self) -> None: ...
async def compact_lsm(self) -> None: ...
async def get_lsm_stats(self, include_generation_rows: bool) -> Optional[dict]: ...
async def close_lsm_writers(self) -> None: ...
@property
def tags(self) -> Tags: ...
@@ -711,10 +579,9 @@ class LsmWriteSpec:
def identity(column: str) -> "LsmWriteSpec": ...
@staticmethod
def unsharded() -> "LsmWriteSpec": ...
def with_maintained_indexes(self, indexes: Optional[List[str]]) -> "LsmWriteSpec":
"""Set which indexes the MemWAL keeps up to date. None resolves every
index on the table at install, failing if one cannot be maintained;
a list is verbatim, empty means none."""
def with_maintained_indexes(self, indexes: List[str]) -> "LsmWriteSpec":
"""Return a copy of this spec asking the MemWAL to keep the named
indexes up to date as rows are appended."""
...
def with_writer_config_defaults(self, defaults: Dict[str, str]) -> "LsmWriteSpec":
"""Return a copy of this spec recording the given default
@@ -729,9 +596,7 @@ class LsmWriteSpec:
@property
def num_buckets(self) -> Optional[int]: ...
@property
def maintained_indexes(self) -> Optional[List[str]]:
"""Indexes the MemWAL keeps up to date, or None for every supported one."""
...
def maintained_indexes(self) -> List[str]: ...
@property
def writer_config_defaults(self) -> Dict[str, str]: ...
-538
View File
@@ -1,538 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Local authoring declaration surface for first-class UDFs.
This module snapshots declaration metadata onto a Python function and privately
validates packagable callables into a source snapshot. It does not mint durable
identity or register anything with a database.
"""
from __future__ import annotations
import ast
import inspect
import stat
import symtable
import sys
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import CodeType, FunctionType
from typing import NoReturn, ParamSpec, TypeVar
import pyarrow as pa
from . import _lancedb
__all__ = ["FunctionCapability", "udf"]
_P = ParamSpec("_P")
_R = TypeVar("_R")
_CONFIG_ATTR = "__lancedb_udf_config__"
_SYNTHETIC_SOURCE_FILENAME = "<lancedb-udf>"
_PACKAGING_ERROR = "udf is not packagable"
_ALLOWED_PARAM_KINDS = (
inspect.Parameter.POSITIONAL_OR_KEYWORD,
inspect.Parameter.KEYWORD_ONLY,
)
class FunctionCapability:
"""Local capability declaration for a first-class UDF.
Construct via :meth:`network` or :meth:`secret`. Direct construction is
rejected so callers cannot create an uninitialized capability.
"""
__slots__ = ("_kind", "_origin", "_reference", "_environment_variable")
def __new__(cls, *args: object, **kwargs: object) -> FunctionCapability:
raise TypeError(
"FunctionCapability cannot be constructed directly; "
"use FunctionCapability.network() or FunctionCapability.secret()"
)
@classmethod
def _create(
cls,
kind: str,
origin: str | None,
reference: str | None,
environment_variable: str | None,
) -> FunctionCapability:
obj = object.__new__(cls)
object.__setattr__(obj, "_kind", kind)
object.__setattr__(obj, "_origin", origin)
object.__setattr__(obj, "_reference", reference)
object.__setattr__(obj, "_environment_variable", environment_variable)
return obj
@classmethod
def network(cls, origin: str) -> FunctionCapability:
if not isinstance(origin, str):
raise TypeError("origin must be a string")
if origin == "":
raise ValueError("origin must be non-empty")
return cls._create("network", origin, None, None)
@classmethod
def secret(cls, reference: str, *, environment_variable: str) -> FunctionCapability:
if not isinstance(reference, str):
raise TypeError("reference must be a string")
if not isinstance(environment_variable, str):
raise TypeError("environment_variable must be a string")
if reference == "":
raise ValueError("reference must be non-empty")
if environment_variable == "":
raise ValueError("environment_variable must be non-empty")
return cls._create("secret", None, reference, environment_variable)
@property
def kind(self) -> str:
return self._kind
@property
def origin(self) -> str | None:
return self._origin
@property
def reference(self) -> str | None:
return self._reference
@property
def environment_variable(self) -> str | None:
return self._environment_variable
def __setattr__(self, name: str, value: object) -> None:
raise AttributeError(
f"{type(self).__name__!r} object attribute {name!r} is read-only"
)
def __delattr__(self, name: str) -> None:
raise AttributeError(
f"{type(self).__name__!r} object attribute {name!r} is read-only"
)
def __eq__(self, other: object) -> bool:
if not isinstance(other, FunctionCapability):
return NotImplemented
return (
self._kind == other._kind
and self._origin == other._origin
and self._reference == other._reference
and self._environment_variable == other._environment_variable
)
def __hash__(self) -> int:
return hash(
(
self._kind,
self._origin,
self._reference,
self._environment_variable,
)
)
def __repr__(self) -> str:
if self._kind == "network":
return f"FunctionCapability.network({self._origin!r})"
return (
"FunctionCapability.secret("
f"environment_variable={self._environment_variable!r})"
)
@dataclass(frozen=True, slots=True)
class _UdfConfig:
"""Private frozen snapshot of a ``@udf`` declaration."""
inputs: tuple[tuple[str, pa.DataType], ...]
output: pa.DataType
output_nullable: bool
python: str
packages: tuple[str, ...]
capabilities: tuple[FunctionCapability, ...]
@dataclass(frozen=True, slots=True)
class _PackagedUdf:
"""Private frozen snapshot of a validated packagable UDF."""
source: str
module: str
callable_name: str
config: _UdfConfig
def __repr__(self) -> str:
return (
f"_PackagedUdf(source=<redacted>, module={self.module!r}, "
f"callable_name={self.callable_name!r}, config={self.config!r})"
)
def _validate_inputs(
inputs: object,
) -> tuple[tuple[str, pa.DataType], ...]:
if not isinstance(inputs, Mapping):
raise TypeError("udf inputs must be a Mapping of name to pyarrow DataType")
snapshot: list[tuple[str, pa.DataType]] = []
for key, value in inputs.items():
if not isinstance(key, str):
raise TypeError("udf input names must be strings")
if key == "":
raise ValueError("udf input names must be non-empty")
if not isinstance(value, pa.DataType):
raise TypeError("udf input types must be pyarrow DataType values")
snapshot.append((key, value))
return tuple(snapshot)
def _validate_packages(packages: object) -> tuple[str, ...]:
if isinstance(packages, (str, bytes, bytearray)):
raise TypeError("udf packages must be a sequence of strings, not a string")
if not isinstance(packages, Sequence):
raise TypeError("udf packages must be a sequence of strings")
snapshot: list[str] = []
seen: set[str] = set()
for package in packages:
if not isinstance(package, str):
raise TypeError("udf packages must contain only strings")
if package == "":
raise ValueError("udf packages must be non-empty strings")
if package in seen:
raise ValueError(f"duplicate udf package: {package}")
seen.add(package)
snapshot.append(package)
return tuple(snapshot)
def _reject_non_exact_capability() -> NoReturn:
# Exact-type only: subclasses are authoring inputs we never accept. Keep the
# message fixed so hostile markers never enter exception text.
raise TypeError(
"udf capabilities must contain only FunctionCapability values"
) from None
def _require_exact_capability(capability: object) -> FunctionCapability:
if type(capability) is not FunctionCapability:
_reject_non_exact_capability()
return capability
def _validate_capabilities(capabilities: object) -> tuple[FunctionCapability, ...]:
if isinstance(capabilities, (str, bytes, bytearray)):
raise TypeError(
"udf capabilities must be a sequence of FunctionCapability, not a string"
)
if not isinstance(capabilities, Sequence):
raise TypeError("udf capabilities must be a sequence of FunctionCapability")
return tuple(_require_exact_capability(capability) for capability in capabilities)
def udf(
*,
inputs: Mapping[str, pa.DataType],
output: pa.DataType,
python: str,
packages: Sequence[str] = (),
output_nullable: bool = True,
capabilities: Sequence[FunctionCapability] = (),
) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]:
"""Declare a local UDF without packaging or registration.
Applying the returned decorator attaches a private frozen config snapshot
and returns the exact same function object.
"""
input_snapshot = _validate_inputs(inputs)
if not isinstance(output, pa.DataType):
raise TypeError("udf output must be a pyarrow DataType")
if not isinstance(python, str):
raise TypeError("udf python must be a string")
if python == "":
raise ValueError("udf python must be a non-empty string")
package_snapshot = _validate_packages(packages)
if not isinstance(output_nullable, bool):
raise TypeError("udf output_nullable must be a bool")
capability_snapshot = _validate_capabilities(capabilities)
config = _UdfConfig(
inputs=input_snapshot,
output=output,
output_nullable=output_nullable,
python=python,
packages=package_snapshot,
capabilities=capability_snapshot,
)
def decorator(fn: Callable[_P, _R]) -> Callable[_P, _R]:
if not inspect.isfunction(fn):
raise TypeError("udf can only decorate a Python function")
if hasattr(fn, _CONFIG_ATTR):
raise ValueError("function is already decorated with @udf")
setattr(fn, _CONFIG_ATTR, config)
return fn
return decorator
def _get_udf_config(fn: object) -> _UdfConfig:
"""Return the private declaration snapshot for a ``@udf``-decorated function."""
config = getattr(fn, _CONFIG_ATTR, None)
if config is None:
raise TypeError("function is not decorated with @udf")
if not isinstance(config, _UdfConfig):
raise TypeError("function is not decorated with @udf")
return config
def _packaging_reject() -> NoReturn:
raise ValueError(_PACKAGING_ERROR) from None
def _is_ordinary_function(fn: FunctionType) -> bool:
if fn.__name__ == "<lambda>":
return False
if fn.__qualname__ != fn.__name__:
return False
if inspect.iscoroutinefunction(fn) or inspect.isasyncgenfunction(fn):
return False
if inspect.isgeneratorfunction(fn):
return False
return True
def _resolve_source_path(fn: FunctionType, module: object) -> Path:
try:
fn_source: str | None = inspect.getsourcefile(fn)
except TypeError:
fn_source = None
source_lookup_failed = True
else:
source_lookup_failed = False
if source_lookup_failed:
_packaging_reject()
module_file = vars(module).get("__file__")
if not fn_source or not isinstance(module_file, str) or module_file == "":
_packaging_reject()
try:
resolved_paths: tuple[Path, Path] | None = (
Path(fn_source).resolve(),
Path(module_file).resolve(),
)
except (OSError, RuntimeError):
resolved_paths = None
if resolved_paths is None:
_packaging_reject()
fn_path, module_path = resolved_paths
if fn_path != module_path:
_packaging_reject()
if fn_path.suffix != ".py":
_packaging_reject()
try:
mode: int | None = fn_path.stat().st_mode
except OSError:
mode = None
if mode is None:
_packaging_reject()
if not stat.S_ISREG(mode):
_packaging_reject()
return fn_path
def _validate_source(
source: str, callable_name: str
) -> tuple[CodeType, symtable.SymbolTable]:
try:
module_code = compile(
source,
_SYNTHETIC_SOURCE_FILENAME,
"exec",
optimize=sys.flags.optimize,
)
ast.parse(source, filename=_SYNTHETIC_SOURCE_FILENAME, mode="exec")
table = symtable.symtable(source, _SYNTHETIC_SOURCE_FILENAME, "exec")
parsed: tuple[CodeType, symtable.SymbolTable] | None = (module_code, table)
except Exception:
parsed = None
if parsed is None:
_packaging_reject()
module_code, table = parsed
for child in table.get_children():
if child.get_name() == callable_name and child.get_type() == "function":
return module_code, table
_packaging_reject()
def _source_bound_names(table: symtable.SymbolTable) -> set[str]:
names: set[str] = set()
for symbol in table.get_symbols():
if symbol.is_imported() or symbol.is_assigned() or symbol.is_namespace():
names.add(symbol.get_name())
return names
def _code_fingerprint(code: CodeType) -> tuple[object, ...]:
"""Structural fingerprint ignoring only location/debug fields."""
constants = tuple(
_code_fingerprint(constant) if isinstance(constant, CodeType) else constant
for constant in code.co_consts
)
return (
code.co_name,
getattr(code, "co_qualname", code.co_name),
code.co_argcount,
code.co_posonlyargcount,
code.co_kwonlyargcount,
code.co_flags,
code.co_code,
code.co_names,
code.co_varnames,
code.co_freevars,
code.co_cellvars,
getattr(code, "co_exceptiontable", b""),
constants,
)
def _toplevel_code_candidates(
module_code: CodeType, callable_name: str
) -> list[CodeType]:
candidates: list[CodeType] = []
for constant in module_code.co_consts:
if not isinstance(constant, CodeType):
continue
if constant.co_name != callable_name:
continue
if getattr(constant, "co_qualname", callable_name) != callable_name:
continue
candidates.append(constant)
return candidates
def _validate_loaded_code_matches_source(
fn: FunctionType, module_code: CodeType
) -> None:
candidates = _toplevel_code_candidates(module_code, fn.__name__)
if not candidates:
_packaging_reject()
target = _code_fingerprint(fn.__code__)
if not any(_code_fingerprint(candidate) == target for candidate in candidates):
_packaging_reject()
def _validate_signature(fn: FunctionType, config: _UdfConfig) -> None:
try:
signature: inspect.Signature | None = inspect.signature(fn)
except (TypeError, ValueError):
signature = None
if signature is None:
_packaging_reject()
parameters = list(signature.parameters.values())
expected = [name for name, _ in config.inputs]
actual = [parameter.name for parameter in parameters]
if actual != expected:
_packaging_reject()
for parameter in parameters:
if parameter.kind not in _ALLOWED_PARAM_KINDS:
_packaging_reject()
def _validate_ambient_globals(fn: FunctionType, table: symtable.SymbolTable) -> None:
try:
closure_vars: inspect.ClosureVars | None = inspect.getclosurevars(fn)
except (TypeError, ValueError):
closure_vars = None
if closure_vars is None:
_packaging_reject()
if closure_vars.nonlocals:
_packaging_reject()
bound_names = _source_bound_names(table)
for name in closure_vars.globals:
if name not in bound_names:
_packaging_reject()
def _package_udf(fn: object) -> _PackagedUdf:
"""Validate and snapshot a packagable ``@udf``-decorated function."""
config = _get_udf_config(fn)
if not isinstance(fn, FunctionType) or not _is_ordinary_function(fn):
_packaging_reject()
module_name = fn.__module__
if (
not isinstance(module_name, str)
or module_name == ""
or module_name == "__main__"
):
_packaging_reject()
module = sys.modules.get(module_name)
if module is None:
_packaging_reject()
callable_name = fn.__name__
if vars(module).get(callable_name) is not fn:
_packaging_reject()
source_path = _resolve_source_path(fn, module)
try:
source: str | None = source_path.read_text(encoding="utf-8")
except (OSError, UnicodeError):
source = None
if source is None:
_packaging_reject()
module_code, table = _validate_source(source, callable_name)
_validate_signature(fn, config)
_validate_ambient_globals(fn, table)
_validate_loaded_code_matches_source(fn, module_code)
return _PackagedUdf(
source=source,
module=module_name,
callable_name=callable_name,
config=config,
)
def _normalize_capability_triple(
capability: FunctionCapability,
) -> tuple[str, str, str | None]:
"""Normalize a local capability declaration to the native triple shape."""
# Private config is untrusted; re-check exact type before any property access.
capability = _require_exact_capability(capability)
if capability.kind == "network":
origin = capability.origin
if origin is None:
raise ValueError("invalid network capability") from None
return ("network", origin, None)
if capability.kind == "secret":
reference = capability.reference
environment_variable = capability.environment_variable
if reference is None or environment_variable is None:
raise ValueError("invalid secret capability") from None
return ("secret", reference, environment_variable)
# Fail closed without echoing the unknown kind.
raise ValueError("unsupported capability kind") from None
def _build_function_definition(fn: object) -> _lancedb._FunctionDefinition:
"""Package a ``@udf`` and bridge it to the private native definition."""
packaged = _package_udf(fn)
config = packaged.config
capabilities = [
_normalize_capability_triple(capability) for capability in config.capabilities
]
return _lancedb._new_function_definition(
parameters=list(config.inputs),
output_type=config.output,
output_nullable=config.output_nullable,
module=packaged.module,
callable_name=packaged.callable_name,
source=packaged.source,
python=config.python,
packages=list(config.packages),
capabilities=capabilities,
)
+8 -309
View File
@@ -45,7 +45,6 @@ from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
from . import __version__
from ._lancedb import connect as lancedb_connect # type: ignore
from .job import AsyncJob, Job
from .table import (
AsyncTable,
LanceTable,
@@ -63,12 +62,7 @@ if TYPE_CHECKING:
import pyarrow as pa
from .pydantic import LanceModel
from ._functions import _AsyncFunctions, _SyncFunctions
from ._lancedb import Connection as LanceDbConnection
from ._lancedb import Function
from ._lancedb import Job as NativeJob
from ._lancedb import JobDescription, JobInfo
from ._lancedb import _FunctionDefinition
from .common import DATA, URI
from .embeddings import EmbeddingFunctionConfig
from ._lancedb import Session
@@ -184,51 +178,6 @@ class DBConnection(EnforceOverrides):
"Namespace operations are not supported for this connection type"
)
def namespace_exists(self, namespace_id: List[str]) -> bool:
"""Check if a namespace exists.
Parameters
----------
namespace_id: List[str]
The namespace identifier to check.
Returns
-------
bool
True if the namespace exists, False otherwise.
Raises
------
NotImplementedError
If the connection type does not support namespace operations.
"""
raise NotImplementedError(
"Namespace operations are not supported for this connection type"
)
def table_exists(self, table_id: List[str]) -> bool:
"""Check if a table exists.
Parameters
----------
table_id: List[str]
The table identifier to check (full path including namespace
segments and table name).
Returns
-------
bool
True if the table exists, False otherwise.
Raises
------
NotImplementedError
If the connection type does not support namespace operations.
"""
raise NotImplementedError(
"Namespace operations are not supported for this connection type"
)
def list_tables(
self,
namespace_path: Optional[List[str]] = None,
@@ -614,111 +563,6 @@ class DBConnection(EnforceOverrides):
"""
raise NotImplementedError("serialize is not supported for this connection type")
def job(self, job_id: str) -> Job:
"""A [Job][lancedb.job.Job] handle for a server-side job by id.
The handle is constructed without a server round trip; an unknown id
surfaces when the handle is used. Dropping the handle has no effect
on the job itself.
"""
raise NotImplementedError("job is not supported for this connection type")
def list_jobs(self) -> List[JobInfo]:
"""List server-side jobs across the database's tables."""
raise NotImplementedError("list_jobs is not supported for this connection type")
def get_job(self, job_id: str) -> Optional[JobDescription]:
"""Describe a single server-side job by id.
Returns None when the server has no such job.
"""
raise NotImplementedError("get_job is not supported for this connection type")
def cancel_job(self, job_id: str) -> bool:
"""Request cancellation of a server-side job by id.
Returns True if the server accepted the cancellation, False if no
such job exists. Cancelling an already-terminal job is a no-op
success.
"""
raise NotImplementedError(
"cancel_job is not supported for this connection type"
)
def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
"""The lifecycle event history of a server-side job, as Arrow batches.
Lists history across all jobs when `job_id` is None.
"""
raise NotImplementedError(
"job_history is not supported for this connection type"
)
@property
def functions(self) -> "_SyncFunctions":
"""First-class Function operations for this connection."""
from ._functions import _SyncFunctions
return _SyncFunctions(self)
def _submit_register_function(
self, name: str, definition: "_FunctionDefinition"
) -> NativeJob:
"""Submit a Function registration job via the native connection.
Connection subclasses that support registration override this hook.
"""
raise NotImplementedError(
"function registration is not supported for this connection type"
)
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
"""Submit a Function conditional replace job via the native connection.
Connection subclasses that support registration override this hook.
"""
raise NotImplementedError(
"function replace is not supported for this connection type"
)
def _lookup_function_by_name(self, name: str) -> "Function":
"""Look up a Function by database-scoped name via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function lookup is not supported for this connection type"
)
def _lookup_function_by_id(self, function_id: str) -> "Function":
"""Look up a Function by exact Function ID via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function lookup is not supported for this connection type"
)
def _remove_function_name(self, name: str, current: "Function") -> None:
"""Conditionally remove a Function catalog name via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function name removal is not supported for this connection type"
)
def _revoke_function(self, function: "Function") -> None:
"""Revoke an exact immutable Function via the native connection.
Connection subclasses that share the native Connection override this hook.
"""
raise NotImplementedError(
"function revocation is not supported for this connection type"
)
class LanceDBConnection(DBConnection):
"""
@@ -776,9 +620,6 @@ class LanceDBConnection(DBConnection):
self._namespace_client_properties = namespace_client_properties
if _inner is not None:
self._conn = _inner
# Native-derived wrappers resolve this in their async reconstruction
# path so construction never synchronously re-enters LOOP.
self._read_consistency_interval = read_consistency_interval
self._cached_namespace_client = None
return
@@ -828,14 +669,11 @@ class LanceDBConnection(DBConnection):
# storage_options. Also, this class really shouldn't be holding any state
# beyond _conn.
self._conn = AsyncConnection(LOOP.run(do_connect()))
# Keep property access synchronous so debugger introspection cannot wait on
# the background loop while that thread is suspended at a breakpoint.
self._read_consistency_interval = read_consistency_interval
self._cached_namespace_client: Optional[LanceNamespace] = None
@property
def read_consistency_interval(self) -> Optional[timedelta]:
return self._read_consistency_interval
return LOOP.run(self._conn.get_read_consistency_interval())
@property
def session(self) -> Optional[Session]:
@@ -846,19 +684,15 @@ class LanceDBConnection(DBConnection):
return self._conn.uri
@classmethod
def from_inner(
cls,
inner: LanceDbConnection,
read_consistency_interval: Optional[timedelta],
):
return cls(
None,
read_consistency_interval=read_consistency_interval,
_inner=inner,
)
def from_inner(cls, inner: LanceDbConnection):
return cls(None, _inner=inner)
def __repr__(self) -> str:
return f"{self.__class__.__name__}(uri={self._conn.uri!r})"
val = f"{self.__class__.__name__}(uri={self._conn.uri!r}"
if self.read_consistency_interval is not None:
val += f", read_consistency_interval={repr(self.read_consistency_interval)}"
val += ")"
return val
@override
def serialize(self) -> str:
@@ -1295,75 +1129,6 @@ class LanceDBConnection(DBConnection):
)
)
@override
def job(self, job_id: str) -> Job:
"""A [Job][lancedb.job.Job] handle for a server-side job by id.
The handle is constructed without a server round trip; an unknown id
surfaces when the handle is used. Dropping the handle has no effect
on the job itself.
"""
return Job(self._conn.job(job_id))
@override
def list_jobs(self) -> List[JobInfo]:
"""List server-side jobs across the database's tables."""
return LOOP.run(self._conn.list_jobs())
@override
def get_job(self, job_id: str) -> Optional[JobDescription]:
"""Describe a single server-side job by id.
Returns None when the server has no such job.
"""
return LOOP.run(self._conn.get_job(job_id))
@override
def cancel_job(self, job_id: str) -> bool:
"""Request cancellation of a server-side job by id.
Returns True if the server accepted the cancellation, False if no
such job exists. Cancelling an already-terminal job is a no-op
success.
"""
return LOOP.run(self._conn.cancel_job(job_id))
@override
def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
"""The lifecycle event history of a server-side job, as Arrow batches.
Lists history across all jobs when `job_id` is None.
"""
return LOOP.run(self._conn.job_history(job_id))
@override
def _submit_register_function(
self, name: str, definition: "_FunctionDefinition"
) -> NativeJob:
return LOOP.run(self._conn._register_function(name, definition))
@override
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
return LOOP.run(self._conn._replace_function(name, current, definition))
@override
def _lookup_function_by_name(self, name: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_name(name))
@override
def _lookup_function_by_id(self, function_id: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_id(function_id))
@override
def _remove_function_name(self, name: str, current: "Function") -> None:
return LOOP.run(self._conn._remove_function_name(name, current))
@override
def _revoke_function(self, function: "Function") -> None:
return LOOP.run(self._conn._revoke_function(function))
@override
def namespace_client(self) -> LanceNamespace:
"""Get the equivalent namespace client for this connection.
@@ -2073,72 +1838,6 @@ class AsyncConnection(object):
namespace_path = []
await self._inner.drop_all_tables(namespace_path=namespace_path)
def job(self, job_id: str) -> AsyncJob:
"""An [AsyncJob][lancedb.job.AsyncJob] handle for a server-side job
by id.
The handle is constructed without a server round trip; an unknown id
surfaces when the handle is used. Dropping the handle has no effect
on the job itself.
"""
return AsyncJob(self._inner.job(job_id))
async def list_jobs(self) -> List[JobInfo]:
"""List server-side jobs across the database's tables."""
return await self._inner.list_jobs()
async def get_job(self, job_id: str) -> Optional[JobDescription]:
"""Describe a single server-side job by id.
Returns None when the server has no such job.
"""
return await self._inner.get_job(job_id)
async def cancel_job(self, job_id: str) -> bool:
"""Request cancellation of a server-side job by id.
Returns True if the server accepted the cancellation, False if no
such job exists. Cancelling an already-terminal job is a no-op
success.
"""
return await self._inner.cancel_job(job_id)
async def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
"""The lifecycle event history of a server-side job, as Arrow batches.
Lists history across all jobs when `job_id` is None.
"""
return await self._inner.job_history(job_id)
@property
def functions(self) -> "_AsyncFunctions":
"""First-class Function operations for this connection."""
from ._functions import _AsyncFunctions
return _AsyncFunctions(self)
async def _register_function(
self, name: str, definition: "_FunctionDefinition"
) -> NativeJob:
return await self._inner._register_function(name, definition)
async def _replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
return await self._inner._replace_function(name, current, definition)
async def _lookup_function_by_name(self, name: str) -> "Function":
return await self._inner._lookup_function_by_name(name)
async def _lookup_function_by_id(self, function_id: str) -> "Function":
return await self._inner._lookup_function_by_id(function_id)
async def _remove_function_name(self, name: str, current: "Function") -> None:
return await self._inner._remove_function_name(name, current)
async def _revoke_function(self, function: "Function") -> None:
return await self._inner._revoke_function(function)
async def namespace_client(self) -> LanceNamespace:
"""Get the equivalent namespace client for this connection.
@@ -101,7 +101,8 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
@weak_lru(maxsize=1)
def ndims(self):
return len(self.generate_embeddings([[self.source_instruction, "foo"]])[0])
model = self.get_model()
return model.encode("foo").shape[0]
def compute_query_embeddings(self, query: str, *args, **kwargs) -> List[np.array]:
return self.generate_embeddings([[self.query_instruction, query]])
+3 -4
View File
@@ -87,13 +87,12 @@ class JinaEmbeddings(EmbeddingFunction):
if isinstance(image, bytes):
image_dict = {"image": base64.b64encode(image).decode("utf-8")}
elif isinstance(image, (str, Path)):
parsed = urlparse(str(image))
parsed = urlparse.urlparse(image)
# TODO handle drive letter on windows.
PIL_Image = attempt_import_or_raise("PIL.Image", "pillow")
if parsed.scheme == "file":
pil_image = PIL_Image.open(parsed.path)
elif parsed.scheme == "" or (os.name == "nt" and len(parsed.scheme) == 1):
# A Windows drive letter parses as a one-character scheme
# ("C:\\img.png" -> scheme="c"), so treat it as a local path.
elif parsed.scheme == "":
pil_image = PIL_Image.open(image if os.name == "nt" else parsed.path)
elif parsed.scheme.startswith("http"):
pil_image = PIL_Image.open(io.BytesIO(url_retrieve(image)))
-49
View File
@@ -3,8 +3,6 @@
"""Custom exception handling"""
from typing import Optional
class MissingValueError(ValueError):
"""Exception raised when a required value is missing."""
@@ -25,50 +23,3 @@ class MissingColumnError(KeyError):
return (
f"Error: Column '{self.column_name}' does not exist in the DataFrame object"
)
class JobFailedError(RuntimeError):
"""Exception raised when an asynchronous job reaches the failed state.
``error_code`` is the optional exact category string projected from the
native job failure when the backend supplied one. The RuntimeError
message remains the existing diagnostic text and must not be used to
recover or override the code.
"""
__slots__ = ("_error_code",)
def __init__(self, message: str, error_code: Optional[str] = None) -> None:
super().__init__(message)
self._error_code = error_code
@property
def error_code(self) -> Optional[str]:
"""Exact job failure error category string, when supplied."""
return self._error_code
class JobCancelledError(RuntimeError):
"""Exception raised when an asynchronous job was cancelled."""
pass
class FunctionError(RuntimeError):
"""Exception raised when a first-class Function operation fails.
``code`` is the stable semantic category from the native error. The
message is a sanitized client diagnostic and must not be used to recover
or override the code.
"""
__slots__ = ("_code",)
def __init__(self, message: str, code: str) -> None:
super().__init__(message)
self._code = code
@property
def code(self) -> str:
"""Stable Function error category string."""
return self._code
-114
View File
@@ -1,114 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Handles to operations a server may run asynchronously."""
import asyncio
from datetime import timedelta
from typing import Optional
from lancedb.background_loop import LOOP
from . import _lancedb
from ._lancedb import Function
class AsyncJob:
"""A handle to an operation that may still be running.
The operation may already be complete when the handle is created.
"""
def __init__(self, inner: Optional["_lancedb.Job"]):
self._inner = inner
@property
def id(self) -> Optional[str]:
"""Identifies the operation on the server that is running it.
Returned for correlating with server logs or the jobs API. Operations
that run in this process have no server id and return `None`. The value
is opaque: parsing it or storing it to resume the job later is not
supported.
"""
return self._inner.id if self._inner is not None else None
async def status(self) -> str:
"""The operation's current lifecycle state: "running", "finished",
"failed", or "cancelled".
A point snapshot; unlike `wait` it does not block or raise on a
terminal failure state. States a newer server reports that this
client version does not know pass through as-is.
"""
if self._inner is None:
return "finished"
return await self._inner.status()
async def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
"""Wait until the operation reaches a terminal state.
Returns the success result when present (currently a
:class:`~lancedb.Function`), or `None` when the job finished without
one.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
if self._inner is None:
return None
if timeout is None:
return await self._inner.wait()
else:
return await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
async def cancel(self):
"""Request cancellation. Cancelling a finished operation is a no-op."""
if self._inner is None:
return
await self._inner.cancel()
class Job:
"""Synchronous counterpart of `AsyncJob`."""
def __init__(self, inner: Optional[AsyncJob]):
self._inner = inner
@property
def id(self) -> Optional[str]:
"""Identifies the operation on the server that is running it.
See :attr:`AsyncJob.id`.
"""
return self._inner.id if self._inner is not None else None
def status(self) -> str:
"""The operation's current lifecycle state: "running", "finished",
"failed", or "cancelled".
See :meth:`AsyncJob.status`.
"""
if self._inner is None:
return "finished"
return LOOP.run(self._inner.status())
def wait(self, timeout: Optional[timedelta] = None) -> Optional[Function]:
"""Block until the operation reaches a terminal state.
Returns the success result when present (currently a
:class:`~lancedb.Function`), or `None` when the job finished without
one.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
if self._inner is None:
return None
return LOOP.run(self._inner.wait(timeout))
def cancel(self):
"""Request cancellation. Cancelling a finished operation is a no-op."""
if self._inner is None:
return
LOOP.run(self._inner.cancel())
+1 -3
View File
@@ -92,10 +92,8 @@ class LanceMergeInsertBuilder(object):
self._when_not_matched_by_source_delete = True
if isinstance(condition, Expr):
self._when_not_matched_by_source_condition_expr = condition._inner
self._when_not_matched_by_source_condition = None
else:
elif condition is not None:
self._when_not_matched_by_source_condition = condition
self._when_not_matched_by_source_condition_expr = None
return self
def use_index(self, use_index: bool) -> LanceMergeInsertBuilder:
+1 -95
View File
@@ -38,11 +38,7 @@ from lance_namespace_urllib3_client.models.query_table_request_vector import (
QueryTableRequestVector,
)
from lance_namespace_urllib3_client.models.string_fts_query import StringFtsQuery
from lance_namespace.errors import (
NamespaceNotEmptyError,
NamespaceNotFoundError,
TableNotFoundError,
)
from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
from lancedb._lancedb import (
connect_namespace as _connect_namespace,
connect_namespace_client as _connect_namespace_client,
@@ -57,8 +53,6 @@ from lance_namespace import (
DropNamespaceResponse,
ListNamespacesResponse,
ListTablesResponse,
NamespaceExistsRequest,
TableExistsRequest,
)
from lancedb.table import AsyncTable, LanceTable, Table
from lancedb.util import validate_table_name
@@ -786,51 +780,6 @@ class LanceNamespaceDBConnection(DBConnection):
"""
return LOOP.run(self._inner.describe_namespace(namespace_path))
@override
def namespace_exists(self, namespace_id: List[str]) -> bool:
"""
Check if a namespace exists.
Parameters
----------
namespace_id : List[str]
The namespace identifier to check.
Returns
-------
bool
True if the namespace exists, False otherwise.
"""
request = NamespaceExistsRequest(id=namespace_id)
try:
self._namespace_client.namespace_exists(request)
return True
except NamespaceNotFoundError:
return False
@override
def table_exists(self, table_id: List[str]) -> bool:
"""
Check if a table exists.
Parameters
----------
table_id : List[str]
The table identifier to check (full path including namespace
segments and table name).
Returns
-------
bool
True if the table exists, False otherwise.
"""
request = TableExistsRequest(id=table_id)
try:
self._namespace_client.table_exists(request)
return True
except TableNotFoundError:
return False
@override
def list_tables(
self,
@@ -1284,49 +1233,6 @@ class AsyncLanceNamespaceDBConnection:
"""
return await self._inner.describe_namespace(namespace_path)
async def namespace_exists(self, namespace_id: List[str]) -> bool:
"""
Check if a namespace exists.
Parameters
----------
namespace_id : List[str]
The namespace identifier to check.
Returns
-------
bool
True if the namespace exists, False otherwise.
"""
request = NamespaceExistsRequest(id=namespace_id)
try:
self._namespace_client.namespace_exists(request)
return True
except NamespaceNotFoundError:
return False
async def table_exists(self, table_id: List[str]) -> bool:
"""
Check if a table exists.
Parameters
----------
table_id : List[str]
The table identifier to check (full path including namespace
segments and table name).
Returns
-------
bool
True if the table exists, False otherwise.
"""
request = TableExistsRequest(id=table_id)
try:
self._namespace_client.table_exists(request)
return True
except TableNotFoundError:
return False
async def list_tables(
self,
namespace_path: Optional[List[str]] = None,
+1 -1
View File
@@ -226,7 +226,7 @@ class PermutationBuilder:
async def do_execute():
inner_tbl = await self._async.execute()
return await LanceTable.from_inner(inner_tbl)
return LanceTable.from_inner(inner_tbl)
return LOOP.run(do_execute())
-1
View File
@@ -1 +0,0 @@
-10
View File
@@ -153,16 +153,6 @@ def Vector(
return FixedSizeList
def _raise_bare_vector_error(*_args):
raise TypeError("Vector must be parameterized with a dimension, e.g. Vector(128).")
# Pydantic v1 and v2 otherwise treat the bare Vector factory as a field validator
# and inspect its signature, which produces misleading errors about internal types.
setattr(Vector, "__get_validators__", _raise_bare_vector_error)
setattr(Vector, "__get_pydantic_core_schema__", _raise_bare_vector_error)
def MultiVector(
dim: int, value_type: pa.DataType = pa.float32(), nullable: bool = True
) -> Type:
+10 -3
View File
@@ -52,6 +52,7 @@ from ._blob import (
finalize_blob_query_table,
replace_v2_blob_columns_with_bytes,
replace_v2_blob_columns_with_bytes_sync,
supports_blob_auto_row_id,
validate_blob_mode,
)
from .types import BlobMode, QueryProjection
@@ -1279,7 +1280,10 @@ class LanceQueryBuilder(ABC):
return self._with_row_id is True
def _blob_auto_row_id_enabled(self) -> bool:
if not supports_blob_auto_row_id(self._table):
return False
return blob_auto_row_id_for_scan(
self._table,
self._table.schema,
self._columns,
with_row_id=self._with_row_id,
@@ -2697,7 +2701,7 @@ class LanceHybridQueryBuilder(LanceQueryBuilder):
self._fts_query.phrase_query(True)
if self._distance_type:
self._vector_query.metric(self._distance_type)
if self._minimum_nprobes is not None:
if self._minimum_nprobes:
self._vector_query.minimum_nprobes(self._minimum_nprobes)
if self._maximum_nprobes is not None:
self._vector_query.maximum_nprobes(self._maximum_nprobes)
@@ -2770,7 +2774,7 @@ class AsyncQueryBase(object):
)
async def _maybe_add_blob_row_id(self) -> None:
if self._table is None:
if self._table is None or not supports_blob_auto_row_id(self._table):
self._blob_auto_row_id = False
self._blob_paths = ()
return
@@ -2778,6 +2782,7 @@ class AsyncQueryBase(object):
req = self._inner.to_query_request()
schema = await self._table.schema()
self._blob_auto_row_id = blob_auto_row_id_for_scan(
self._table,
schema,
req.select,
with_row_id=self._with_row_id,
@@ -3029,6 +3034,7 @@ class AsyncQueryBase(object):
schema = await self._table.schema()
blob_auto_row_id = blob_auto_row_id_for_scan(
self._table,
schema,
query.columns,
with_row_id=self._with_row_id,
@@ -3874,9 +3880,10 @@ class AsyncHybridQuery(AsyncStandardQuery, AsyncVectorQueryBase):
req = fts_query._inner.to_query_request()
blob_auto_row_id = False
blob_paths: tuple[str, ...] = ()
if self._table is not None:
if self._table is not None and supports_blob_auto_row_id(self._table):
schema = await self._table.schema()
blob_auto_row_id = blob_auto_row_id_for_scan(
self._table,
schema,
req.select,
with_row_id=self._with_row_id,
+1 -81
View File
@@ -7,7 +7,7 @@ import json
import logging
from concurrent.futures import ThreadPoolExecutor
import sys
from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Union
from typing import Any, Dict, Iterable, List, Optional, Union
from urllib.parse import urlparse
import warnings
@@ -23,12 +23,6 @@ import pyarrow as pa
from ..common import DATA
from ..db import DBConnection, LOOP
from ..job import Job
if TYPE_CHECKING:
from .._lancedb import Function
from .._lancedb import Job as NativeJob
from .._lancedb import JobDescription, JobInfo, _FunctionDefinition
from ..embeddings import EmbeddingFunctionConfig
from lance_namespace import (
LanceNamespace,
@@ -421,11 +415,6 @@ class RemoteDBConnection(DBConnection):
if namespace_path is None:
namespace_path = []
if storage_options is not None:
logging.info(
"storage_options is ignored in LanceDb Cloud"
" (storage is managed; set storage_options on connect() instead)"
)
if index_cache_size is not None:
logging.info(
"index_cache_size is ignored in LanceDb Cloud"
@@ -695,75 +684,6 @@ class RemoteDBConnection(DBConnection):
)
)
@override
def job(self, job_id: str) -> Job:
"""A [Job][lancedb.job.Job] handle for a server-side job by id.
The handle is constructed without a server round trip; an unknown id
surfaces when the handle is used. Dropping the handle has no effect
on the job itself.
"""
return Job(self._conn.job(job_id))
@override
def list_jobs(self) -> List["JobInfo"]:
"""List server-side jobs across the database's tables."""
return LOOP.run(self._conn.list_jobs())
@override
def get_job(self, job_id: str) -> Optional["JobDescription"]:
"""Describe a single server-side job by id.
Returns None when the server has no such job.
"""
return LOOP.run(self._conn.get_job(job_id))
@override
def cancel_job(self, job_id: str) -> bool:
"""Request cancellation of a server-side job by id.
Returns True if the server accepted the cancellation, False if no
such job exists. Cancelling an already-terminal job is a no-op
success.
"""
return LOOP.run(self._conn.cancel_job(job_id))
@override
def job_history(self, job_id: Optional[str] = None) -> List[pa.RecordBatch]:
"""The lifecycle event history of a server-side job, as Arrow batches.
Lists history across all jobs when `job_id` is None.
"""
return LOOP.run(self._conn.job_history(job_id))
@override
def _submit_register_function(
self, name: str, definition: "_FunctionDefinition"
) -> "NativeJob":
return LOOP.run(self._conn._register_function(name, definition))
@override
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> "NativeJob":
return LOOP.run(self._conn._replace_function(name, current, definition))
@override
def _lookup_function_by_name(self, name: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_name(name))
@override
def _lookup_function_by_id(self, function_id: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_id(function_id))
@override
def _remove_function_name(self, name: str, current: "Function") -> None:
return LOOP.run(self._conn._remove_function_name(name, current))
@override
def _revoke_function(self, function: "Function") -> None:
return LOOP.run(self._conn._revoke_function(function))
@override
def namespace_client(self) -> LanceNamespace:
"""Get the equivalent namespace client for this connection.
+9 -82
View File
@@ -7,7 +7,6 @@ import logging
from functools import cached_property
import os
from typing import (
TYPE_CHECKING,
Any,
Callable,
Dict,
@@ -21,7 +20,6 @@ from typing import (
import warnings
from lancedb import __version__
from lancedb._blob import BlobFile
from lancedb._lancedb import (
AddColumnsResult,
@@ -49,7 +47,6 @@ from lancedb.index import (
IvfSq,
LabelList,
)
from lancedb.job import Job
from lancedb.remote.db import LOOP
from lancedb.table import IndexConfigType, KNOWN_METRICS
import pyarrow as pa
@@ -68,9 +65,6 @@ from ..query import (
from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Table, Tags
from ..types import BaseTokenizerType
if TYPE_CHECKING:
from lancedb._lancedb import _FunctionCall
class RemoteTable(Table):
def __init__(
@@ -546,73 +540,6 @@ class RemoteTable(Table):
)
)
def create_index_async(
self,
column: str,
*,
config: IndexConfigType,
replace: Optional[bool] = None,
wait_timeout: Optional[timedelta] = None,
name: Optional[str] = None,
train: bool = True,
) -> Job:
"""Create an index, returning a handle to the indexing job.
The job may already be complete when returned; callers must not assume
the index exists until :meth:`Job.wait` returns.
"""
return Job(
LOOP.run(
self._table.create_index_async(
column,
replace=replace,
config=config,
wait_timeout=wait_timeout,
name=name,
train=train,
)
)
)
def add_generated_column(self, column_name: str, call: "_FunctionCall") -> Job:
"""Add a generated column from an authored Function call.
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
of the Job does not publish the column; callers must wait and re-read
the table to observe the new definition and values.
"""
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
return LOOP.run(self._table.generated_column_status(column_name))
def refresh_generated_column(self, column_name: str) -> Job:
"""Refresh values for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
of the Job does not publish new values; callers must wait and re-read
the table to observe refreshed results.
"""
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
def alter_generated_column(
self, column_name: str, new_call: "_FunctionCall"
) -> Job:
"""Alter the Function call for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
of the Job does not publish the new definition; callers must wait and
re-read the table to observe the updated column.
"""
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
def _is_legacy_create_index_call(
self,
first_arg: str,
@@ -1112,22 +1039,22 @@ class RemoteTable(Table):
)
def blob_columns(self) -> list[str]:
return LOOP.run(self._table.blob_columns())
raise NotImplementedError(
"blob_columns() is not yet supported on the LanceDB Cloud"
)
def fetch_blobs(
self, column: str, row_ids: Union[list[int], pa.Table]
) -> pa.LargeBinaryArray:
return LOOP.run(self._table.fetch_blobs(column, row_ids))
def fetch_blobs(self, column: str, row_ids) -> pa.LargeBinaryArray:
raise NotImplementedError("fetch_blobs() is not supported on LanceDB Cloud")
def fetch_blob_ranges(self, column: str, requests) -> pa.LargeBinaryArray:
raise NotImplementedError(
"fetch_blob_ranges() is not supported on LanceDB Cloud"
)
def fetch_blob_files(
self, column: str, row_ids: Union[list[int], pa.Table]
) -> "list[Optional[BlobFile]]":
return LOOP.run(self._table.fetch_blob_files(column, row_ids))
def fetch_blob_files(self, column: str, row_ids):
raise NotImplementedError(
"fetch_blob_files() is not supported on LanceDB Cloud"
)
def head(self, n=5) -> pa.Table:
"""
+27 -315
View File
@@ -11,11 +11,6 @@ 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
@@ -27,7 +22,7 @@ import time
from collections import deque
from concurrent.futures import ThreadPoolExecutor
from multiprocessing import RawArray
from typing import Any, Callable, Iterator, Optional, Union
from typing import Any, Callable, Iterator, Optional
from torch.utils.data import IterableDataset, get_worker_info
@@ -132,49 +127,6 @@ 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
@@ -200,7 +152,6 @@ 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,
):
@@ -216,13 +167,6 @@ 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
@@ -238,7 +182,6 @@ 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
@@ -256,28 +199,19 @@ 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,
# rows_skipped]
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 8)
# bytes_loaded, fetch_time_us, transform_time_us]
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 7)
# 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)
@@ -341,7 +275,6 @@ 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
@@ -349,20 +282,14 @@ class StreamingDataset(IterableDataset):
if self._columns is not None:
perm = perm.select_columns(self._columns)
perm = perm.with_transform(lambda batch: batch)
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)
if self._resume_offset > 0:
perm = perm.with_skip(self._resume_offset)
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
@@ -375,14 +302,12 @@ class StreamingDataset(IterableDataset):
self._transform if self._transform is not None else Transforms.arrow2python
)
# 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.
# Per-split pipeline state.
fetch_head = [0] * n
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
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
# Limit simultaneous transforms to transform_workers across all splits.
tx_semaphore = threading.Semaphore(transform_workers)
@@ -405,8 +330,7 @@ class StreamingDataset(IterableDataset):
fetch_head[i] += fetch
perm_i = permutations[i]
indices = list(range(start, start + fetch))
abs_start = initial_positions[i] + start
io_pending[i].append((abs_start, io_pool.submit(_io_call, perm_i, indices)))
io_pending[i].append(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]:
@@ -414,72 +338,15 @@ 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][1].done():
abs_start, fut = io_pending[i].popleft()
raw_batches[i].append((abs_start, fut.result()))
while io_pending[i] and io_pending[i][0].done():
raw_batches[i].append(io_pending[i].popleft().result())
# ── Stage 2 helpers ───────────────────────────────────────────────────
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):
def _tx_call_guarded(batch):
try:
t0 = time.perf_counter()
result = _transform_batch(abs_start, batch)
result = final_transform(batch)
self._transform_time += time.perf_counter() - t0
return result
finally:
@@ -488,8 +355,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):
abs_start, batch = raw_batches[i].popleft()
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, abs_start, batch))
batch = raw_batches[i].popleft()
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch))
def _drain_tx(i: int) -> None:
"""Move completed transform futures into cooked non-blockingly."""
@@ -517,14 +384,11 @@ class StreamingDataset(IterableDataset):
# Acquire a transform slot (may block briefly if all
# transform_workers are busy with other splits).
tx_semaphore.acquire()
abs_start, batch = raw_batches[i].popleft()
tx_pending[i].append(
tx_pool.submit(_tx_call_guarded, abs_start, batch)
)
batch = raw_batches[i].popleft()
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, batch))
elif io_pending[i]:
# Block on the oldest in-flight I/O fetch.
abs_start, fut = io_pending[i].popleft()
raw_batches[i].append((abs_start, fut.result()))
raw_batches[i].append(io_pending[i].popleft().result())
_advance(i)
else:
break # split exhausted
@@ -543,28 +407,15 @@ class StreamingDataset(IterableDataset):
_fill_io(i)
while True:
# 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:
# 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)):
break
for i in range(n):
pos, row = cooked[i].popleft()
_ensure_cooked(i)
row = cooked[i].popleft()
local_consumed[i] += 1
pos_consumed[i] = pos + 1
_advance(i)
# After the last split in each cycle: update the
@@ -573,39 +424,21 @@ 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
@@ -659,7 +492,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
@@ -689,19 +522,6 @@ 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.
@@ -767,27 +587,12 @@ 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:
@@ -813,96 +618,3 @@ 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
+21 -337
View File
@@ -40,7 +40,6 @@ from ._blob import (
from .types import BlobMode
from lancedb.arrow import peek_reader
from lancedb.background_loop import LOOP, embedding_executor
from lancedb.job import AsyncJob, Job
from .dependencies import (
_check_for_hugging_face,
_check_for_lance,
@@ -108,11 +107,6 @@ def _should_push_down_query_table(
return namespace_client is not None and "QueryTable" in pushdown_operations
def _polars_predicate_pushdown_barrier(frame: Any) -> Any:
"""Return a Polars frame unchanged while blocking predicate pushdown."""
return frame
_MODEL_BACKED_TOKENIZER_PREFIXES = ("jieba", "lindera")
_MODEL_BACKED_TOKENIZER_ERRORS = (
"unknown base tokenizer",
@@ -185,7 +179,6 @@ if TYPE_CHECKING:
LsmWriteSpec,
MergeResult,
UpdateResult,
_FunctionCall,
)
from .index import IndexConfig
import pandas
@@ -870,18 +863,12 @@ class Table(ABC):
"""
raise NotImplementedError
def to_polars(self, **kwargs) -> "pl.LazyFrame":
"""Return the table as a Polars LazyFrame.
Note
----
The Polars streaming engine is not supported because it does not currently
implement Python PyArrow dataset scans. Use the default engine when collecting
this LazyFrame.
def to_polars(self, **kwargs) -> "pl.DataFrame":
"""Return the table as a polars.DataFrame.
Returns
-------
polars.LazyFrame
polars.DataFrame
"""
raise NotImplementedError
@@ -990,61 +977,6 @@ class Table(ABC):
"""
raise NotImplementedError
def create_index_async(
self,
column: str,
*,
config: IndexConfigType,
replace: Optional[bool] = None,
wait_timeout: Optional[timedelta] = None,
name: Optional[str] = None,
train: bool = True,
) -> Job:
"""Create an index, returning a handle to the indexing job.
Takes the same arguments as :meth:`create_index`. The job may already
be complete when returned; callers must not assume the index exists
until :meth:`Job.wait` returns.
"""
raise NotImplementedError
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
"""Add a generated column from an authored Function call.
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
of the Job does not publish the column; callers must wait and re-read
the table to observe the new definition and values.
"""
raise NotImplementedError
def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
raise NotImplementedError
def refresh_generated_column(self, column_name: str) -> Job:
"""Refresh values for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
of the Job does not publish new values; callers must wait and re-read
the table to observe refreshed results.
"""
raise NotImplementedError
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
"""Alter the Function call for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
of the Job does not publish the new definition; callers must wait and
re-read the table to observe the updated column.
"""
raise NotImplementedError
def drop_index(self, name: str) -> None:
"""
Drop an index from the table.
@@ -1642,10 +1574,8 @@ class Table(ABC):
"""Open lazy, seekable :class:`~lancedb._blob.BlobFile` handles.
Prefer this over :meth:`fetch_blobs` for large payloads. ``row_ids`` is
a ``list[int]`` or a query ``pyarrow.Table`` carrying row identity via
``_rowid`` or a ``_lance_row_id`` field on the blob descriptor. Null
rows are ``None``. Remote tables require LanceDB Cloud server 0.5.0 or
newer.
a ``list[int]`` or query ``pyarrow.Table`` with ``_rowid`` (or stashed
row-id metadata). Null rows are ``None``. Local tables only.
"""
@abstractmethod
@@ -2231,15 +2161,11 @@ class LanceTable(Table):
return self.name
@classmethod
async def from_inner(cls, tbl: LanceDBTable):
from .db import AsyncConnection, LanceDBConnection
def from_inner(cls, tbl: LanceDBTable):
from .db import LanceDBConnection
async_tbl = AsyncTable(tbl)
inner_conn = tbl.database()
read_consistency_interval = await AsyncConnection(
inner_conn
).get_read_consistency_interval()
conn = LanceDBConnection.from_inner(inner_conn, read_consistency_interval)
conn = LanceDBConnection.from_inner(tbl.database())
return cls(
conn,
async_tbl.name,
@@ -2543,7 +2469,13 @@ class LanceTable(Table):
return LOOP.run(self._table.count_rows(filter))
def __repr__(self) -> str:
return f"{self.__class__.__name__}(name={self.name!r}, _conn={self._conn!r})"
val = f"{self.__class__.__name__}(name={self.name!r}"
if self._conn.read_consistency_interval is not None:
val += ", read_consistency_interval={!r}".format(
self._conn.read_consistency_interval
)
val += f", _conn={self._conn!r})"
return val
def __str__(self) -> str:
return self.__repr__()
@@ -2618,9 +2550,6 @@ class LanceTable(Table):
2. Currently we've disabled push-down of the filters from polars
because polars pushdown into pyarrow uses pyarrow compute
expressions rather than SQl strings (which LanceDB supports)
3. The Polars streaming engine is not supported because it does not
currently implement Python PyArrow dataset scans. Use the default
engine when collecting this LazyFrame.
Returns
-------
@@ -2629,12 +2558,8 @@ class LanceTable(Table):
from lancedb.integrations.pyarrow import PyarrowDatasetAdapter
dataset = PyarrowDatasetAdapter(self)
# Polars 1.32's non-PyArrow callback path passes batch_size twice. Keep
# the compatible PyArrow path, but block predicates because this adapter
# cannot translate PyArrow expressions into LanceDB filters.
return pl.scan_pyarrow_dataset(dataset, batch_size=batch_size).map_batches(
_polars_predicate_pushdown_barrier,
predicate_pushdown=False,
return pl.scan_pyarrow_dataset(
dataset, allow_pyarrow_filter=False, batch_size=batch_size
)
# New unified API overload
@@ -2859,71 +2784,6 @@ class LanceTable(Table):
)
)
def create_index_async(
self,
column: str,
*,
config: IndexConfigType,
replace: Optional[bool] = None,
wait_timeout: Optional[timedelta] = None,
name: Optional[str] = None,
train: bool = True,
) -> Job:
"""Create an index, returning a handle to the indexing job.
The job may already be complete when returned; callers must not assume
the index exists until :meth:`Job.wait` returns.
"""
return Job(
LOOP.run(
self._table.create_index_async(
column,
replace=replace,
config=config,
wait_timeout=wait_timeout,
name=name,
train=train,
)
)
)
def add_generated_column(self, column_name: str, call: _FunctionCall) -> Job:
"""Add a generated column from an authored Function call.
Returns a :class:`~lancedb.job.Job` for the create operation. Acceptance
of the Job does not publish the column; callers must wait and re-read
the table to observe the new definition and values.
"""
return Job(LOOP.run(self._table.add_generated_column(column_name, call)))
def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
return LOOP.run(self._table.generated_column_status(column_name))
def refresh_generated_column(self, column_name: str) -> Job:
"""Refresh values for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the refresh operation. Acceptance
of the Job does not publish new values; callers must wait and re-read
the table to observe refreshed results.
"""
return Job(LOOP.run(self._table.refresh_generated_column(column_name)))
def alter_generated_column(self, column_name: str, new_call: _FunctionCall) -> Job:
"""Alter the Function call for an existing generated column.
Returns a :class:`~lancedb.job.Job` for the change operation. Acceptance
of the Job does not publish the new definition; callers must wait and
re-read the table to observe the updated column.
"""
return Job(LOOP.run(self._table.alter_generated_column(column_name, new_call)))
def _is_legacy_create_index_call(
self,
first_arg: str,
@@ -4051,28 +3911,6 @@ class LanceTable(Table):
[`AsyncTable.get_lsm_write_spec`][lancedb.AsyncTable.get_lsm_write_spec]."""
return LOOP.run(self._table.get_lsm_write_spec())
def checkpoint_lsm(self) -> None:
"""Synchronous version of
[`AsyncTable.checkpoint_lsm`][lancedb.AsyncTable.checkpoint_lsm]."""
return LOOP.run(self._table.checkpoint_lsm())
def flush_lsm(self) -> None:
"""Synchronous version of
[`AsyncTable.flush_lsm`][lancedb.AsyncTable.flush_lsm]."""
return LOOP.run(self._table.flush_lsm())
def compact_lsm(self) -> None:
"""Synchronous version of
[`AsyncTable.compact_lsm`][lancedb.AsyncTable.compact_lsm]."""
return LOOP.run(self._table.compact_lsm())
def get_lsm_stats(self, *, include_generation_rows: bool = False) -> Optional[dict]:
"""Synchronous version of
[`AsyncTable.get_lsm_stats`][lancedb.AsyncTable.get_lsm_stats]."""
return LOOP.run(
self._table.get_lsm_stats(include_generation_rows=include_generation_rows)
)
def close_lsm_writers(self) -> None:
"""Close cached MemWAL shard writers. See
[`AsyncTable.close_lsm_writers`][lancedb.AsyncTable.close_lsm_writers]."""
@@ -4751,13 +4589,6 @@ class AsyncTable:
via [`set_unenforced_primary_key`]; bucket sharding additionally
requires it to be the single column being bucketed.
By default the MemWAL maintains every index on the table, resolved
here a snapshot, so an index created afterwards needs the spec unset
and set again. This fails if one cannot be maintained; name the set
with ``with_maintained_indexes`` to install anyway. That pins an exact
set (a still-building index is rejected, not omitted); ``[]`` maintains
none.
Parameters
----------
spec : LsmWriteSpec
@@ -4784,73 +4615,12 @@ class AsyncTable:
Returns ``None`` when the MemWAL LSM write path is not enabled (no
spec has been set, or it was removed with `unset_lsm_write_spec`).
The returned spec mirrors what was passed to `set_lsm_write_spec`,
except that ``maintained_indexes`` always reports the concrete list
resolved when the spec was set ``None`` never round-trips.
The returned spec including its ``maintained_indexes`` and
``writer_config_defaults`` mirrors what was passed to
`set_lsm_write_spec`.
"""
return await self._inner.get_lsm_write_spec()
async def checkpoint_lsm(self) -> None:
"""Converge this table's LSM write path into its base table.
One flush, sealing every memtable into L0, then compaction triggers
until every generation that existed at that moment has reached base.
The loop runs client-side, reading progress from ``get_lsm_stats``.
Best-effort: generations created *while* it runs are deliberately not
waited on, which is what lets it terminate on a table taking writes.
Idempotent and safe on a cadence.
There is no deadline, and the caller owns that. It returns when the
target generations are gone, raises on a terminal server fault, and
otherwise waits however long the server takes. A slow table and a
stuck one are the same picture from the client: the compactor pool is
shared across every table on the node, so a checkpoint queued behind
unrelated work looks exactly like one that is merging. Wrap this in
``asyncio.wait_for`` for a wall-clock bound; abandoning it partway
costs nothing.
"""
return await self._inner.checkpoint_lsm()
async def flush_lsm(self) -> None:
"""Seal every bucket's active memtable into L0.
Does not touch the base table moving L0 into base is
`compact_lsm`. On a node that has not claimed this table, this claims
it and replays its WAL log first.
"""
return await self._inner.flush_lsm()
async def compact_lsm(self) -> None:
"""Trigger a background L0 to base compaction pass per bucket.
Returns once the passes are dispatched, not once they finish: watch
``get_lsm_stats`` for progress, or use ``checkpoint_lsm`` to loop
until the current L0 has reached base.
"""
return await self._inner.compact_lsm()
async def get_lsm_stats(
self, *, include_generation_rows: bool = False
) -> Optional[dict]:
"""Read live per-bucket LSM state.
Answers "how far behind is my fresh tier", "which bucket is hot", and
"why is my fresh-tier vector search brute-force". Mutates no table
state, though on a node that has not claimed this table it claims it,
exactly as a read would.
Returns ``None`` only when the LSM write path is not enabled.
Parameters
----------
include_generation_rows
Report a row count per L0 generation. Off by default: each count
opens an uncached Lance dataset, and ``checkpoint_lsm`` polls this
needing only generation numbers.
"""
return await self._inner.get_lsm_stats(include_generation_rows)
async def close_lsm_writers(self) -> None:
"""Drain and close any cached MemWAL shard writers for this table.
@@ -5101,90 +4871,6 @@ class AsyncTable:
)
raise e
async def create_index_async(
self,
column: str,
*,
replace: Optional[bool] = None,
config: Optional[
Union[
IvfFlat,
IvfPq,
IvfRq,
HnswPq,
HnswSq,
HnswFlat,
BTree,
Bitmap,
LabelList,
Fm,
FTS,
]
] = None,
wait_timeout: Optional[timedelta] = None,
name: Optional[str] = None,
train: bool = True,
) -> AsyncJob:
"""Create an index, returning a handle to the indexing job.
Takes the same arguments as :meth:`create_index`. The job may already
be complete when returned; callers must not assume the index exists
until :meth:`AsyncJob.wait` resolves.
"""
job = await self._inner.create_index_async(
column,
index=config,
replace=replace,
wait_timeout=wait_timeout,
name=name,
train=train,
)
return AsyncJob(job)
async def add_generated_column(
self, column_name: str, call: _FunctionCall
) -> AsyncJob:
"""Add a generated column from an authored Function call.
Returns an :class:`~lancedb.job.AsyncJob` for the create operation.
Acceptance of the Job does not publish the column; callers must wait
and re-read the table to observe the new definition and values.
"""
job = await self._inner._add_generated_column(column_name, call)
return AsyncJob(job)
async def generated_column_status(
self, column_name: str
) -> Literal["complete", "incomplete"]:
"""Return ``"complete"`` or ``"incomplete"`` for a generated column.
Projection-only: reads the named column's stored definition status.
Does not refresh values, submit a Job, or mutate table state.
"""
return await self._inner._generated_column_status(column_name)
async def refresh_generated_column(self, column_name: str) -> AsyncJob:
"""Refresh values for an existing generated column.
Returns an :class:`~lancedb.job.AsyncJob` for the refresh operation.
Acceptance of the Job does not publish new values; callers must wait
and re-read the table to observe refreshed results.
"""
job = await self._inner._refresh_generated_column(column_name)
return AsyncJob(job)
async def alter_generated_column(
self, column_name: str, new_call: _FunctionCall
) -> AsyncJob:
"""Alter the Function call for an existing generated column.
Returns an :class:`~lancedb.job.AsyncJob` for the change operation.
Acceptance of the Job does not publish the new definition; callers must
wait and re-read the table to observe the updated column.
"""
job = await self._inner._alter_generated_column(column_name, new_call)
return AsyncJob(job)
async def drop_index(self, name: str) -> None:
"""
Drop an index from the table.
@@ -6460,9 +6146,7 @@ class TableStatistics:
Attributes
----------
total_bytes: int
The total size, in bytes, of the table's data files, index files, and
overlay files. Read from the manifest, so this excludes deletion files
and manifests.
The total number of bytes in the table.
num_rows: int
The total number of rows in the table.
num_indices: int
-5
View File
@@ -395,11 +395,6 @@ def _(value: dict):
)
@value_to_sql.register(pa.Scalar)
def _(value: pa.Scalar):
return value_to_sql(value.as_py())
@value_to_sql.register(np.ndarray)
def _(value: np.ndarray):
return value_to_sql(value.tolist())
+3 -3
View File
@@ -226,13 +226,13 @@ def test_fetch_blob_ranges_validates_requests():
table = _blob_table("range_validation", [{"id": 1, "image": b"abc"}])
row_id = _row_ids_by_id(table)[1]
with pytest.raises(ValueError, match="exceeds blob size"):
with pytest.raises(RuntimeError, match="exceeds blob size"):
table.fetch_blob_ranges("image", [(row_id, 2, 2)])
with pytest.raises(ValueError, match="offset \\+ length overflowed"):
with pytest.raises(RuntimeError, match="offset \\+ length overflowed"):
table.fetch_blob_ranges("image", [(row_id, 2**64 - 1, 1)])
with pytest.raises(ValueError, match="row IDs"):
with pytest.raises(ValueError, match="row ids"):
table.fetch_blob_ranges("image", [(2**64 - 1, 0, 1)])
-44
View File
@@ -2,11 +2,9 @@
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import inspect
import re
import sys
from datetime import timedelta
from importlib import resources
import os
from types import SimpleNamespace
@@ -19,10 +17,6 @@ from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
from lancedb.pydantic import LanceModel, Vector
def test_package_includes_pep_561_marker():
assert resources.files(lancedb).joinpath("py.typed").is_file()
def test_basic(tmp_path):
db = lancedb.connect(tmp_path)
@@ -68,44 +62,6 @@ def test_basic(tmp_path):
assert db.open_table("test").name == db["test"].name
def test_sync_debugger_inspection_does_not_use_background_loop(tmp_path, monkeypatch):
from lancedb.background_loop import LOOP
db = lancedb.connect(tmp_path)
table = db.create_table("test", data=[{"id": 1}])
def fail_run(*args, **kwargs):
raise AssertionError("debugger inspection should not use the background loop")
monkeypatch.setattr(LOOP, "run", fail_run)
# Debuggers enumerate and evaluate every exposed attribute when expanding a
# variable. This must remain safe while their breakpoint suspends LOOP's thread.
members = dict(inspect.getmembers(db))
assert members["uri"] == str(tmp_path)
assert members["read_consistency_interval"] is None
assert repr(db) == f"LanceDBConnection(uri={str(tmp_path)!r})"
assert repr(table) == f"LanceTable(name='test', _conn={db!r})"
def test_read_consistency_interval_does_not_use_background_loop(tmp_path, monkeypatch):
from lancedb.background_loop import LOOP
from lancedb.db import LanceDBConnection
consistency_interval = timedelta(seconds=5)
db = lancedb.connect(tmp_path, read_consistency_interval=consistency_interval)
db_from_inner = LanceDBConnection.from_inner(db._inner, consistency_interval)
def fail_run(*args, **kwargs):
raise AssertionError("properties should not use the Python background loop")
monkeypatch.setattr(LOOP, "run", fail_run)
assert db.read_consistency_interval == consistency_interval
assert db_from_inner.read_consistency_interval == consistency_interval
def test_ingest_pd(tmp_path):
db = lancedb.connect(tmp_path)
@@ -1456,408 +1456,6 @@ 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(
+27 -51
View File
@@ -64,23 +64,6 @@ def test_embedding_function(tmp_path):
assert np.allclose(actual, expected)
def test_instructor_ndims_uses_instruction():
instructor = get_registry().get("instructor").create()
model = MagicMock()
model.encode.return_value = np.zeros((1, 384))
with patch.object(type(instructor), "get_model", return_value=model):
assert instructor.ndims() == 384
model.encode.assert_called_once_with(
[[instructor.source_instruction, "foo"]],
batch_size=instructor.batch_size,
show_progress_bar=instructor.show_progress_bar,
normalize_embeddings=instructor.normalize_embeddings,
device=instructor.device,
)
def test_embedding_function_variables():
@register("variable-testing")
class VariableTestingFunction(TextEmbeddingFunction):
@@ -132,16 +115,34 @@ def test_embedding_function_variables():
assert func.safe_model_dump()["secret_key"] == "$var:secret"
def test_openai_variables_survive_metadata_round_trip():
def test_parse_functions_with_variables():
@register("variable-parsing-test")
class VariableParsingFunction(TextEmbeddingFunction):
api_key: str
base_url: Optional[str] = None
@staticmethod
def sensitive_keys():
return ["api_key"]
def ndims(self):
return 10
def generate_embeddings(self, texts):
# Mock implementation that just returns random embeddings
# In real usage, this would use the api_key to call an API
return [np.random.rand(self.ndims()).tolist() for _ in texts]
registry = EmbeddingFunctionRegistry.get_instance()
registry.set_var("test_api_key", "sk-test-key-12345")
registry.set_var("test_base_url", "https://api.example.com")
conf = EmbeddingFunctionConfig(
source_column="text",
vector_column="vector",
function=registry.get("openai").create(
api_key="$var:test_api_key", base_url="https://api.example.com"
function=registry.get("variable-parsing-test").create(
api_key="$var:test_api_key", base_url="$var:test_base_url"
),
)
@@ -149,10 +150,7 @@ def test_openai_variables_survive_metadata_round_trip():
# Create a mock arrow table with the metadata
schema = pa.schema(
[
pa.field("text", pa.string()),
pa.field("vector", pa.list_(pa.float32(), 1536)),
]
[pa.field("text", pa.string()), pa.field("vector", pa.list_(pa.float32(), 10))]
)
table = pa.table({"text": [], "vector": []}, schema=schema)
table = table.replace_schema_metadata(metadata)
@@ -166,15 +164,13 @@ def test_openai_variables_survive_metadata_round_trip():
assert parsed_func.api_key == "sk-test-key-12345"
assert parsed_func.base_url == "https://api.example.com"
embeddings = parsed_func.generate_embeddings(["test text"])
assert len(embeddings) == 1
assert len(embeddings[0]) == 10
assert parsed_func.safe_model_dump()["api_key"] == "$var:test_api_key"
with patch("lancedb.embeddings.openai.attempt_import_or_raise") as import_openai:
parsed_func._openai_client
import_openai.return_value.OpenAI.assert_called_once_with(
api_key="sk-test-key-12345", base_url="https://api.example.com"
)
def test_embedding_with_bad_results(tmp_path):
@register("null-embedding")
@@ -631,23 +627,3 @@ def test_url_retrieve_downloads_image():
image_bytes = url_retrieve(image_url)
img = Image.open(io.BytesIO(image_bytes))
assert img.size[0] > 0 and img.size[1] > 0
def test_jina_generate_image_input_dict_local_path(tmp_path):
"""
JinaEmbeddings._generate_image_input_dict must accept a local image path
(str or Path), not just bytes. Previously it crashed with
`AttributeError: 'function' object has no attribute 'urlparse'` on any
str/Path input because it called `urlparse.urlparse(image)` instead of
`urlparse(image)` (urlparse was imported as a function, not a module).
"""
Image = pytest.importorskip("PIL.Image")
from lancedb.embeddings.jinaai import JinaEmbeddings
image_path = tmp_path / "test.png"
Image.new("RGB", (4, 4), color="red").save(image_path, format="PNG")
for image in (str(image_path), image_path):
image_dict = JinaEmbeddings._generate_image_input_dict(image)
assert "image" in image_dict
assert isinstance(image_dict["image"], str) and len(image_dict["image"]) > 0
@@ -1,372 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python exact Function handle call authoring (FF-028)."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.expr import Expr, col, func, lit
_CALL_PATH = "/v1/functions/lookup"
_CALL_CATALOG_NAME = "text.normalize.call-name"
_CALL_FUNCTION_ID = "fn.exact.call-handle"
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_call_xyz_42"
_INT_PAYLOAD_SENTINEL = 2_147_000_123
# Pinned Rust-canonical schema-only type IPC (base64).
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_LIST_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////+4AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAABAAAANz///8c"
"AAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECH"
"AAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////8AAAAAFAAAAAAAAAAMABQAEgAMAAgABAAMAAAAnAAAAKAAAAAQAAAAAAAEAAgACAAAAAQACAAAAAQAAAA"
"BAAAABAAAANz///8cAAAADAAAAAAAAQxcAAAAAQAAABwAAAAEAAQABAAAABAAFAAQAA4ADwAEAAAACAAQAAAA"
"GAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAQAAABpdGVtAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAAAwAAAAEFSUk9XMQ=="
)
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"name",
"connection",
"table",
"snapshot",
"field_id",
"field_ids",
"job",
"job_id",
"artifact",
"digest",
"retry_key",
"idempotency_key",
"user_version",
"execute",
"status",
"wait",
"cancel",
"to_json",
"_to_json",
"serialize",
"geneva",
)
def _sample_function_wire(
*,
function_id: str = _CALL_FUNCTION_ID,
parameters: list[dict[str, str]] | None = None,
output_type_ipc: str = _UTF8_TYPE_IPC_B64,
) -> dict[str, Any]:
return {
"format_version": 1,
"id": function_id,
"signature": {
"parameters": parameters
or [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": output_type_ipc,
"nullable": True,
},
},
}
def _lookup_success_body(function: dict[str, Any] | None = None) -> bytes:
return json.dumps({"function": function or _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _lookup_function(function: dict[str, Any] | None = None):
body = _lookup_success_body(function)
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _CALL_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(body)
with _mock_remote_db(handler) as db:
return db.functions.get(_CALL_CATALOG_NAME)
def _authored_call_type():
cls = getattr(_native, "_FunctionCall", None)
if cls is None:
pytest.fail("lancedb._lancedb._FunctionCall is missing")
return cls
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def test_function_keyword_call_returns_private_frozen_authored_value():
function = _lookup_function()
assert callable(function)
authored = function(text=col("text"), limit=8)
authored_type = _authored_call_type()
assert type(authored) is authored_type
assert authored_type.__module__ == "lancedb._lancedb"
assert authored_type.__name__ == "_FunctionCall"
# Keyword order must not matter; bindings store/render in signature order.
authored_reversed = function(limit=8, text=col("text"))
assert type(authored_reversed) is authored_type
rendered = repr(authored_reversed)
assert rendered.index("text=") < rendered.index("limit=")
assert 'text=field("text")' in rendered
assert "limit=literal(Int32, null=false)" in rendered
def test_function_call_rejects_positional_missing_and_unknown_args():
function = _lookup_function()
with pytest.raises(TypeError, match="keyword"):
function(col("text"), 8)
with pytest.raises((TypeError, ValueError), match="limit"):
function(text=col("text"))
with pytest.raises((TypeError, ValueError), match="text"):
function(limit=8)
with pytest.raises((TypeError, ValueError), match="unknown|extra"):
function(text=col("text"), limit=8, extra=1)
def test_function_call_accepts_direct_case_sensitive_column_and_rejects_complex_exprs():
function = _lookup_function()
authored = function(text=col("firstName"), limit=1)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert 'text=field("firstName")' in rendered
assert "limit=literal(Int32, null=false)" in rendered
complex_exprs = (
col("text") + lit("x"),
col("text").cast(pa.string()),
func("lower", col("text")),
col("text") == lit("x"),
col("text").lower(),
)
for expr in complex_exprs:
with pytest.raises((TypeError, ValueError)):
function(text=expr, limit=1)
# Raw native PyExpr is not the public col() wrapper.
with pytest.raises((TypeError, ValueError)):
function(text=col("text")._inner, limit=1)
# Non-expression / non-literal objects are rejected for field-shaped misuse
# when a column binding is required; plain strings are literals for utf8.
with pytest.raises((TypeError, ValueError)):
function(text=object(), limit=1)
def test_function_call_plain_literal_declared_type_null_and_nested():
function = _lookup_function()
authored = function(text="hello", limit=7)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
# Plain Python int normalizes to declared Int32 and non-null.
authored_int32 = function(text="hello", limit=2_147_483_647)
assert type(authored_int32) is _authored_call_type()
rendered_int32 = repr(authored_int32)
assert "limit=literal(Int32, null=false)" in rendered_int32
assert "Int64" not in rendered_int32
assert "2147483647" not in rendered_int32
# Plain None keeps each declared parameter type with null=true.
authored_null = function(text=None, limit=None)
assert type(authored_null) is _authored_call_type()
rendered_null = repr(authored_null)
assert "text=literal(Utf8, null=true)" in rendered_null
assert "limit=literal(Int32, null=true)" in rendered_null
list_function = _lookup_function(
_sample_function_wire(
parameters=[
{"name": "values", "data_type_ipc": _LIST_INT32_TYPE_IPC_B64},
]
)
)
authored_list = list_function(values=[1, 2, 3])
assert type(authored_list) is _authored_call_type()
rendered_list = repr(authored_list)
assert "values=literal(List(Int32), null=false)" in rendered_list
assert "[1, 2, 3]" not in rendered_list
authored_list_null = list_function(values=None)
assert type(authored_list_null) is _authored_call_type()
rendered_list_null = repr(authored_list_null)
assert "values=literal(List(Int32), null=true)" in rendered_list_null
def test_function_call_direct_literal_expr_exact_type_only():
function = _lookup_function()
# lit(int) is Int64 in the expression builder; int32 parameter must reject it.
with pytest.raises((TypeError, ValueError), match="limit|int32|type") as raised:
function(text="hello", limit=lit(8))
reject_text = _exception_text(raised.value)
assert "Int64" in reject_text or "int64" in reject_text.lower()
assert "Int32" in reject_text or "int32" in reject_text.lower()
# Exact utf8 literal expression is accepted and stored as Utf8/non-null.
authored = function(text=lit("hello"), limit=8)
assert type(authored) is _authored_call_type()
rendered = repr(authored)
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
assert "hello" not in rendered
# Cast / arithmetic around a literal is not a direct Literal node.
with pytest.raises((TypeError, ValueError)):
function(text=lit("hello").cast(pa.string()), limit=8)
def test_function_call_conversion_error_and_repr_are_payload_free():
function = _lookup_function()
with pytest.raises((TypeError, ValueError)) as raised:
function(text="ok", limit=_LITERAL_PAYLOAD_SENTINEL)
text = _exception_text(raised.value)
assert _LITERAL_PAYLOAD_SENTINEL not in text
assert "limit" in text
assert "int32" in text.lower() or "Int32" in text
authored = function(text=_LITERAL_PAYLOAD_SENTINEL, limit=_INT_PAYLOAD_SENTINEL)
rendered = f"{authored!r}\n{authored!s}"
assert _LITERAL_PAYLOAD_SENTINEL not in rendered
assert str(_INT_PAYLOAD_SENTINEL) not in rendered
assert "text=literal(Utf8, null=false)" in rendered
assert "limit=literal(Int32, null=false)" in rendered
assert type(authored).__name__ == "_FunctionCall"
assert "_FunctionCall" in rendered
def test_function_call_private_type_nonconstructible_immutable_and_not_exported():
function = _lookup_function()
authored = function(text=col("text"), limit=1)
authored_type = _authored_call_type()
assert "_FunctionCall" not in getattr(lancedb, "__all__", [])
assert not hasattr(lancedb, "_FunctionCall")
assert getattr(_native, "_FunctionCall", None) is authored_type
with pytest.raises(TypeError):
authored_type()
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(authored, attr)
for attr in ("function", "bindings", "arguments", "parameters", "text", "limit"):
with pytest.raises(AttributeError):
setattr(authored, attr, None)
# Existing Function handle stays frozen / connection-free / name-free.
assert not hasattr(function, "name")
assert not hasattr(function, "connection")
with pytest.raises(AttributeError):
function.id = "mutated"
def test_function_call_does_not_change_col_query_expression_behavior():
# Regression guard: authoring must not alter public col()/Expr query behavior.
expr = col("firstName") > lit(1)
assert isinstance(expr, Expr)
assert expr.to_sql() == "(`firstName` > 1)"
@@ -1,181 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
from __future__ import annotations
import pyarrow as pa
from lancedb import udf
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def double_nullable(value):
if value is None:
return None
return value * 2
def test_first_class_function_enterprise_lifecycle():
import json
import os
import uuid
from datetime import timedelta
import pytest
import lancedb
from lancedb.exceptions import FunctionError
from lancedb.expr import col
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
if not host:
pytest.skip("LANCEDB_FCF_E2E_HOST is required for the live enterprise test")
database_uri = os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local")
api_key = os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake")
run_suffix = uuid.uuid4().hex[:12]
table_name = f"fcf_e2e_{run_suffix}"
function_name = f"fcf_e2e.double_{run_suffix}"
job_timeout = timedelta(minutes=5)
query_timeout = timedelta(seconds=30)
def connect():
return lancedb.connect(
database_uri,
api_key=api_key,
host_override=host,
)
setup_db = connect()
setup_db.create_table(
table_name,
data=pa.Table.from_pylist(
[
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 5},
{"row_id": 3, "value": None},
],
schema=pa.schema(
[
pa.field("row_id", pa.int64(), nullable=False),
pa.field("value", pa.int64(), nullable=True),
]
),
),
)
registration_job = setup_db.functions.register(function_name, double_nullable)
registration_job_id = registration_job.id
assert isinstance(registration_job_id, str) and registration_job_id
registered_function = registration_job.wait(timeout=job_timeout)
assert type(registered_function) is lancedb.Function
assert isinstance(registered_function.id, str) and registered_function.id
with pytest.raises(AttributeError):
registered_function.id = "mutated"
catalog_reader = connect()
function_by_name = catalog_reader.functions.get(function_name)
function_by_id = catalog_reader.functions.get_by_id(registered_function.id)
expected_signature = ((("value", pa.int64()),), pa.int64(), True)
expected_identity = (
registered_function.id,
*expected_signature,
)
for function in (registered_function, function_by_name, function_by_id):
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
generated_column_table = catalog_reader.open_table(table_name)
generated_column_job = generated_column_table.add_generated_column(
"derived",
registered_function(value=col("value")),
)
generated_column_job_id = generated_column_job.id
assert isinstance(generated_column_job_id, str) and generated_column_job_id
assert generated_column_job.wait(timeout=job_timeout) is None
complete_reader = connect().open_table(table_name)
complete_status = complete_reader.generated_column_status("derived")
assert complete_status == "complete"
initial_rows = sorted(
complete_reader.search()
.select(["row_id", "value", "derived"])
.limit(3)
.to_list(timeout=query_timeout),
key=lambda row: row["row_id"],
)
assert initial_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 5, "derived": 10},
{"row_id": 3, "value": None, "derived": None},
]
update_result = complete_reader.update(
where="row_id = 2",
values={"value": 7},
)
assert update_result.rows_updated == 1
incomplete_reader = connect().open_table(table_name)
incomplete_status = incomplete_reader.generated_column_status("derived")
assert incomplete_status == "incomplete"
with pytest.raises(FunctionError) as raised:
(
incomplete_reader.search()
.select(["row_id", "derived"])
.limit(3)
.to_list(timeout=query_timeout)
)
assert raised.value.code == "generated_column_incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
assert refresh_job.wait(timeout=job_timeout) is None
refreshed_reader = connect().open_table(table_name)
refreshed_status = refreshed_reader.generated_column_status("derived")
assert refreshed_status == "complete"
final_rows = sorted(
refreshed_reader.search()
.select(["row_id", "value", "derived"])
.limit(3)
.to_list(timeout=query_timeout),
key=lambda row: row["row_id"],
)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
evidence = {
"run_suffix": run_suffix,
"database": database_uri.removeprefix("db://"),
"table": table_name,
"function": function_name,
"function_id": registered_function.id,
"job_ids": {
"register": registration_job_id,
"add_generated_column": generated_column_job_id,
"refresh_generated_column": refresh_job_id,
},
"status_transitions": [
complete_status,
incomplete_status,
refreshed_status,
],
"final_rows": final_rows,
}
print(json.dumps(evidence, sort_keys=True, separators=(",", ":")))
@@ -1,595 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
from __future__ import annotations
import pyarrow as pa
from lancedb import udf
_RUNNING_DEADLINE_SECONDS = 30
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def reliable_double(value):
if value is None:
return None
return value * 2
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=True,
)
def terminate_worker_on_input(value):
if value is None:
return None
try:
if len(value) == 0:
return value
except TypeError:
pass
import os
os._exit(73)
@udf(
inputs={"value": pa.int64()},
output=pa.int64(),
python="3.12",
packages=["pyarrow==24.0.0"],
output_nullable=False,
)
def slow_triple(value):
import time
time.sleep(0.02)
return value * 3
def _require_live() -> str:
import os
import pytest
host = os.environ.get("LANCEDB_FCF_E2E_HOST")
if not host:
pytest.skip(
"LANCEDB_FCF_E2E_HOST is required for live enterprise reliability tests"
)
return host
def _job_timeout():
from datetime import timedelta
return timedelta(minutes=5)
def _query_timeout():
from datetime import timedelta
return timedelta(seconds=30)
def _connect():
import os
import lancedb
return lancedb.connect(
os.environ.get("LANCEDB_FCF_E2E_DB_URI", "db://fcf-e2e-local"),
api_key=os.environ.get("LANCEDB_FCF_E2E_API_KEY", "fake"),
host_override=_require_live(),
)
def _run_names(case: str) -> tuple[str, str]:
import uuid
suffix = uuid.uuid4().hex[:12]
return f"fcf_rel_{case}_{suffix}", f"fcf_rel.{case}_{suffix}"
def _read_rows(table, columns: list[str], row_count: int) -> list[dict]:
return sorted(
table.search()
.select(columns)
.limit(row_count)
.to_list(timeout=_query_timeout()),
key=lambda row: row["row_id"],
)
def _emit_evidence(case: str, evidence: dict) -> None:
import json
print(
json.dumps(
{"case": case, **evidence},
sort_keys=True,
separators=(",", ":"),
)
)
def test_enterprise_reliability_core_lifecycle():
import pytest
import lancedb
from lancedb.exceptions import FunctionError
from lancedb.expr import col
_require_live()
table_name, function_name = _run_names("lifecycle")
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.Table.from_pylist(
[
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 5},
{"row_id": 3, "value": None},
],
schema=pa.schema(
[
pa.field("row_id", pa.int64(), nullable=False),
pa.field("value", pa.int64(), nullable=True),
]
),
),
)
registration_job = setup_db.functions.register(function_name, reliable_double)
registration_job_id = registration_job.id
assert isinstance(registration_job_id, str) and registration_job_id
registered = registration_job.wait(timeout=_job_timeout())
assert type(registered) is lancedb.Function
assert isinstance(registered.id, str) and registered.id
with pytest.raises(AttributeError):
registered.id = "mutated"
catalog_reader = _connect()
by_name = catalog_reader.functions.get(function_name)
by_id = catalog_reader.functions.get_by_id(registered.id)
expected_identity = (
registered.id,
(("value", pa.int64()),),
pa.int64(),
True,
)
for function in (registered, by_name, by_id):
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
table = catalog_reader.open_table(table_name)
create_job = table.add_generated_column(
"derived",
registered(value=col("value")),
)
create_job_id = create_job.id
assert isinstance(create_job_id, str) and create_job_id
assert create_job.wait(timeout=_job_timeout()) is None
complete_reader = _connect().open_table(table_name)
complete_status = complete_reader.generated_column_status("derived")
assert complete_status == "complete"
initial_rows = _read_rows(
complete_reader,
["row_id", "value", "derived"],
3,
)
assert initial_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 5, "derived": 10},
{"row_id": 3, "value": None, "derived": None},
]
complete_reader.update(where="row_id = 2", values={"value": 7})
incomplete_reader = _connect().open_table(table_name)
changed_rows = _read_rows(incomplete_reader, ["row_id", "value"], 3)
assert changed_rows == [
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 7},
{"row_id": 3, "value": None},
]
incomplete_status = incomplete_reader.generated_column_status("derived")
assert incomplete_status == "incomplete"
with pytest.raises(FunctionError) as raised:
(
incomplete_reader.search()
.select(["row_id", "derived"])
.limit(3)
.to_list(timeout=_query_timeout())
)
assert raised.value.code == "generated_column_incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
assert refresh_job.wait(timeout=_job_timeout()) is None
refreshed_reader = _connect().open_table(table_name)
refreshed_status = refreshed_reader.generated_column_status("derived")
assert refreshed_status == "complete"
final_rows = _read_rows(
refreshed_reader,
["row_id", "value", "derived"],
3,
)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
_emit_evidence(
"core_lifecycle",
{
"final_rows": final_rows,
"function_id": registered.id,
"job_ids": {
"create": create_job_id,
"refresh": refresh_job_id,
"register": registration_job_id,
},
"status": [
complete_status,
incomplete_status,
refreshed_status,
],
"table": table_name,
},
)
def test_enterprise_reliability_restart_retention():
import json
import os
import pytest
import lancedb
_require_live()
raw_evidence = os.environ.get("LANCEDB_FCF_E2E_RESTART_EVIDENCE")
if not raw_evidence:
pytest.skip(
"LANCEDB_FCF_E2E_RESTART_EVIDENCE is required for restart retention"
)
try:
evidence = json.loads(raw_evidence)
except json.JSONDecodeError as error:
pytest.fail(f"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be valid JSON: {error.msg}")
assert isinstance(evidence, dict), (
"LANCEDB_FCF_E2E_RESTART_EVIDENCE must be a JSON object"
)
table_name = evidence.get("table")
function_id = evidence.get("function_id")
raw_job_ids = evidence.get("job_ids")
assert isinstance(table_name, str) and table_name, (
"restart evidence must contain a non-empty table"
)
assert isinstance(function_id, str) and function_id, (
"restart evidence must contain a non-empty function_id"
)
assert isinstance(raw_job_ids, dict), (
"restart evidence must contain a job_ids object"
)
job_ids = {}
for job_kind in ("register", "create", "refresh"):
job_id = raw_job_ids.get(job_kind)
assert isinstance(job_id, str) and job_id, (
f"restart evidence must contain a non-empty job_ids.{job_kind}"
)
job_ids[job_kind] = job_id
db = _connect()
function = db.functions.get_by_id(function_id)
expected_identity = (
function_id,
(("value", pa.int64()),),
pa.int64(),
True,
)
assert type(function) is lancedb.Function
assert (
function.id,
function.parameters,
function.output_type,
function.output_nullable,
) == expected_identity
jobs = {}
for job_kind in ("register", "create", "refresh"):
job = db.get_job(job_ids[job_kind])
assert job is not None
assert job.job_id == job_ids[job_kind]
assert job.state == "finished"
assert job.failure is None
jobs[job_kind] = job
registered_result = jobs["register"].result
assert type(registered_result) is lancedb.Function
assert (
registered_result.id,
registered_result.parameters,
registered_result.output_type,
registered_result.output_nullable,
) == expected_identity
assert jobs["create"].result is None
assert jobs["refresh"].result is None
table = db.open_table(table_name)
status = table.generated_column_status("derived")
assert status == "complete"
assert table.count_rows() == 3
final_rows = _read_rows(table, ["row_id", "value", "derived"], 3)
assert final_rows == [
{"row_id": 1, "value": 2, "derived": 4},
{"row_id": 2, "value": 7, "derived": 14},
{"row_id": 3, "value": None, "derived": None},
]
_emit_evidence(
"restart_retention",
{
"final_rows": final_rows,
"function_id": function_id,
"generated_column_status": status,
"job_ids": job_ids,
"job_states": {
job_kind: jobs[job_kind].state
for job_kind in ("register", "create", "refresh")
},
"table": table_name,
},
)
def test_enterprise_reliability_failure_atomicity_and_worker_recovery():
import pytest
import lancedb
from lancedb.exceptions import JobFailedError
from lancedb.expr import col
_require_live()
table_name, failing_function_name = _run_names("worker_failure")
_, healthy_function_name = _run_names("worker_recovery")
row_count = 4
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.table(
{
"row_id": list(range(row_count)),
"value": [1, 2, 3, 4],
}
),
)
registration_job = setup_db.functions.register(
failing_function_name,
terminate_worker_on_input,
)
failing_function = registration_job.wait(timeout=_job_timeout())
assert type(failing_function) is lancedb.Function
table = setup_db.open_table(table_name)
failed_create_job = table.add_generated_column(
"must_not_publish",
failing_function(value=col("value")),
)
failed_job_id = failed_create_job.id
assert isinstance(failed_job_id, str) and failed_job_id
with pytest.raises(JobFailedError) as raised:
failed_create_job.wait(timeout=_job_timeout())
assert raised.value.error_code == "udf_execution_failure"
first_description = _connect().get_job(failed_job_id)
second_description = _connect().get_job(failed_job_id)
for description in (first_description, second_description):
assert description is not None
assert description.job_id == failed_job_id
assert description.state == "failed"
assert description.failure is not None
assert description.failure.error_code == "udf_execution_failure"
atomic_reader = _connect().open_table(table_name)
assert "must_not_publish" not in atomic_reader.schema.names
assert _read_rows(atomic_reader, ["row_id", "value"], row_count) == [
{"row_id": 0, "value": 1},
{"row_id": 1, "value": 2},
{"row_id": 2, "value": 3},
{"row_id": 3, "value": 4},
]
healthy_registration_job = setup_db.functions.register(
healthy_function_name,
reliable_double,
)
healthy_function = healthy_registration_job.wait(timeout=_job_timeout())
assert type(healthy_function) is lancedb.Function
recovery_job = atomic_reader.add_generated_column(
"recovered",
healthy_function(value=col("value")),
)
recovery_job_id = recovery_job.id
assert isinstance(recovery_job_id, str) and recovery_job_id
assert recovery_job.wait(timeout=_job_timeout()) is None
recovered_reader = _connect().open_table(table_name)
assert "must_not_publish" not in recovered_reader.schema.names
assert recovered_reader.generated_column_status("recovered") == "complete"
recovered_rows = _read_rows(
recovered_reader,
["row_id", "value", "recovered"],
row_count,
)
assert recovered_rows == [
{"row_id": 0, "value": 1, "recovered": 2},
{"row_id": 1, "value": 2, "recovered": 4},
{"row_id": 2, "value": 3, "recovered": 6},
{"row_id": 3, "value": 4, "recovered": 8},
]
_emit_evidence(
"failure_atomicity_and_worker_recovery",
{
"failure_code": first_description.failure.error_code,
"failed_job_id": failed_job_id,
"recovered_rows": recovered_rows,
"recovery_job_id": recovery_job_id,
"table": table_name,
},
)
def test_enterprise_reliability_concurrent_refresh_fencing():
import time
import pytest
import lancedb
from lancedb.exceptions import FunctionError, JobFailedError
from lancedb.expr import col
_require_live()
table_name, function_name = _run_names("refresh_fencing")
row_count = 1024
setup_db = _connect()
setup_db.create_table(
table_name,
data=pa.table(
{
"row_id": list(range(row_count)),
"value": list(range(row_count)),
}
),
)
registration_job = setup_db.functions.register(function_name, slow_triple)
function = registration_job.wait(timeout=_job_timeout())
assert type(function) is lancedb.Function
table = setup_db.open_table(table_name)
create_job = table.add_generated_column(
"derived",
function(value=col("value")),
)
assert create_job.wait(timeout=_job_timeout()) is None
initial_reader = _connect().open_table(table_name)
assert initial_reader.generated_column_status("derived") == "complete"
initial_reader.update(where="row_id = 0", values={"value": 10_000})
incomplete_reader = _connect().open_table(table_name)
assert incomplete_reader.generated_column_status("derived") == "incomplete"
refresh_job = incomplete_reader.refresh_generated_column("derived")
refresh_job_id = refresh_job.id
assert isinstance(refresh_job_id, str) and refresh_job_id
deadline = time.monotonic() + _RUNNING_DEADLINE_SECONDS
observed_states = []
running_observations = 0
while running_observations < 2:
state = refresh_job.status()
if not observed_states or observed_states[-1] != state:
observed_states.append(state)
if state == "running":
running_observations += 1
else:
running_observations = 0
assert state not in {"finished", "failed", "cancelled"}
assert time.monotonic() < deadline
if running_observations < 2:
time.sleep(0.05)
concurrent_writer = _connect().open_table(table_name)
concurrent_writer.update(where="row_id = 1", values={"value": 20_000})
with pytest.raises(JobFailedError) as raised:
refresh_job.wait(timeout=_job_timeout())
assert raised.value.error_code == "stale_or_conflicting_input"
stale_job = _connect().get_job(refresh_job_id)
assert stale_job is not None
assert stale_job.job_id == refresh_job_id
assert stale_job.state == "failed"
assert stale_job.failure is not None
assert stale_job.failure.error_code == raised.value.error_code
if observed_states[-1] != stale_job.state:
observed_states.append(stale_job.state)
stale_reader = _connect().open_table(table_name)
stale_rows = _read_rows(stale_reader, ["row_id", "value"], row_count)
assert len(stale_rows) == row_count
for row_id, row in enumerate(stale_rows):
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
assert (row["row_id"], row["value"]) == (row_id, expected_value)
assert stale_reader.generated_column_status("derived") == "incomplete"
with pytest.raises(FunctionError) as incomplete:
(
stale_reader.search()
.select(["row_id", "derived"])
.limit(row_count)
.to_list(timeout=_query_timeout())
)
assert incomplete.value.code == "generated_column_incomplete"
resubmitted_job = stale_reader.refresh_generated_column("derived")
resubmitted_job_id = resubmitted_job.id
assert isinstance(resubmitted_job_id, str) and resubmitted_job_id
assert resubmitted_job.wait(timeout=_job_timeout()) is None
final_reader = _connect().open_table(table_name)
final_status = final_reader.generated_column_status("derived")
assert final_status == "complete"
final_rows = _read_rows(
final_reader,
["row_id", "value", "derived"],
row_count,
)
assert len(final_rows) == row_count
final_checksum = 0
for row_id, row in enumerate(final_rows):
expected_value = 10_000 if row_id == 0 else 20_000 if row_id == 1 else row_id
assert (row["row_id"], row["value"], row["derived"]) == (
row_id,
expected_value,
expected_value * 3,
)
final_checksum += row["derived"]
_emit_evidence(
"concurrent_refresh_fencing",
{
"failure_code": stale_job.failure.error_code,
"final_checksum": final_checksum,
"final_status": final_status,
"observed_states": observed_states,
"resubmitted_job_id": resubmitted_job_id,
"row_count": row_count,
"stale_job_id": refresh_job_id,
"table": table_name,
},
)
@@ -1,268 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract: Python projection of JobFailure.error_code / JobFailedError.error_code.
Public Function failures expose eight stable string categories. Asynchronous
errors remain the unified JobFailedError and JobFailureInfo. Python must
project the optional exact error_code string already supplied structurally by
Rust: preserve a known code, preserve an unknown nonempty future code
byte-for-byte, and return None for legacy failure payloads without error_code.
Never infer or override a code from message, phase, retryable, HTTP status,
job type, or state.
"""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable, Optional
import pytest
import lancedb
from lancedb.exceptions import JobFailedError
_DESCRIBE_PATH = "/v1/jobs/describe"
_KNOWN_CODE = "name_or_function_not_found"
_CONFLICTING_STABLE_IN_MESSAGE = "definition_validation_failure"
_UNKNOWN_CODE = "enterprise_future_category_xyz"
_WAIT_KNOWN_CODE = "unsupported_runtime_or_capability"
_WAIT_CONFLICTING_IN_MESSAGE = "revoked_function"
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _failed_describe_body(
*,
job_id: str,
error_code: Optional[str] = None,
include_error_code: bool = True,
phase: str = "execute",
message: str = "worker died",
retryable: bool = False,
job_type: str = "create_index",
) -> dict[str, Any]:
failure: dict[str, Any] = {
"phase": phase,
"message": message,
"retryable": retryable,
}
if include_error_code:
failure["error_code"] = error_code
return {
"job_id": job_id,
"job_type": job_type,
"job_state": "FAILED",
"creation_ms": 1000,
"spec": {},
"failure": failure,
}
def _describe_handler(bodies_by_job_id: dict[str, dict[str, Any]]):
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _DESCRIBE_PATH
payload = json.loads(_read_body(request).decode("utf-8") or "{}")
job_id = payload["job_id"]
body = bodies_by_job_id.get(job_id)
if body is None:
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
return handler
def test_get_job_failure_error_code_known_not_inferred_from_message():
"""Structural error_code wins; conflicting message text must not override."""
body = _failed_describe_body(
job_id="job-known",
error_code=_KNOWN_CODE,
phase="validate",
message=f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}",
retryable=False,
)
with _mock_remote_db(_describe_handler({"job-known": body})) as db:
description = db.get_job("job-known")
assert description is not None
failure = description.failure
assert failure is not None
assert failure.error_code == _KNOWN_CODE
assert failure.error_code != _CONFLICTING_STABLE_IN_MESSAGE
assert failure.phase == "validate"
assert failure.message == f"looks like {_CONFLICTING_STABLE_IN_MESSAGE}"
assert failure.retryable is False
def test_get_job_failure_error_code_unknown_preserved_byte_for_byte():
body = _failed_describe_body(
job_id="job-unknown",
error_code=_UNKNOWN_CODE,
phase="execute",
message=f"new category mentioning {_KNOWN_CODE}",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-unknown": body})) as db:
failure = db.get_job("job-unknown").failure
assert failure.error_code == _UNKNOWN_CODE
assert failure.error_code != _KNOWN_CODE
def test_get_job_failure_error_code_absent_is_none():
"""Legacy describe payloads without error_code must not invent a category."""
body = _failed_describe_body(
job_id="job-legacy",
include_error_code=False,
phase="execute",
message=f"{_KNOWN_CODE} in logs",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-legacy": body})) as db:
failure = db.get_job("job-legacy").failure
assert failure.error_code is None
assert failure.phase == "execute"
assert failure.retryable is True
def test_sync_job_wait_job_failed_error_code_known_not_inferred():
body = _failed_describe_body(
job_id="job-wait-known",
error_code=_WAIT_KNOWN_CODE,
phase="dispatch",
message=f"{_WAIT_CONFLICTING_IN_MESSAGE} in transport logs",
retryable=False,
)
with _mock_remote_db(_describe_handler({"job-wait-known": body})) as db:
with pytest.raises(JobFailedError) as exc_info:
db.job("job-wait-known").wait(timeout=timedelta(seconds=5))
err = exc_info.value
assert isinstance(err, JobFailedError)
assert err.error_code == _WAIT_KNOWN_CODE
assert err.error_code != _WAIT_CONFLICTING_IN_MESSAGE
def test_sync_job_wait_job_failed_error_code_absent_is_none():
body = _failed_describe_body(
job_id="job-wait-legacy",
include_error_code=False,
phase="execute",
message=f"{_WAIT_KNOWN_CODE} mentioned only in message",
retryable=True,
)
with _mock_remote_db(_describe_handler({"job-wait-legacy": body})) as db:
with pytest.raises(JobFailedError) as exc_info:
db.job("job-wait-legacy").wait(timeout=timedelta(seconds=5))
assert exc_info.value.error_code is None
@pytest.mark.asyncio
async def test_async_job_wait_job_failed_error_code_unknown_preserved():
body = _failed_describe_body(
job_id="job-wait-unknown",
error_code=_UNKNOWN_CODE,
phase="execute",
message=f"future code with {_WAIT_KNOWN_CODE} in text",
retryable=False,
)
async with _mock_remote_db_async(
_describe_handler({"job-wait-unknown": body})
) as db:
with pytest.raises(JobFailedError) as exc_info:
await db.job("job-wait-unknown").wait(timeout=timedelta(seconds=5))
err = exc_info.value
assert err.error_code == _UNKNOWN_CODE
assert err.error_code != _WAIT_KNOWN_CODE
def test_job_failed_error_legacy_message_construction_error_code_is_none():
err = JobFailedError("legacy construction with only a message")
assert err.error_code is None
def test_job_failed_error_error_code_is_read_only():
err = JobFailedError("message")
with pytest.raises(AttributeError):
err.error_code = _KNOWN_CODE
@@ -1,634 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python first-class Function catalog lookup."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
_LOOKUP_PATH = "/v1/functions/lookup"
_LOOKUP_CATALOG_NAME = "text.normalize.lookup-name"
_LOOKUP_FUNCTION_ID = "fn.exact.lookup-handle"
_LOOKUP_SERVER_MESSAGE_MARKER = (
"SERVER_LOOKUP_DIAGNOSTIC_MARKER name=text.normalize.lookup-name "
"id=fn.exact.lookup-handle"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_LOOKUP_BODY_MARKER"
_UNKNOWN_CODE = "enterprise_future_lookup_category_xyz"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as job-result
# tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust FileWriter.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_LOOKUP_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"replace",
"expected_current_function_id",
"list",
"alias",
"lineage",
"FunctionVersion",
)
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _LOOKUP_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body(
*,
function: dict[str, Any] | None = None,
extra_outer: dict[str, Any] | None = None,
) -> bytes:
body: dict[str, Any] = {"function": function or _sample_function_wire()}
if extra_outer:
body.update(extra_outer)
return json.dumps(body).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _LOOKUP_SERVER_MESSAGE_MARKER not in text
assert _LOOKUP_CATALOG_NAME not in text
assert _LOOKUP_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_lookup_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _LOOKUP_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _LOOKUP_CATALOG_NAME not in repr(function)
assert _LOOKUP_CATALOG_NAME not in str(function)
def _assert_name_request(raw: bytes, body: dict[str, Any]) -> None:
assert raw
assert body == {"name": _LOOKUP_CATALOG_NAME}
assert "function_id" not in body
def _assert_id_request(raw: bytes, body: dict[str, Any]) -> None:
assert raw
assert body == {"function_id": _LOOKUP_FUNCTION_ID}
assert "name" not in body
def _assert_native_lookup_methods_present() -> None:
assert hasattr(_native.Connection, "_lookup_function_by_name")
assert hasattr(_native.Connection, "_lookup_function_by_id")
assert callable(getattr(_native.Connection, "_lookup_function_by_name"))
assert callable(getattr(_native.Connection, "_lookup_function_by_id"))
def test_native_connection_exposes_private_lookup_methods():
_assert_native_lookup_methods_present()
def test_sync_remote_get_by_name_exact_request_and_function_shape():
_assert_native_lookup_methods_present()
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
with _mock_remote_db(handler) as db:
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
assert not hasattr(db, "get_function")
function = db.functions.get(_LOOKUP_CATALOG_NAME)
_assert_name_request(seen["raw"], seen["body"])
_assert_exact_lookup_function(function)
def test_sync_remote_get_by_id_exact_request_and_function_shape():
_assert_native_lookup_methods_present()
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
with _mock_remote_db(handler) as db:
function = db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
_assert_id_request(seen["raw"], seen["body"])
_assert_exact_lookup_function(function)
@pytest.mark.asyncio
async def test_async_remote_get_by_name_and_id():
_assert_native_lookup_methods_present()
name_seen: dict[str, Any] = {}
id_seen: dict[str, Any] = {}
stage = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
body = json.loads(raw.decode("utf-8"))
stage["n"] += 1
if stage["n"] == 1:
name_seen["raw"] = raw
name_seen["body"] = body
else:
id_seen["raw"] = raw
id_seen["body"] = body
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
async with _mock_remote_db_async(handler) as db:
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
by_name = await db.functions.get(_LOOKUP_CATALOG_NAME)
by_id = await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
_assert_name_request(name_seen["raw"], name_seen["body"])
_assert_id_request(id_seen["raw"], id_seen["body"])
_assert_exact_lookup_function(by_name)
_assert_exact_lookup_function(by_id)
def test_sync_remote_get_accepts_additive_outer_success_fields():
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
_lookup_success_body(
extra_outer={
"server_extra": {"ok": True},
"request_echo_name": _LOOKUP_CATALOG_NAME,
}
)
)
with _mock_remote_db(handler) as db:
function = db.functions.get(_LOOKUP_CATALOG_NAME)
_assert_exact_lookup_function(function)
def test_empty_name_and_id_reject_before_transport():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(ValueError):
db.functions.get("")
with pytest.raises(ValueError):
db.functions.get_by_id("")
assert received["n"] == 0
def test_local_sync_lookup_not_implemented_without_table_mutation(tmp_path):
_assert_native_lookup_methods_present()
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
with pytest.raises(NotImplementedError):
db.functions.get(_LOOKUP_CATALOG_NAME)
with pytest.raises(NotImplementedError):
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_lookup_not_implemented_without_table_mutation(tmp_path):
_assert_native_lookup_methods_present()
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
with pytest.raises(NotImplementedError):
await db.functions.get(_LOOKUP_CATALOG_NAME)
with pytest.raises(NotImplementedError):
await db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
assert (await db.list_tables()).tables == before
def test_explicit_known_code_is_function_error_with_exact_code():
body = {
"error_code": "name_or_function_not_found",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
"looks_like": "definition_validation_failure",
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_or_function_not_found"
assert err.code != "definition_validation_failure"
_assert_payload_free(err)
def test_explicit_unknown_code_preserved_despite_status_and_message():
body = {
"error_code": _UNKNOWN_CODE,
"message": (
f"{_LOOKUP_SERVER_MESSAGE_MARKER} revoked_function "
"name_or_function_not_found"
),
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
raw = _read_body(request)
assert json.loads(raw.decode("utf-8")) == {"function_id": _LOOKUP_FUNCTION_ID}
request.send_response(409)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get_by_id(_LOOKUP_FUNCTION_ID)
err = exc_info.value
assert err.code == _UNKNOWN_CODE
assert err.code != "revoked_function"
assert err.code != "name_or_function_not_found"
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"missing_code_404",
404,
{
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"empty_code",
400,
{
"error_code": "",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"wrong_type_code",
400,
{
"error_code": 123,
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"null_code",
404,
{
"error_code": None,
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
},
),
(
"non_json",
404,
f"not-json {_LOOKUP_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
],
)
def test_invalid_or_missing_error_code_is_payload_free_http(
label: str, status: int, response_body: object
):
del label # parametrize label for failure diagnosis only
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,response_body",
[
(
"missing_function",
{
"server_extra": True,
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"null_function",
{
"function": None,
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"wrong_type_function",
{
"function": "not-an-object",
_SENSITIVE_BODY_MARKER: _LOOKUP_SERVER_MESSAGE_MARKER,
},
),
(
"invalid_function_shape",
{
"function": {
"format_version": 1,
"id": _LOOKUP_FUNCTION_ID,
# missing signature
_SENSITIVE_BODY_MARKER: True,
}
},
),
],
)
def test_malformed_success_is_payload_free_http(label: str, response_body: dict):
del label
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
def test_function_error_surface_omits_server_marker_name_and_id():
body = {
"error_code": "name_or_function_not_found",
"message": _LOOKUP_SERVER_MESSAGE_MARKER,
"function_id": _LOOKUP_FUNCTION_ID,
"name": _LOOKUP_CATALOG_NAME,
_SENSITIVE_BODY_MARKER: True,
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
_read_body(request)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
with pytest.raises(function_error) as exc_info:
db.functions.get(_LOOKUP_CATALOG_NAME)
err = exc_info.value
_assert_payload_free(err)
assert getattr(err, "code", None) == "name_or_function_not_found"
def test_no_direct_db_lookup_methods_and_no_deleted_keywords():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
assert not hasattr(db, "lookup_function")
assert not hasattr(db, "lookup_function_by_name")
assert not hasattr(db, "lookup_function_by_id")
assert not hasattr(db, "get_function")
assert not hasattr(db.functions, "get_by_name")
assert not hasattr(db.functions, "list")
for keyword in _DELETED_LOOKUP_KEYWORDS:
with pytest.raises(TypeError):
db.functions.get(_LOOKUP_CATALOG_NAME, **{keyword: True})
with pytest.raises(TypeError):
db.functions.get_by_id(_LOOKUP_FUNCTION_ID, **{keyword: True})
assert received["n"] == 0
def test_function_error_is_not_top_level_export():
assert not hasattr(lancedb, "FunctionError")
function_error = _function_error_cls()
assert issubclass(function_error, RuntimeError)
@@ -1,398 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for Python first-class Function registration."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable
from unittest import mock
import pyarrow as pa
import pytest
import lancedb
import lancedb._udf as _udf_mod
import lancedb.job
from lancedb import FunctionCapability, udf
from lancedb.remote.errors import HttpError
_SOURCE_MARKER = "registration-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/registration-redact-token-xyz"
_SECRET_ENV = "REGISTER_API_TOKEN"
_NETWORK_ORIGIN = "https://api.registration-example.com"
_FUNCTION_NAME = "text.normalize"
_FUNCTION_ID_RETRY = "fn.register-retry-1"
_JOB_ID_RETRY = "job-register-retry-1"
_JOB_ID_ASYNC = "job-register-async-1"
_REGISTER_PATH = "/v1/functions/register"
_DESCRIBE_PATH = "/v1/jobs/describe"
_DELETED_REGISTER_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"deterministic",
"null_policy",
"replace",
"expected_current_function_id",
)
_SPEC_KEYS = {
"format_version",
"name",
"definition",
"expected_current_function_id",
}
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_register_normalize(text, limit):
"""registration-source-marker-unique-xyz."""
return text[:limit]
def _definition_json(fn: object) -> dict[str, Any]:
payload = _udf_mod._build_function_definition(fn)._to_json()
if isinstance(payload, bytes):
return json.loads(payload.decode("utf-8"))
assert isinstance(payload, str)
return json.loads(payload)
def _expected_register_spec(name: str, fn: object) -> dict[str, Any]:
return {
"format_version": 1,
"name": name,
"definition": _definition_json(fn),
"expected_current_function_id": None,
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_markers_absent_from_exception(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _SOURCE_MARKER not in text
assert _SECRET_REFERENCE not in text
def _assert_exact_register_spec(body: dict[str, Any], expected: dict[str, Any]) -> None:
assert set(body) == _SPEC_KEYS
assert body == expected
assert body["format_version"] == 1
assert body["expected_current_function_id"] is None
assert _SOURCE_MARKER in json.dumps(body["definition"])
assert any(
capability.get("reference") == _SECRET_REFERENCE
for capability in body["definition"]["capabilities"]
)
def test_sync_remote_register_retries_exact_wire_and_returns_job():
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
attempts: list[dict[str, Any]] = []
describe_calls: list[dict[str, Any]] = []
function_result_wire = {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _FUNCTION_ID_RETRY,
"signature": expected_spec["definition"]["signature"],
},
}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _REGISTER_PATH:
request_id = request.headers.get("x-request-id")
attempts.append(
{
"request_id": request_id,
"raw": raw,
"body": json.loads(raw.decode("utf-8")),
}
)
if len(attempts) == 1:
request.send_response(500)
request.end_headers()
request.wfile.write(b"transient register failure")
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_RETRY}).encode("utf-8"))
return
assert request.path == _DESCRIBE_PATH
body = json.loads(raw.decode("utf-8"))
assert body["job_id"] == _JOB_ID_RETRY
describe_calls.append(body)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_RETRY,
"job_state": "DONE",
"job_type": "register_function",
"creation_ms": 1,
"spec": {},
"result": function_result_wire,
}
).encode("utf-8")
)
package_calls = {"n": 0}
original_package = _udf_mod._package_udf
def counting_package(fn: object):
package_calls["n"] += 1
return original_package(fn)
with _mock_remote_db(handler) as db:
assert not hasattr(db, "register_function")
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
job = db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_RETRY
waited = job.wait(timeout=timedelta(seconds=5))
assert package_calls["n"] == 1
assert len(attempts) == 2
first, second = attempts
assert isinstance(first["request_id"], str) and first["request_id"]
assert first["request_id"] == second["request_id"]
assert first["raw"] == second["raw"]
assert first["raw"]
_assert_exact_register_spec(first["body"], expected_spec)
_assert_exact_register_spec(second["body"], expected_spec)
assert len(describe_calls) == 1
assert describe_calls[0]["job_id"] == _JOB_ID_RETRY
assert type(waited) is lancedb.Function
assert waited.id == _FUNCTION_ID_RETRY
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
assert waited.output_type == pa.string()
assert waited.output_nullable is True
@pytest.mark.asyncio
async def test_async_remote_register_returns_async_job_with_exact_spec():
expected_spec = _expected_register_spec(_FUNCTION_NAME, packable_register_normalize)
seen: dict[str, Any] = {}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
assert request.path == _REGISTER_PATH
raw = _read_body(request)
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
async with _mock_remote_db_async(handler) as db:
assert not hasattr(db, "register_function")
job = await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert seen.get("raw")
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
_assert_exact_register_spec(seen["body"], expected_spec)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
def test_sync_remote_register_http_error_omits_source_and_secret_markers():
echoed = f"register failed with {_SOURCE_MARKER} and {_SECRET_REFERENCE}"
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
assert request.path == _REGISTER_PATH
_read_body(request)
request.send_response(400)
request.end_headers()
request.wfile.write(echoed.encode("utf-8"))
with _mock_remote_db(handler) as db:
with pytest.raises(HttpError) as exc_info:
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert received["n"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert err.status_code == 400
_assert_markers_absent_from_exception(err)
def test_empty_name_rejects_before_http():
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(ValueError):
db.functions.register("", packable_register_normalize)
assert received["n"] == 0
def test_local_sync_register_not_implemented_without_table_mutation(tmp_path):
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "register_function")
with pytest.raises(NotImplementedError):
db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_register_not_implemented_without_table_mutation(tmp_path):
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "register_function")
with pytest.raises(NotImplementedError):
await db.functions.register(_FUNCTION_NAME, packable_register_normalize)
assert (await db.list_tables()).tables == before
@pytest.mark.parametrize("keyword", _DELETED_REGISTER_KEYWORDS)
def test_register_rejects_deleted_overdesign_keywords_before_submission(keyword):
received = {"n": 0}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
received["n"] += 1
_read_body(request)
request.send_response(500)
request.end_headers()
request.wfile.write(b"should not be reached")
with _mock_remote_db(handler) as db:
with pytest.raises(TypeError):
db.functions.register(
_FUNCTION_NAME,
packable_register_normalize,
**{keyword: True},
)
assert received["n"] == 0
@@ -1,719 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python conditional first-class Function name removal."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
_REMOVE_PATH = "/v1/functions/remove"
_LOOKUP_PATH = "/v1/functions/lookup"
_REMOVE_CATALOG_NAME = "text.normalize.remove-name"
_REMOVE_FUNCTION_ID = "fn.exact.remove-handle"
_REMOVE_SERVER_MESSAGE_MARKER = (
"SERVER_REMOVE_DIAGNOSTIC_MARKER name=text.normalize.remove-name "
"id=fn.exact.remove-handle"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_REMOVE_BODY_MARKER"
_CONFLICTING_MESSAGE_CODE = "revoked_function"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# replace tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REMOVE_KEYWORDS = (
"expected_current_function_id",
"function_id",
"idempotency_key",
"retry_key",
"user_version",
"version",
"force",
"if_exists",
"revoke",
"delete",
)
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _REMOVE_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body() -> bytes:
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
def _close_db(db: Any) -> None:
with contextlib.suppress(Exception):
inner = getattr(db, "_conn", None)
if inner is not None:
inner.close()
return
close = getattr(db, "close", None)
if callable(close):
close()
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _REMOVE_SERVER_MESSAGE_MARKER not in text
assert _REMOVE_CATALOG_NAME not in text
assert _REMOVE_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_remove_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _REMOVE_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _REMOVE_CATALOG_NAME not in repr(function)
assert _REMOVE_CATALOG_NAME not in str(function)
def _assert_exact_remove_request(
request: http.server.BaseHTTPRequestHandler,
raw: bytes,
body: dict[str, Any],
*,
expected_id: str,
) -> None:
assert request.command == "POST"
assert request.path == _REMOVE_PATH
assert "?" not in request.path
assert raw
assert body == {
"name": _REMOVE_CATALOG_NAME,
"expected_current_function_id": expected_id,
}
assert set(body) == {"name", "expected_current_function_id"}
assert "format_version" not in body
assert "function_id" not in body
assert "function" not in body
assert "signature" not in body
assert "job_id" not in body
assert "idempotency_key" not in body
assert "user_version" not in body
assert "force" not in body
assert "if_exists" not in body
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
def _assert_native_remove_method_present() -> None:
assert hasattr(_native.Connection, "_remove_function_name")
assert callable(getattr(_native.Connection, "_remove_function_name"))
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact name lookup; optionally continue for remove."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _REMOVE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["remove"] = counters.get("remove", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected remove")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_REMOVE_CATALOG_NAME)
_assert_exact_remove_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_REMOVE_CATALOG_NAME)
_assert_exact_remove_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
def test_native_connection_exposes_private_remove_function_name():
_assert_native_remove_method_present()
def test_sync_remote_remove_exact_body_path_request_id_returns_none():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
remove_attempts: list[dict[str, Any]] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
body = json.loads(payload.decode("utf-8"))
remove_attempts.append(
{
"request": request,
"raw": payload,
"body": body,
"request_id": request.headers.get("x-request-id"),
}
)
# Illegal body on 204 must be ignored; success is status-driven only.
request.send_response(204)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
_SENSITIVE_BODY_MARKER: True,
"message": _REMOVE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup"] == 1
assert counters["remove"] == 1
assert len(remove_attempts) == 1
attempt = remove_attempts[0]
_assert_exact_remove_request(
attempt["request"],
attempt["raw"],
attempt["body"],
expected_id=current.id,
)
assert attempt["body"]["expected_current_function_id"] == current.id
@pytest.mark.asyncio
async def test_async_remote_remove_exact_body_returns_none():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
seen["request"] = request
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(204)
request.end_headers()
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
result = await db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup"] == 1
assert counters["remove"] == 1
assert seen.get("raw")
_assert_exact_remove_request(
seen["request"],
seen["raw"],
seen["body"],
expected_id=current.id,
)
def test_after_remove_name_lookup_not_found_id_lookup_same_function():
"""Catalog-pointer SDK sequence via a stateful fixture; not server atomicity."""
_assert_native_remove_method_present()
counters: dict[str, int] = {
"lookup_name": 0,
"lookup_id": 0,
"remove": 0,
}
removed = {"yes": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
body = json.loads(raw.decode("utf-8"))
if "name" in body:
counters["lookup_name"] += 1
assert body == {"name": _REMOVE_CATALOG_NAME}
if removed["yes"]:
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"error_code": "name_or_function_not_found",
"message": _REMOVE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
counters["lookup_id"] += 1
assert body == {"function_id": _REMOVE_FUNCTION_ID}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
assert request.path == _REMOVE_PATH
counters["remove"] += 1
body = json.loads(raw.decode("utf-8"))
_assert_exact_remove_request(
request, raw, body, expected_id=_REMOVE_FUNCTION_ID
)
removed["yes"] = True
request.send_response(204)
request.end_headers()
function_error = _function_error_cls()
with _mock_remote_db(handler) as db:
current = _observe_current(db)
assert counters["lookup_name"] == 1
assert counters["lookup_id"] == 0
assert counters["remove"] == 0
result = db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert result is None
assert counters["lookup_name"] == 1
assert counters["remove"] == 1
with pytest.raises(function_error) as exc_info:
db.functions.get(_REMOVE_CATALOG_NAME)
err = exc_info.value
assert err.code == "name_or_function_not_found"
_assert_payload_free(err)
by_id = db.functions.get_by_id(_REMOVE_FUNCTION_ID)
assert counters["lookup_name"] == 2
assert counters["lookup_id"] == 1
assert counters["remove"] == 1
_assert_exact_remove_function(by_id)
assert by_id.id == current.id
assert by_id.parameters == current.parameters
assert by_id.output_type == current.output_type
assert by_id.output_nullable is current.output_nullable
def test_explicit_name_conflict_is_function_error_payload_free():
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
body = {
"error_code": "name_conflict",
"message": (
f"{_REMOVE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE}"
),
"name": _REMOVE_CATALOG_NAME,
"function_id": _REMOVE_FUNCTION_ID,
_SENSITIVE_BODY_MARKER: True,
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_remove_request(
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
)
request.send_response(409)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(function_error) as exc_info:
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert counters["lookup"] == 1
assert counters["remove"] == 1
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_conflict"
assert err.code != _CONFLICTING_MESSAGE_CODE
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"200_with_body",
200,
{
"ok": True,
"message": _REMOVE_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
"job_id": "must-not-infer-job",
},
),
(
"202_empty",
202,
f"{_REMOVE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
("200_empty", 200, ""),
],
)
def test_http_200_202_cannot_return_success(
label: str, status: int, response_body: object
):
del label
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REMOVE_PATH
counters["remove"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_remove_request(
request, payload, parsed, expected_id=_REMOVE_FUNCTION_ID
)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(HttpError) as exc_info:
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert counters["lookup"] == 1
assert counters["remove"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
def test_empty_name_rejects_before_remove_transport():
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(ValueError):
db.functions.remove("", current)
assert counters["lookup"] == 1
assert counters["remove"] == 0
@pytest.mark.parametrize(
"bad_current",
[
_REMOVE_FUNCTION_ID,
{"id": _REMOVE_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_current_rejected_without_remove(bad_current):
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-current path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(TypeError):
db.functions.remove(_REMOVE_CATALOG_NAME, bad_current)
assert counters["lookup"] == 1
assert counters["remove"] == 0
def test_local_sync_remove_not_implemented_without_table_mutation(tmp_path):
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
with pytest.raises(NotImplementedError):
db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert db.list_tables().tables == before
_close_db(db)
@pytest.mark.asyncio
async def test_local_async_remove_not_implemented_without_table_mutation(tmp_path):
_assert_native_remove_method_present()
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
with pytest.raises(NotImplementedError):
await db.functions.remove(_REMOVE_CATALOG_NAME, current)
assert (await db.list_tables()).tables == before
db.close()
@pytest.mark.parametrize("keyword", _DELETED_REMOVE_KEYWORDS)
def test_remove_rejects_deleted_cas_retry_version_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["remove"] == 0
with pytest.raises(TypeError):
db.functions.remove(
_REMOVE_CATALOG_NAME,
current,
**{keyword: True},
)
assert counters["lookup"] == 1
assert counters["remove"] == 0
def test_no_direct_remove_methods_and_function_has_no_remove_facade_private():
counters: dict[str, int] = {"lookup": 0, "remove": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "remove_function")
assert not hasattr(db, "remove_function_name")
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
assert callable(getattr(db.functions, "remove", None))
assert not hasattr(lancedb, "_SyncFunctions")
assert not hasattr(lancedb, "_AsyncFunctions")
assert type(db.functions).__name__.startswith("_")
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
assert counters["lookup"] == 1
assert counters["remove"] == 0
@@ -1,579 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python conditional first-class Function replacement."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable
from unittest import mock
import pyarrow as pa
import pytest
import lancedb
import lancedb._udf as _udf_mod
import lancedb.job
from lancedb import FunctionCapability, udf
from lancedb.exceptions import JobFailedError
_SOURCE_MARKER = "replace-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/replace-redact-token-xyz"
_SECRET_ENV = "REPLACE_API_TOKEN"
_NETWORK_ORIGIN = "https://api.replace-example.com"
_FUNCTION_NAME = "text.normalize"
_CURRENT_FUNCTION_ID = "fn.replace-current-1"
_REPLACED_FUNCTION_ID = "fn.replace-result-1"
_JOB_ID_SYNC = "job-replace-sync-1"
_JOB_ID_ASYNC = "job-replace-async-1"
_JOB_ID_CONFLICT = "job-replace-conflict-1"
_REGISTER_PATH = "/v1/functions/register"
_LOOKUP_PATH = "/v1/functions/lookup"
_DESCRIBE_PATH = "/v1/jobs/describe"
_CONFLICTING_MESSAGE_CODE = "definition_validation_failure"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# job-result tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REPLACE_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"version",
"deterministic",
"null_policy",
"replace",
"expected_current_function_id",
"alias",
"lineage",
)
_SPEC_KEYS = {
"format_version",
"name",
"definition",
"expected_current_function_id",
}
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_replace_normalize(text, limit):
"""replace-source-marker-unique-xyz."""
return text[:limit]
def _current_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _CURRENT_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _definition_json(fn: object) -> dict[str, Any]:
payload = _udf_mod._build_function_definition(fn)._to_json()
if isinstance(payload, bytes):
return json.loads(payload.decode("utf-8"))
assert isinstance(payload, str)
return json.loads(payload)
def _expected_replace_spec(name: str, current_id: str, fn: object) -> dict[str, Any]:
return {
"format_version": 1,
"name": name,
"definition": _definition_json(fn),
"expected_current_function_id": current_id,
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _assert_exact_replace_spec(
body: dict[str, Any], expected: dict[str, Any], current_id: str
) -> None:
assert set(body) == _SPEC_KEYS
assert body == expected
assert body["format_version"] == 1
assert body["expected_current_function_id"] == current_id
assert body["expected_current_function_id"] is not None
assert _SOURCE_MARKER in json.dumps(body["definition"])
assert any(
capability.get("reference") == _SECRET_REFERENCE
for capability in body["definition"]["capabilities"]
)
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact lookup; optionally continue for register/describe."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _FUNCTION_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"function": _current_function_wire()}).encode("utf-8")
)
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["register"] = counters.get("register", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected register")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
def test_sync_remote_replace_exact_body_one_package_job_and_function_result():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
register_attempts: list[dict[str, Any]] = []
function_result_wire = {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _REPLACED_FUNCTION_ID,
"signature": expected_spec["definition"]["signature"],
},
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
register_attempts.append(
{
"request_id": request.headers.get("x-request-id"),
"raw": payload,
"body": json.loads(payload.decode("utf-8")),
}
)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_SYNC}).encode("utf-8"))
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(payload.decode("utf-8"))
assert body["job_id"] == _JOB_ID_SYNC
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_SYNC,
"job_state": "DONE",
"job_type": "register_function",
"creation_ms": 1,
"spec": {},
"result": function_result_wire,
}
).encode("utf-8")
)
package_calls = {"n": 0}
original_package = _udf_mod._package_udf
def counting_package(fn: object):
package_calls["n"] += 1
return original_package(fn)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
job = db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
waited = job.wait(timeout=timedelta(seconds=5))
assert package_calls["n"] == 1
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
assert len(register_attempts) == 1
attempt = register_attempts[0]
assert isinstance(attempt["request_id"], str) and attempt["request_id"]
assert attempt["raw"]
_assert_exact_replace_spec(attempt["body"], expected_spec, current.id)
assert type(waited) is lancedb.Function
assert waited.id == _REPLACED_FUNCTION_ID
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
assert waited.output_type == pa.string()
assert waited.output_nullable is True
@pytest.mark.asyncio
async def test_async_remote_replace_exact_body_returns_async_job():
counters: dict[str, int] = {"lookup": 0, "register": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REGISTER_PATH
counters["register"] += 1
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = await db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 1
assert seen.get("raw")
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
_assert_exact_replace_spec(seen["body"], expected_spec, current.id)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
def test_sync_remote_replace_failed_name_conflict_raises_job_failed_error_code():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"job_id": _JOB_ID_CONFLICT}).encode("utf-8")
)
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(raw.decode("utf-8"))
assert body["job_id"] == _JOB_ID_CONFLICT
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_CONFLICT,
"job_type": "register_function",
"job_state": "FAILED",
"creation_ms": 1,
"spec": {},
"failure": {
"phase": "validate",
"message": (
f"looks like {_CONFLICTING_MESSAGE_CODE} during CAS"
),
"retryable": False,
"error_code": "name_conflict",
},
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert type(job) is lancedb.job.Job
with pytest.raises(JobFailedError) as exc_info:
job.wait(timeout=timedelta(seconds=5))
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
err = exc_info.value
assert isinstance(err, JobFailedError)
assert err.error_code == "name_conflict"
assert err.error_code != _CONFLICTING_MESSAGE_CODE
def test_empty_name_rejects_before_register_transport():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(ValueError):
db.functions.replace("", current, packable_replace_normalize)
assert counters["lookup"] == 1
assert counters["register"] == 0
@pytest.mark.parametrize(
"bad_current",
[
_CURRENT_FUNCTION_ID,
{"id": _CURRENT_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_current_rejected_without_register(bad_current):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-current path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME, bad_current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_local_sync_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
await db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert (await db.list_tables()).tables == before
@pytest.mark.parametrize("keyword", _DELETED_REPLACE_KEYWORDS)
def test_replace_rejects_deleted_cas_retry_version_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME,
current,
packable_replace_normalize,
**{keyword: True},
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_no_direct_replace_function_methods_and_function_has_no_replace():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "replace_function")
assert not hasattr(db, "register_function")
assert not hasattr(current, "replace")
assert not hasattr(current, "replace_function")
assert callable(getattr(db.functions, "replace", None))
assert counters["lookup"] == 1
assert counters["register"] == 0
@@ -1,728 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python exact first-class Function revocation."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pyarrow as pa
import pytest
import lancedb
from lancedb import _lancedb as _native
from lancedb.remote.errors import HttpError
_REVOKE_PATH = "/v1/functions/revoke"
_LOOKUP_PATH = "/v1/functions/lookup"
_REVOKE_CATALOG_NAME = "text.normalize.revoke-name"
_REVOKE_FUNCTION_ID = "fn.exact.revoke-handle"
_REVOKE_SERVER_MESSAGE_MARKER = (
"SERVER_REVOKE_DIAGNOSTIC_MARKER id=fn.exact.revoke-handle "
"name=text.normalize.revoke-name"
)
_SENSITIVE_BODY_MARKER = "SENSITIVE_REVOKE_BODY_MARKER"
_CONFLICTING_MESSAGE_CODE = "revoked_function"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# remove tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REVOKE_KEYWORDS = (
"function_id",
"name",
"idempotency_key",
"retry_key",
"user_version",
"version",
"reason",
"expiry",
"force",
"if_exists",
"remove",
"delete",
)
def _function_error_cls(*, required: bool = True):
"""Resolve FunctionError from the live module (records RED when absent)."""
from lancedb import exceptions as exc_mod
cls = getattr(exc_mod, "FunctionError", None)
if cls is None:
if required:
pytest.fail("lancedb.exceptions.FunctionError is missing")
return type("MissingFunctionError", (), {})
return cls
def _sample_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _REVOKE_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _lookup_success_body() -> bytes:
return json.dumps({"function": _sample_function_wire()}).encode("utf-8")
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
def _close_db(db: Any) -> None:
with contextlib.suppress(Exception):
inner = getattr(db, "_conn", None)
if inner is not None:
inner.close()
return
close = getattr(db, "close", None)
if callable(close):
close()
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
db = None
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
if db is not None:
_close_db(db)
server.shutdown()
thread.join()
def _exception_chain_text(exc: BaseException) -> str:
parts: list[str] = []
seen: set[int] = set()
current: BaseException | None = exc
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(str(current))
parts.append(repr(current))
current = current.__cause__
return "\n".join(parts)
def _assert_payload_free(exc: BaseException) -> None:
text = _exception_chain_text(exc)
assert _REVOKE_SERVER_MESSAGE_MARKER not in text
assert _REVOKE_CATALOG_NAME not in text
assert _REVOKE_FUNCTION_ID not in text
assert _SENSITIVE_BODY_MARKER not in text
def _assert_exact_revoke_function(function: object) -> None:
assert type(function) is lancedb.Function
assert function.id == _REVOKE_FUNCTION_ID
assert not hasattr(function, "name")
assert function.parameters == (
("text", pa.string()),
("limit", pa.int32()),
)
assert function.output_type == pa.string()
assert function.output_nullable is True
assert _REVOKE_CATALOG_NAME not in repr(function)
assert _REVOKE_CATALOG_NAME not in str(function)
def _assert_exact_revoke_request(
request: http.server.BaseHTTPRequestHandler,
raw: bytes,
body: dict[str, Any],
*,
expected_id: str,
) -> None:
assert request.command == "POST"
assert request.path == _REVOKE_PATH
assert "?" not in request.path
assert "remove" not in request.path
assert raw
assert body == {"function_id": expected_id}
assert set(body) == {"function_id"}
assert "name" not in body
assert "expected_current_function_id" not in body
assert "format_version" not in body
assert "function" not in body
assert "signature" not in body
assert "job_id" not in body
assert "idempotency_key" not in body
assert "user_version" not in body
assert "reason" not in body
assert "expiry" not in body
assert "force" not in body
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
def _assert_native_revoke_method_present() -> None:
assert hasattr(_native.Connection, "_revoke_function")
assert callable(getattr(_native.Connection, "_revoke_function"))
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact name lookup; optionally continue for revoke."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _REVOKE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["revoke"] = counters.get("revoke", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected revoke")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_REVOKE_CATALOG_NAME)
_assert_exact_revoke_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_REVOKE_CATALOG_NAME)
_assert_exact_revoke_function(current)
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
return current
def test_native_connection_exposes_private_revoke_function():
_assert_native_revoke_method_present()
def test_sync_remote_revoke_exact_body_path_request_id_returns_none():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
revoke_attempts: list[dict[str, Any]] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(payload.decode("utf-8"))
revoke_attempts.append(
{
"request": request,
"raw": payload,
"body": body,
"request_id": request.headers.get("x-request-id"),
}
)
# Illegal body on 204 must be ignored; success is status-driven only.
request.send_response(204)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
_SENSITIVE_BODY_MARKER: True,
"message": _REVOKE_SERVER_MESSAGE_MARKER,
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "revoke_function")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
result = db.functions.revoke(current)
assert result is None
assert counters["lookup"] == 1
assert counters["revoke"] == 1
assert len(revoke_attempts) == 1
attempt = revoke_attempts[0]
_assert_exact_revoke_request(
attempt["request"],
attempt["raw"],
attempt["body"],
expected_id=current.id,
)
assert attempt["body"]["function_id"] == current.id
_assert_exact_revoke_function(current)
@pytest.mark.asyncio
async def test_async_remote_revoke_exact_body_returns_none():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
seen["request"] = request
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(204)
request.end_headers()
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "revoke_function")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
result = await db.functions.revoke(current)
assert result is None
assert counters["lookup"] == 1
assert counters["revoke"] == 1
assert seen.get("raw")
_assert_exact_revoke_request(
seen["request"],
seen["raw"],
seen["body"],
expected_id=current.id,
)
def test_repeated_remote_revoke_204_both_return_none():
"""Two logical calls each receiving 204 both succeed (Python outcome only)."""
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
revoke_request_ids: list[str] = []
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, body, expected_id=_REVOKE_FUNCTION_ID
)
request_id = request.headers.get("x-request-id")
assert isinstance(request_id, str) and request_id
revoke_request_ids.append(request_id)
request.send_response(204)
request.end_headers()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
first = db.functions.revoke(current)
second = db.functions.revoke(current)
assert first is None
assert second is None
assert counters["lookup"] == 1
assert counters["revoke"] == 2
assert len(revoke_request_ids) == 2
_assert_exact_revoke_function(current)
def test_after_revoke_name_and_id_lookup_still_return_function():
"""Revoke does not unlink names; SDK-visible sequence only, not Sophon proof."""
_assert_native_revoke_method_present()
counters: dict[str, int] = {
"lookup_name": 0,
"lookup_id": 0,
"revoke": 0,
}
revoked = {"yes": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
body = json.loads(raw.decode("utf-8"))
if "name" in body:
counters["lookup_name"] += 1
assert body == {"name": _REVOKE_CATALOG_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
counters["lookup_id"] += 1
assert body == {"function_id": _REVOKE_FUNCTION_ID}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(_lookup_success_body())
return
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
body = json.loads(raw.decode("utf-8"))
_assert_exact_revoke_request(
request, raw, body, expected_id=_REVOKE_FUNCTION_ID
)
revoked["yes"] = True
request.send_response(204)
request.end_headers()
with _mock_remote_db(handler) as db:
current = _observe_current(db)
assert counters["lookup_name"] == 1
assert counters["lookup_id"] == 0
assert counters["revoke"] == 0
assert not revoked["yes"]
result = db.functions.revoke(current)
assert result is None
assert counters["lookup_name"] == 1
assert counters["revoke"] == 1
assert revoked["yes"]
by_name = db.functions.get(_REVOKE_CATALOG_NAME)
by_id = db.functions.get_by_id(_REVOKE_FUNCTION_ID)
assert counters["lookup_name"] == 2
assert counters["lookup_id"] == 1
assert counters["revoke"] == 1
_assert_exact_revoke_function(by_name)
_assert_exact_revoke_function(by_id)
assert by_name.id == current.id
assert by_id.id == current.id
assert by_name.parameters == current.parameters
assert by_id.parameters == current.parameters
assert by_name.output_type == current.output_type
assert by_id.output_type == current.output_type
assert by_name.output_nullable is current.output_nullable
assert by_id.output_nullable is current.output_nullable
def test_explicit_name_or_function_not_found_is_function_error_payload_free():
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
body = {
"error_code": "name_or_function_not_found",
"message": (
f"{_REVOKE_SERVER_MESSAGE_MARKER} looks_like {_CONFLICTING_MESSAGE_CODE} "
"name_conflict"
),
"name": _REVOKE_CATALOG_NAME,
"function_id": _REVOKE_FUNCTION_ID,
_SENSITIVE_BODY_MARKER: True,
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
)
request.send_response(404)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(body).encode("utf-8"))
function_error = _function_error_cls()
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(function_error) as exc_info:
db.functions.revoke(current)
assert counters["lookup"] == 1
assert counters["revoke"] == 1
err = exc_info.value
assert isinstance(err, function_error)
assert err.code == "name_or_function_not_found"
assert err.code != _CONFLICTING_MESSAGE_CODE
assert err.code != "name_conflict"
_assert_payload_free(err)
@pytest.mark.parametrize(
"label,status,response_body",
[
(
"200_with_body",
200,
{
"ok": True,
"message": _REVOKE_SERVER_MESSAGE_MARKER,
_SENSITIVE_BODY_MARKER: True,
"job_id": "must-not-infer-job",
},
),
(
"202_empty",
202,
f"{_REVOKE_SERVER_MESSAGE_MARKER} {_SENSITIVE_BODY_MARKER}",
),
("200_empty", 200, ""),
],
)
def test_http_200_202_cannot_return_success(
label: str, status: int, response_body: object
):
del label
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
assert request.path == _REVOKE_PATH
counters["revoke"] += 1
parsed = json.loads(payload.decode("utf-8"))
_assert_exact_revoke_request(
request, payload, parsed, expected_id=_REVOKE_FUNCTION_ID
)
request.send_response(status)
request.send_header("Content-Type", "application/json")
request.end_headers()
if isinstance(response_body, str):
request.wfile.write(response_body.encode("utf-8"))
else:
request.wfile.write(json.dumps(response_body).encode("utf-8"))
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(HttpError) as exc_info:
db.functions.revoke(current)
assert counters["lookup"] == 1
assert counters["revoke"] == 1
err = exc_info.value
assert isinstance(err, HttpError)
assert not isinstance(err, _function_error_cls(required=False))
_assert_payload_free(err)
@pytest.mark.parametrize(
"bad_function",
[
_REVOKE_FUNCTION_ID,
{"id": _REVOKE_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_function_rejected_without_revoke(bad_function):
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-function path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(TypeError):
db.functions.revoke(bad_function)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
def test_local_sync_revoke_not_implemented_without_table_mutation(tmp_path):
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "revoke_function")
with pytest.raises(NotImplementedError):
db.functions.revoke(current)
assert db.list_tables().tables == before
_close_db(db)
@pytest.mark.asyncio
async def test_local_async_revoke_not_implemented_without_table_mutation(tmp_path):
_assert_native_revoke_method_present()
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "revoke_function")
with pytest.raises(NotImplementedError):
await db.functions.revoke(current)
assert (await db.list_tables()).tables == before
db.close()
@pytest.mark.parametrize("keyword", _DELETED_REVOKE_KEYWORDS)
def test_revoke_rejects_overdesigned_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["revoke"] == 0
with pytest.raises(TypeError):
db.functions.revoke(current, **{keyword: True})
assert counters["lookup"] == 1
assert counters["revoke"] == 0
def test_no_direct_revoke_methods_and_function_has_no_revoke_facade_private():
counters: dict[str, int] = {"lookup": 0, "revoke": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "revoke_function")
assert not hasattr(current, "remove")
assert not hasattr(current, "delete")
assert not hasattr(current, "revoke")
assert callable(getattr(db.functions, "revoke", None))
assert not hasattr(lancedb, "_SyncFunctions")
assert not hasattr(lancedb, "_AsyncFunctions")
assert type(db.functions).__name__.startswith("_")
assert "_SyncFunctions" not in getattr(lancedb, "__all__", [])
assert "_AsyncFunctions" not in getattr(lancedb, "__all__", [])
assert counters["lookup"] == 1
assert counters["revoke"] == 0
File diff suppressed because it is too large Load Diff
@@ -1,899 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python ``table.add_generated_column`` (FF-032).
Public user shape under test:
job = table.add_generated_column(
"normalized_text",
normalize(text=col("text")),
)
job.wait()
These tests exercise the live worktree PyO3 extension and public sync/async
wrappers. While the public methods and hidden native bridge are absent they
fail against that extension; once present they freeze the public contract
below. They must not fake success paths.
"""
from __future__ import annotations
import contextlib
import http.server
import inspect
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable
import pytest
import lancedb
import lancedb.job
from lancedb import _lancedb as _native
from lancedb.expr import col
from lancedb.remote.table import RemoteTable
from lancedb.table import AsyncTable, LanceTable, Table
_LOOKUP_PATH = "/v1/functions/lookup"
_JOB_DESCRIBE_PATH = "/v1/jobs/describe"
_TABLE_NAME = "articles"
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/generated_columns/create/"
_BRANCHES_CREATE_PATH = f"/v1/table/{_TABLE_NAME}/branches/create/"
_BRANCHES_LIST_PATH = f"/v1/table/{_TABLE_NAME}/branches/list/"
_CATALOG_NAME = "text.normalize"
_FUNCTION_ID = "fn.exact.normalize.gen-col"
_JOB_ID_SYNC = "job-create-gen-col-sync-1"
_JOB_ID_ASYNC = "job-create-gen-col-async-1"
_JOB_ID_BRANCH = "job-create-gen-col-branch-1"
_SOURCE_TABLE_VERSION = 42
_TEXT_FIELD_ID = 7
_BRANCH_NAME = "exp"
_BRANCH_SOURCE_VERSION = 9
_BRANCH_TEXT_FIELD_ID = 11
_DESCRIBE_BODY_MARKER = "SENSITIVE_DESCRIBE_BODY_MARKER_gen_col_xyz"
_CREATE_RESPONSE_MARKER = "SENSITIVE_CREATE_RESPONSE_MARKER_gen_col_xyz"
_LITERAL_PAYLOAD_SENTINEL = "LITERAL_PAYLOAD_SENTINEL_gen_col_xyz"
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_FORBIDDEN_PUBLIC_NAMES = (
"FunctionCall",
"BoundFunctionCall",
"AuthoredFunctionCall",
"CreateGeneratedColumnRequest",
"CreateGeneratedColumnJobSpec",
"GeneratedColumnBindingSnapshot",
"GeneratedColumnCreateRequest",
"geneva",
"GenevaFunction",
"VirtualColumnDefinition",
)
_FORBIDDEN_METHOD_KWARGS = (
"source_table_version",
"version",
"field_id",
"field_ids",
"output",
"output_type",
"output_nullable",
"nullable",
"spec",
"retry_key",
"idempotency_key",
"request",
"envelope",
"table_ref",
"branch",
)
def _sample_function_wire(
*,
function_id: str = _FUNCTION_ID,
parameters: list[dict[str, str]] | None = None,
) -> dict[str, Any]:
return {
"format_version": 1,
"id": function_id,
"signature": {
"parameters": parameters
or [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _text_schema_fields(
*, arrow_type: str = "string", nullable: bool = True
) -> dict[str, Any]:
return {
"fields": [
{
"name": "text",
"type": {"type": arrow_type},
"nullable": nullable,
}
]
}
def _describe_body(
*,
version: int = _SOURCE_TABLE_VERSION,
field_ids: list[int] | None = None,
arrow_type: str = "string",
include_marker: bool = True,
) -> dict[str, Any]:
body: dict[str, Any] = {
"version": version,
"schema": _text_schema_fields(arrow_type=arrow_type),
"field_ids": field_ids if field_ids is not None else [_TEXT_FIELD_ID],
}
if include_marker:
body["server_diagnostic"] = _DESCRIBE_BODY_MARKER
return body
def _create_gen_column_done_body(job_id: str) -> dict[str, Any]:
# DONE with omitted result: create_gen_column projects JobResult::None.
return {
"job_id": job_id,
"job_state": "DONE",
"job_type": "create_gen_column",
"creation_ms": 1,
"spec": {},
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def _json_response(
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
) -> None:
payload = json.dumps(body).encode("utf-8")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(payload)
def _lookup_function(db: Any) -> lancedb.Function:
return db.functions.get(_CATALOG_NAME)
class _RequestLog:
"""Track lookup/describe/create after setup; setup traffic is excluded."""
def __init__(self) -> None:
self.lookup: list[dict[str, Any]] = []
self.describe: list[dict[str, Any]] = []
self.create: list[dict[str, Any]] = []
self.other_table: list[str] = []
self.recording = False
def start(self) -> None:
# Drop setup's explicit Function lookup and open_table describe so
# operation accounting cannot be polluted by fixture traffic.
self.lookup.clear()
self.describe.clear()
self.create.clear()
self.other_table.clear()
self.recording = True
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
if not self.recording:
return
if path == _LOOKUP_PATH:
self.lookup.append(body or {})
elif path == _DESCRIBE_PATH:
self.describe.append(body or {})
elif path == _CREATE_PATH:
self.create.append(body or {})
elif path.startswith(f"/v1/table/{_TABLE_NAME}/"):
self.other_table.append(path)
def _assert_no_operation_traffic(log: _RequestLog) -> None:
assert log.lookup == []
assert log.describe == []
assert log.create == []
assert log.other_table == []
def _assert_exact_public_signature(method: Any) -> None:
"""Freeze ``(self, column_name, call)`` with no varargs/kwargs escape hatches."""
params = list(inspect.signature(method).parameters.values())
assert [p.name for p in params] == ["self", "column_name", "call"]
for param in params:
assert param.kind in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
assert param.default is inspect.Parameter.empty
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
assert param.kind is not inspect.Parameter.VAR_KEYWORD
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
def _open_table_and_function(
*,
describe_body: dict[str, Any] | None = None,
on_create: Callable[[dict[str, Any], http.server.BaseHTTPRequestHandler], None]
| None = None,
job_id: str = _JOB_ID_SYNC,
support_branch_create: bool = False,
function_wire: dict[str, Any] | None = None,
):
"""Open remote table + immutable Function; return (db, table, function, log, cm)."""
log = _RequestLog()
binding_describe = describe_body or _describe_body()
open_describe = {
"version": 1,
"schema": _text_schema_fields(),
}
state = {"opened": False}
wire = function_wire or _sample_function_wire()
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": wire})
return
if request.path == _JOB_DESCRIBE_PATH:
assert body["job_id"] == job_id
_json_response(request, _create_gen_column_done_body(job_id))
return
if support_branch_create and request.path == _BRANCHES_CREATE_PATH:
log.note(request.path, body)
_json_response(request, {})
return
if support_branch_create and request.path == _BRANCHES_LIST_PATH:
log.note(request.path, body)
_json_response(
request,
{
"branches": {
_BRANCH_NAME: {
"parentBranch": None,
"parentVersion": 1,
"createAt": 1,
"manifestSize": 1,
}
}
},
)
return
if request.path == _DESCRIBE_PATH:
# First describe seeds open_table; later ones are binding snapshots.
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
if on_create is not None:
on_create(body, request)
return
_json_response(
request,
{
"job_id": job_id,
"server_extra": {"marker": _CREATE_RESPONSE_MARKER},
},
)
return
if request.path.startswith(f"/v1/table/{_TABLE_NAME}/"):
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db(handler)
db = cm.__enter__()
function = _lookup_function(db)
table = db.open_table(_TABLE_NAME)
assert isinstance(table, RemoteTable)
# open_table consumed the seed describe; binding/create accounting starts now.
# Setup's one explicit lookup is cleared here and must not pollute counts.
log.start()
return db, table, function, log, cm
def _assert_exact_create_envelope(
body: dict[str, Any],
*,
source_table_version: int,
column_name: str,
field_id: int,
branch: str | None = None,
) -> None:
expected_keys = {"source_table_version", "spec"}
if branch is not None:
expected_keys.add("branch")
assert set(body) == expected_keys
assert body["source_table_version"] == source_table_version
assert "table_ref" not in body
if branch is None:
assert "branch" not in body
else:
assert body["branch"] == branch
spec = body["spec"]
assert set(spec) == {"format_version", "column_name", "function_call"}
assert spec["format_version"] == 1
assert spec["column_name"] == column_name
for forbidden in (
"table_ref",
"source_table_version",
"version",
"output",
"output_type",
"output_field_id",
"dependency_epoch",
"materialized_epoch",
"idempotency_key",
"retry_key",
"name",
"handle",
"artifact",
"geneva",
):
assert forbidden not in spec
call = spec["function_call"]
assert set(call) == {"function_id", "arguments"}
assert call["function_id"] == _FUNCTION_ID
assert len(call["arguments"]) == 1
binding = call["arguments"][0]
assert binding["parameter"] == "text"
value = binding["value"]
assert value["kind"] == "field"
assert value["field_id"] == field_id
assert value["data_type_ipc"] == _UTF8_TYPE_IPC_B64
assert "name" not in value
assert "column_name" not in value
assert "text" not in value
# Serialized call must not late-bind by column name anywhere relevant.
dumped = json.dumps(call)
assert '"column_name"' not in dumped
assert "normalized_text" not in dumped
def test_public_and_native_add_generated_column_seams_must_exist():
"""Public sync/async methods and the private native bridge must exist."""
assert hasattr(_native.Table, "_add_generated_column"), (
"native private bridge Table._add_generated_column is missing"
)
assert hasattr(AsyncTable, "add_generated_column"), (
"AsyncTable.add_generated_column is missing"
)
assert hasattr(Table, "add_generated_column"), (
"Table.add_generated_column is missing"
)
assert hasattr(LanceTable, "add_generated_column"), (
"LanceTable.add_generated_column is missing"
)
assert hasattr(RemoteTable, "add_generated_column"), (
"RemoteTable.add_generated_column is missing"
)
# Once present, freeze the exact public positional surface.
_assert_exact_public_signature(Table.add_generated_column)
_assert_exact_public_signature(LanceTable.add_generated_column)
_assert_exact_public_signature(RemoteTable.add_generated_column)
_assert_exact_public_signature(AsyncTable.add_generated_column)
def test_sync_remote_add_generated_column_returns_job_without_eager_wrapper_mutation():
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
try:
# Capture public wrapper state before the operation window.
schema_before = table.schema
version_before = table.version
log.start()
call = normalize(text=col("text"))
# Exact public argument order from the frozen user example.
job = table.add_generated_column(
"normalized_text",
call,
)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
# Exact success path stops after submit: one binding describe, one create,
# and no catalog re-lookup. Do not wait yet.
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.other_table == []
# Public schema/version through the existing wrapper must still reflect
# the pre-submit table: generated column is not published by Job accept.
# Access both before wait so eager wrapper cache invalidation / refresh /
# version advancement is observable.
schema_after = table.schema
assert "normalized_text" not in schema_after.names
assert schema_after == schema_before
# Schema must be served from the existing wrapper cache — no extra
# describe beyond the one binding snapshot.
assert len(log.describe) == 1
assert len(log.lookup) == 0
assert len(log.create) == 1
version_after = table.version
assert version_after == version_before
# Public Remote ``version`` always describes once by design; that probe
# must not drag a schema-cache miss, create, or catalog lookup with it.
assert len(log.describe) == 2
assert len(log.lookup) == 0
assert len(log.create) == 1
assert log.other_table == []
waited = job.wait()
assert waited is None
assert len(log.lookup) == 0
assert len(log.create) == 1
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_remote_add_generated_column_returns_async_job_and_wait_none():
log = _RequestLog()
state = {"opened": False}
binding_describe = _describe_body()
open_describe = {"version": 1, "schema": _text_schema_fields()}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": _sample_function_wire()})
return
if request.path == _JOB_DESCRIBE_PATH:
assert body["job_id"] == _JOB_ID_ASYNC
_json_response(request, _create_gen_column_done_body(_JOB_ID_ASYNC))
return
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
_json_response(request, {"job_id": _JOB_ID_ASYNC})
return
request.send_response(404)
request.end_headers()
async with _mock_remote_db_async(handler) as db:
normalize = await db.functions.get(_CATALOG_NAME)
table = await db.open_table(_TABLE_NAME)
log.start()
call = normalize(text=col("text"))
job = await table.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
waited = await job.wait()
assert waited is None
assert len(log.lookup) == 0
assert len(log.create) == 1
def test_remote_add_generated_column_one_describe_one_create_exact_envelope():
db, table, normalize, log, cm = _open_table_and_function(job_id=_JOB_ID_SYNC)
try:
call = normalize(text=col("text"))
job = table.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.other_table == []
_assert_exact_create_envelope(
log.create[0],
source_table_version=_SOURCE_TABLE_VERSION,
column_name="normalized_text",
field_id=_TEXT_FIELD_ID,
)
finally:
cm.__exit__(None, None, None)
def test_remote_branch_add_generated_column_includes_exact_branch_identity():
branch_describe = _describe_body(
version=_BRANCH_SOURCE_VERSION,
field_ids=[_BRANCH_TEXT_FIELD_ID],
)
db, table, normalize, log, cm = _open_table_and_function(
describe_body=branch_describe,
job_id=_JOB_ID_BRANCH,
support_branch_create=True,
)
try:
branched = table.branches.create(_BRANCH_NAME)
assert isinstance(branched, RemoteTable)
assert branched.current_branch() == _BRANCH_NAME
log.start()
call = normalize(text=col("text"))
job = branched.add_generated_column("normalized_text", call)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_BRANCH
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert len(log.create) == 1
assert log.describe[0].get("branch") == _BRANCH_NAME
_assert_exact_create_envelope(
log.create[0],
source_table_version=_BRANCH_SOURCE_VERSION,
column_name="normalized_text",
field_id=_BRANCH_TEXT_FIELD_ID,
branch=_BRANCH_NAME,
)
finally:
cm.__exit__(None, None, None)
def test_empty_column_name_fails_locally_with_zero_table_requests():
db, table, normalize, log, cm = _open_table_and_function()
try:
# Authored call owns a real literal so payload-free failure is not vacuous.
call = normalize(text=_LITERAL_PAYLOAD_SENTINEL)
with pytest.raises((ValueError, TypeError)) as raised:
table.add_generated_column("", call)
text = _exception_text(raised.value)
lowered = text.lower()
assert "column" in lowered or "empty" in lowered or "non-empty" in lowered
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert _LITERAL_PAYLOAD_SENTINEL not in text
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.parametrize(
("column_ref", "expected_token"),
[
("missing_text", "missing_text"),
("Text", "Text"), # exact-case mismatch against schema field "text"
],
)
def test_missing_or_case_mismatch_column_one_describe_zero_create(
column_ref: str, expected_token: str
):
db, table, normalize, log, cm = _open_table_and_function()
try:
call = normalize(text=col(column_ref))
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert expected_token in text
assert "text" in text # parameter name from the Function signature
assert "missing" in text.lower() or "field" in text.lower()
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
def test_type_mismatch_one_describe_zero_create_identifies_parameter():
db, table, normalize, log, cm = _open_table_and_function(
describe_body=_describe_body(arrow_type="int32"),
)
try:
call = normalize(text=col("text"))
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "text" in text
assert "type" in text.lower() or "mismatch" in text.lower()
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
def test_literal_payload_stays_out_of_field_binding_failure():
"""Authored call owns a real literal; later field binding fails payload-free."""
wire = _sample_function_wire(
parameters=[
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "prefix", "data_type_ipc": _UTF8_TYPE_IPC_B64},
]
)
db, table, normalize, log, cm = _open_table_and_function(function_wire=wire)
try:
call = normalize(text=col("missing_text"), prefix=_LITERAL_PAYLOAD_SENTINEL)
with pytest.raises(ValueError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "missing_text" in text
assert _LITERAL_PAYLOAD_SENTINEL not in text
assert _DESCRIBE_BODY_MARKER not in text
assert _CREATE_RESPONSE_MARKER not in text
assert len(log.lookup) == 0
assert len(log.describe) == 1
assert log.create == []
assert log.other_table == []
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_closed_async_table_fails_with_zero_operation_requests():
log = _RequestLog()
state = {"opened": False}
binding_describe = _describe_body()
open_describe = {"version": 1, "schema": _text_schema_fields()}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _LOOKUP_PATH:
log.note(request.path, body)
_json_response(request, {"function": _sample_function_wire()})
return
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding_describe)
return
if request.path == _CREATE_PATH:
log.note(request.path, body)
_json_response(request, {"job_id": _JOB_ID_ASYNC})
return
request.send_response(404)
request.end_headers()
async with _mock_remote_db_async(handler) as db:
normalize = await db.functions.get(_CATALOG_NAME)
table = await db.open_table(_TABLE_NAME)
call = normalize(text=col("text"))
# Public close only — do not mutate private implementation fields.
table.close()
log.start()
try:
await table.add_generated_column("normalized_text", call)
except AttributeError:
# Method missing: re-raise so the failure names the public seam.
raise
except Exception as exc:
text = _exception_text(exc)
assert "closed" in text.lower()
else:
pytest.fail("closed AsyncTable must fail before transport")
_assert_no_operation_traffic(log)
def test_rejects_non_authored_call_before_any_operation_request():
db, table, normalize, log, cm = _open_table_and_function()
try:
bad_values = (
normalize, # exact Function handle itself
{"text": "x"},
col("text"), # direct query Expr
object(),
)
for bad in bad_values:
with pytest.raises(TypeError):
table.add_generated_column("normalized_text", bad)
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
def test_native_valid_call_returns_not_supported_without_mutation(tmp_path):
# Immutable Function handle is connection-free; obtain it via remote lookup.
def lookup_only(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.path == _LOOKUP_PATH
_read_body(request)
_json_response(request, {"function": _sample_function_wire()})
with _mock_remote_db(lookup_only) as remote_db:
normalize = _lookup_function(remote_db)
db = lancedb.connect(tmp_path)
table = db.create_table(_TABLE_NAME, [{"text": "Hello"}, {"text": "World"}])
assert isinstance(table, LanceTable)
version_before = table.version
schema_before = table.schema
rows_before = table.to_arrow().to_pylist()
call = normalize(text=col("text"))
with pytest.raises(NotImplementedError) as raised:
table.add_generated_column("normalized_text", call)
text = _exception_text(raised.value)
assert "not supported" in text.lower() or "submit_create_generated_column" in text
assert "add_columns" not in text.lower()
assert table.version == version_before
assert table.schema == schema_before
assert "normalized_text" not in table.schema.names
assert table.to_arrow().to_pylist() == rows_before
def test_public_surface_is_minimal_and_private_call_stays_opaque():
for name in _FORBIDDEN_PUBLIC_NAMES:
assert name not in getattr(lancedb, "__all__", [])
assert not hasattr(lancedb, name)
assert not hasattr(lancedb, "_FunctionCall")
authored_type = getattr(_native, "_FunctionCall", None)
assert authored_type is not None
with pytest.raises(TypeError):
authored_type()
# When the public method exists, reject overdesign kwargs and keep the frozen
# positional surface: (self, column_name, call).
if hasattr(Table, "add_generated_column"):
_assert_exact_public_signature(Table.add_generated_column)
for keyword in _FORBIDDEN_METHOD_KWARGS:
assert (
keyword not in inspect.signature(Table.add_generated_column).parameters
)
db, table, normalize, log, cm = _open_table_and_function()
try:
call = normalize(text=col("text"))
for keyword in _FORBIDDEN_METHOD_KWARGS:
with pytest.raises(TypeError):
table.add_generated_column(
"normalized_text",
call,
**{keyword: object()},
)
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
if hasattr(LanceTable, "add_generated_column"):
_assert_exact_public_signature(LanceTable.add_generated_column)
if hasattr(RemoteTable, "add_generated_column"):
_assert_exact_public_signature(RemoteTable.add_generated_column)
if hasattr(AsyncTable, "add_generated_column"):
_assert_exact_public_signature(AsyncTable.add_generated_column)
for keyword in _FORBIDDEN_METHOD_KWARGS:
assert (
keyword
not in inspect.signature(AsyncTable.add_generated_column).parameters
)
File diff suppressed because it is too large Load Diff
@@ -1,672 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python ``table.generated_column_status`` (B3d2).
Public user shape under test:
status = table.generated_column_status("complete_col") # "complete" | "incomplete"
These tests exercise the live worktree PyO3 extension and public sync/async
wrappers. While the public methods and hidden native bridge are absent they
fail against that extension; once present they freeze the public contract
below. They must not fake success paths.
"""
from __future__ import annotations
import contextlib
import http.server
import inspect
import json
import threading
from collections.abc import AsyncIterator, Iterator
from typing import Any, Callable, Literal, get_type_hints
import pytest
import lancedb
import lancedb.table
from lancedb import _lancedb as _native
from lancedb.remote.table import RemoteTable
from lancedb.table import AsyncTable, LanceTable, Table
_TABLE_NAME = "articles"
_DESCRIBE_PATH = f"/v1/table/{_TABLE_NAME}/describe/"
_ORDINARY_FIELD_ID = 1
_COMPLETE_FIELD_ID = 5
_INCOMPLETE_FIELD_ID = 7
_STABLE_FIELD_IDS = [_ORDINARY_FIELD_ID, _COMPLETE_FIELD_ID, _INCOMPLETE_FIELD_ID]
_STATUS_FUNCTION_ID = "fn.exact.status.projection"
_METADATA_KEY = "lancedb::generated_column"
_RAW_METADATA_MARKER = "SENSITIVE_STATUS_METADATA_MARKER_b3d2_py_9f2e"
# Pinned Rust-canonical schema-only Utf8 type IPC (base64), shared with FF-028.
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_EXPECTED_RETURN = Literal["complete", "incomplete"]
_FORBIDDEN_PUBLIC_NAMES = (
"GeneratedColumnStatus",
"GeneratedColumnDefinition",
"GeneratedColumnBindingSnapshot",
"GeneratedColumnBindingEntry",
)
_FORBIDDEN_BRIDGE_KWARGS = (
"epoch",
"dependency_epoch",
"materialized_epoch",
"function_id",
"field_id",
"field_ids",
"version",
"branch",
"wait",
"job",
"request",
"backend",
)
def _definition_metadata_json(
output_field_id: int,
dependency_epoch: int,
materialized_epoch: int,
*,
text_field_id: int = _ORDINARY_FIELD_ID,
) -> str:
"""Exact JSON stored under Arrow field metadata ``lancedb::generated_column``."""
return json.dumps(
{
"format_version": 1,
"output_field_id": output_field_id,
"function_call": {
"function_id": _STATUS_FUNCTION_ID,
"arguments": [
{
"parameter": "text",
"value": {
"kind": "field",
"field_id": text_field_id,
"data_type_ipc": _UTF8_TYPE_IPC_B64,
},
}
],
},
"dependency_epoch": dependency_epoch,
"materialized_epoch": materialized_epoch,
},
separators=(",", ":"),
)
def _field(
name: str,
*,
arrow_type: str = "string",
nullable: bool = True,
metadata: dict[str, str] | None = None,
) -> dict[str, Any]:
body: dict[str, Any] = {
"name": name,
"type": {"type": arrow_type},
"nullable": nullable,
}
if metadata is not None:
body["metadata"] = metadata
return body
def _status_schema_fields(
*,
complete_meta: str | None = None,
incomplete_meta: str | None = None,
bad_name: str | None = None,
bad_meta: str | None = None,
) -> dict[str, Any]:
fields = [
_field("ordinary", arrow_type="string"),
_field(
"complete_col",
arrow_type="int32",
metadata={
_METADATA_KEY: complete_meta
if complete_meta is not None
else _definition_metadata_json(_COMPLETE_FIELD_ID, 3, 3)
},
),
_field(
"incomplete_col",
arrow_type="int32",
metadata={
_METADATA_KEY: incomplete_meta
if incomplete_meta is not None
else _definition_metadata_json(_INCOMPLETE_FIELD_ID, 4, 1)
},
),
]
if bad_name is not None and bad_meta is not None:
fields.append(
_field(
bad_name,
arrow_type="int32",
metadata={_METADATA_KEY: bad_meta},
)
)
return {"fields": fields}
def _describe_body(
*,
version: int = 11,
field_ids: list[int] | None = _STABLE_FIELD_IDS,
schema: dict[str, Any] | None = None,
) -> dict[str, Any]:
body: dict[str, Any] = {
"version": version,
"schema": schema if schema is not None else _status_schema_fields(),
}
if field_ids is not None:
body["field_ids"] = field_ids
return body
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _exception_text(exc: BaseException) -> str:
parts = [str(exc), repr(exc)]
current: BaseException | None = exc
seen: set[int] = set()
while current is not None and id(current) not in seen:
seen.add(id(current))
parts.append(f"{type(current).__name__}: {current}")
current = current.__cause__ or current.__context__
return "\n".join(parts)
def _json_response(
request: http.server.BaseHTTPRequestHandler, body: dict[str, Any]
) -> None:
payload = json.dumps(body).encode("utf-8")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(payload)
class _RequestLog:
"""Track post-open describe and any non-describe operation traffic."""
def __init__(self) -> None:
self.describe: list[dict[str, Any]] = []
self.other: list[str] = []
self.recording = False
def start(self) -> None:
self.describe.clear()
self.other.clear()
self.recording = True
def note(self, path: str, body: dict[str, Any] | None = None) -> None:
if not self.recording:
return
if path == _DESCRIBE_PATH:
self.describe.append(body or {})
else:
self.other.append(path)
def _assert_no_operation_traffic(log: _RequestLog) -> None:
assert log.describe == []
assert log.other == []
def _assert_one_status_describe(log: _RequestLog) -> None:
assert len(log.describe) == 1, f"expected one status describe, got {log.describe!r}"
assert log.other == [], f"unexpected non-describe traffic: {log.other!r}"
def _assert_exact_public_signature(method: Any) -> None:
"""Freeze ``(self, column_name)`` with no varargs/kwargs/keyword-only escape."""
params = list(inspect.signature(method).parameters.values())
assert [p.name for p in params] == ["self", "column_name"]
for param in params:
assert param.kind in (
inspect.Parameter.POSITIONAL_ONLY,
inspect.Parameter.POSITIONAL_OR_KEYWORD,
)
assert param.default is inspect.Parameter.empty
assert param.kind is not inspect.Parameter.VAR_POSITIONAL
assert param.kind is not inspect.Parameter.VAR_KEYWORD
assert param.kind is not inspect.Parameter.KEYWORD_ONLY
def _assert_status_string(value: Any, expected: str) -> None:
assert value == expected
assert type(value) is str
assert value in ("complete", "incomplete")
def _open_remote_table(
*,
status_describe: dict[str, Any] | None = None,
):
"""Open sync RemoteTable; return (table, log, cm)."""
log = _RequestLog()
binding = status_describe if status_describe is not None else _describe_body()
open_describe = {
"version": 1,
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
}
state = {"opened": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding)
return
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db(handler)
db = cm.__enter__()
table = db.open_table(_TABLE_NAME)
assert isinstance(table, RemoteTable)
log.start()
return table, log, cm
async def _open_remote_table_async(
*,
status_describe: dict[str, Any] | None = None,
):
"""Open async table under a live mock server; return (table, log, cm)."""
log = _RequestLog()
binding = status_describe if status_describe is not None else _describe_body()
open_describe = {
"version": 1,
"schema": {"fields": [_field("ordinary", arrow_type="string")]},
}
state = {"opened": False}
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
body = json.loads(raw.decode("utf-8")) if raw else {}
if request.path == _DESCRIBE_PATH:
if not state["opened"]:
state["opened"] = True
_json_response(request, open_describe)
return
log.note(request.path, body)
_json_response(request, binding)
return
log.note(request.path, body)
request.send_response(404)
request.end_headers()
request.wfile.write(b"unexpected path")
cm = _mock_remote_db_async(handler)
db = await cm.__aenter__()
table = await db.open_table(_TABLE_NAME)
assert isinstance(table, AsyncTable)
log.start()
return table, log, cm
def test_no_public_generated_column_status_resource_exported():
"""Baseline: no public status class/enum/resource is exported."""
for mod in (lancedb, lancedb.table, _native):
for name in _FORBIDDEN_PUBLIC_NAMES:
assert not hasattr(mod, name), f"{mod.__name__}.{name} must not be public"
def test_public_surface_signatures_annotations_and_hidden_bridge():
"""Four public methods + hidden native bridge must exist with frozen shape."""
assert hasattr(_native.Table, "_generated_column_status"), (
"native private bridge Table._generated_column_status is missing"
)
assert hasattr(Table, "generated_column_status"), (
"Table.generated_column_status is missing"
)
assert hasattr(LanceTable, "generated_column_status"), (
"LanceTable.generated_column_status is missing"
)
assert hasattr(RemoteTable, "generated_column_status"), (
"RemoteTable.generated_column_status is missing"
)
assert hasattr(AsyncTable, "generated_column_status"), (
"AsyncTable.generated_column_status is missing"
)
bridge = _native.Table._generated_column_status
_assert_exact_public_signature(bridge)
for keyword in _FORBIDDEN_BRIDGE_KWARGS:
assert keyword not in inspect.signature(bridge).parameters
for method in (
Table.generated_column_status,
LanceTable.generated_column_status,
RemoteTable.generated_column_status,
):
_assert_exact_public_signature(method)
assert not inspect.iscoroutinefunction(method)
assert get_type_hints(method)["return"] == _EXPECTED_RETURN
async_method = AsyncTable.generated_column_status
_assert_exact_public_signature(async_method)
assert inspect.iscoroutinefunction(async_method)
assert get_type_hints(async_method)["return"] == _EXPECTED_RETURN
def test_sync_remote_complete_and_incomplete_one_describe_each():
table, log, cm = _open_remote_table()
try:
complete = table.generated_column_status("complete_col")
_assert_status_string(complete, "complete")
_assert_one_status_describe(log)
log.start()
incomplete = table.generated_column_status("incomplete_col")
_assert_status_string(incomplete, "incomplete")
_assert_one_status_describe(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_remote_complete_and_incomplete_one_describe_each():
table, log, cm = await _open_remote_table_async()
try:
complete = await table.generated_column_status("complete_col")
_assert_status_string(complete, "complete")
_assert_one_status_describe(log)
log.start()
incomplete = await table.generated_column_status("incomplete_col")
_assert_status_string(incomplete, "incomplete")
_assert_one_status_describe(log)
finally:
await cm.__aexit__(None, None, None)
@pytest.mark.parametrize(
("column_name", "status_describe", "expected_exc"),
[
(
"missing",
_describe_body(),
ValueError,
),
(
"Complete_Col",
_describe_body(),
ValueError,
),
(
"ordinary",
_describe_body(),
ValueError,
),
(
"complete_col",
_describe_body(
schema=_status_schema_fields(
complete_meta=_definition_metadata_json(
_COMPLETE_FIELD_ID + 1, 3, 3
)
)
),
ValueError,
),
(
"gen_bad",
_describe_body(
field_ids=[*_STABLE_FIELD_IDS, 9],
schema=_status_schema_fields(
bad_name="gen_bad",
bad_meta=(
'{"format_version":1,"output_field_id":9,'
f'"function_call":{_RAW_METADATA_MARKER},'
'"dependency_epoch":1,"materialized_epoch":1}'
),
),
),
ValueError,
),
(
"complete_col",
_describe_body(
schema=_status_schema_fields(
complete_meta=_definition_metadata_json(
_COMPLETE_FIELD_ID, 1, 1
).replace('"format_version":1', '"format_version":2')
)
),
ValueError,
),
(
"incomplete_col",
_describe_body(
schema=_status_schema_fields(
incomplete_meta=_definition_metadata_json(
_INCOMPLETE_FIELD_ID, 1, 2
)
)
),
ValueError,
),
(
"complete_col",
_describe_body(field_ids=None),
NotImplementedError,
),
],
ids=[
"missing",
"case_mismatch",
"ordinary",
"output_id_mismatch",
"malformed_metadata",
"unknown_format_version",
"reversed_epochs",
"old_server_missing_field_ids",
],
)
def test_remote_fail_closed_matrix_one_describe(
column_name: str,
status_describe: dict[str, Any],
expected_exc: type[BaseException],
):
table, log, cm = _open_remote_table(status_describe=status_describe)
try:
with pytest.raises(expected_exc) as raised:
table.generated_column_status(column_name)
text = _exception_text(raised.value)
assert _RAW_METADATA_MARKER not in text
_assert_one_status_describe(log)
finally:
cm.__exit__(None, None, None)
def test_sync_empty_name_zero_post_open_requests():
table, log, cm = _open_remote_table()
try:
with pytest.raises(ValueError):
table.generated_column_status("")
_assert_no_operation_traffic(log)
finally:
cm.__exit__(None, None, None)
@pytest.mark.asyncio
async def test_async_empty_name_zero_post_open_requests():
table, log, cm = await _open_remote_table_async()
try:
with pytest.raises(ValueError):
await table.generated_column_status("")
_assert_no_operation_traffic(log)
finally:
await cm.__aexit__(None, None, None)
@pytest.mark.asyncio
async def test_async_closed_status_empty_validation_wins_and_nonempty_closed():
"""Publicly closed AsyncTable: empty validates first; nonempty is closed."""
table, log, cm = await _open_remote_table_async()
try:
table.close()
log.start()
try:
await table.generated_column_status("complete_col")
except AttributeError:
raise
except Exception as exc:
text = _exception_text(exc)
assert "closed" in text.lower()
else:
pytest.fail("closed AsyncTable must fail before transport")
_assert_no_operation_traffic(log)
log.start()
with pytest.raises(ValueError) as raised:
await table.generated_column_status("")
text = _exception_text(raised.value)
assert "closed" not in text.lower()
_assert_no_operation_traffic(log)
finally:
await cm.__aexit__(None, None, None)
def test_local_sync_ordinary_column_fails_without_side_effects(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table(
"ordinary_only",
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
)
assert isinstance(table, LanceTable)
version_before = table.version
schema_before = table.schema
data_before = table.to_arrow()
with pytest.raises(ValueError):
table.generated_column_status("ordinary")
assert table.version == version_before
assert table.schema == schema_before
assert table.to_arrow().equals(data_before)
@pytest.mark.asyncio
async def test_local_async_ordinary_column_fails_without_side_effects(tmp_path):
db = await lancedb.connect_async(tmp_path)
table = await db.create_table(
"ordinary_only_async",
[{"ordinary": "alpha", "id": 1}, {"ordinary": "beta", "id": 2}],
)
assert isinstance(table, AsyncTable)
version_before = await table.version()
schema_before = await table.schema()
data_before = await table.to_arrow()
with pytest.raises(ValueError):
await table.generated_column_status("ordinary")
assert await table.version() == version_before
assert await table.schema() == schema_before
assert (await table.to_arrow()).equals(data_before)
-291
View File
@@ -1,291 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for the local @udf declaration surface."""
from __future__ import annotations
import importlib
import inspect
import types
import pyarrow as pa
import pytest
import lancedb
from lancedb import Function, Job, udf
from lancedb._udf import _get_udf_config
_REMOVED_AUTHORING_KNOBS = (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"artifact",
"digest",
"geneva",
)
def _decorate(fn, **overrides):
kwargs = {
"inputs": {"x": pa.int32()},
"output": pa.int64(),
"python": "3.12",
}
kwargs.update(overrides)
return udf(**kwargs)(fn)
def test_udf_top_level_export_and_identity_metadata_behavior():
assert "udf" in lancedb.__all__
assert udf is lancedb.udf
assert isinstance(importlib.import_module("lancedb._udf"), types.ModuleType)
assert not isinstance(lancedb.udf, types.ModuleType)
def add(x, y=1):
"""Add locally."""
return x + y
original = add
decorated = _decorate(
add,
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
)
assert decorated is original
assert decorated.__name__ == "add"
assert decorated.__doc__ == "Add locally."
assert str(inspect.signature(decorated)) == "(x, y=1)"
assert decorated(2) == 3
assert decorated(2, 5) == 7
assert decorated(x=4, y=6) == 10
def test_udf_config_snapshot_order_defaults_and_immutability():
inputs = {"z": pa.string(), "a": pa.int32()}
packages = ["pkg-b==2", "pkg-a==1"]
def combine(z, a):
return f"{z}:{a}"
decorated = udf(
inputs=inputs,
output=pa.string(),
python="3.11",
packages=packages,
output_nullable=False,
)(combine)
config = _get_udf_config(decorated)
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
assert isinstance(config.inputs, tuple)
assert config.output == pa.string()
assert config.output_nullable is False
assert config.python == "3.11"
assert config.packages == ("pkg-b==2", "pkg-a==1")
assert isinstance(config.packages, tuple)
inputs["extra"] = pa.bool_()
del inputs["z"]
packages.append("pkg-c==3")
packages[0] = "mutated==0"
assert config.inputs == (("z", pa.string()), ("a", pa.int32()))
assert config.packages == ("pkg-b==2", "pkg-a==1")
for attr in ("inputs", "output", "output_nullable", "python", "packages"):
with pytest.raises(AttributeError):
setattr(config, attr, None)
def defaults_only(x):
return x
defaulted = udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
)(defaults_only)
default_config = _get_udf_config(defaulted)
assert default_config.packages == ()
assert default_config.output_nullable is True
def test_udf_accepts_lambda_and_closure_for_local_declaration():
ambient = "ambient-secret-value-xyz"
lam = udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda n: n + 1)
assert lam(3) == 4
assert _get_udf_config(lam).inputs == (("n", pa.int32()),)
def factory(offset):
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
packages=["demo==0.1"],
)
def closed(n):
return n + offset + len(ambient)
return closed
closed = factory(10)
assert closed(2) == 12 + len(ambient)
assert _get_udf_config(closed).packages == ("demo==0.1",)
def test_udf_declaration_defers_signature_and_implementation_packaging():
"""Declaration must not validate callable signature or embed implementation."""
def local_add(left, right=1):
return left + right
decorated = udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)(local_add)
assert decorated is local_add
assert str(inspect.signature(decorated)) == "(left, right=1)"
assert decorated(2) == 3
assert decorated(2, 5) == 7
config = _get_udf_config(decorated)
assert config.inputs == (("x", pa.int32()), ("y", pa.int32()))
for attr in (
"source",
"module",
"callable",
"function",
"implementation",
"bundle",
"artifact",
"digest",
):
assert not hasattr(config, attr)
def test_udf_lookup_double_decoration_and_non_function_target():
def plain(x):
return x
with pytest.raises((TypeError, ValueError)):
_get_udf_config(plain)
decorated = _decorate(plain)
with pytest.raises((TypeError, ValueError)):
_decorate(decorated)
with pytest.raises(TypeError):
udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(object())
with pytest.raises(TypeError):
udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(42)
def test_udf_config_validation_errors():
def target(x):
return x
with pytest.raises(TypeError):
udf({"x": pa.int32()}, pa.int32(), "3.12")(target)
with pytest.raises(TypeError):
_decorate(target, inputs=[("x", pa.int32())])
with pytest.raises(TypeError):
_decorate(target, inputs={1: pa.int32()})
with pytest.raises(ValueError):
_decorate(target, inputs={"": pa.int32()})
with pytest.raises(TypeError):
_decorate(target, inputs={"x": "int32"})
with pytest.raises(TypeError):
_decorate(target, output="int64")
with pytest.raises(TypeError):
_decorate(target, python=3.12)
with pytest.raises(ValueError):
_decorate(target, python="")
with pytest.raises(TypeError):
_decorate(target, packages="pkg==1")
with pytest.raises(ValueError):
_decorate(target, packages=["pkg==1", ""])
with pytest.raises(ValueError):
_decorate(target, packages=["pkg==1", "pkg==1"])
with pytest.raises(TypeError):
_decorate(target, packages=["pkg==1", 2])
with pytest.raises(TypeError):
_decorate(target, output_nullable=1)
with pytest.raises(TypeError):
_decorate(target, output_nullable="true")
def test_udf_rejects_removed_overdesign_and_has_no_durable_side_effects():
params = inspect.signature(udf).parameters
for name in _REMOVED_AUTHORING_KNOBS:
assert name not in params
def score(x):
"""score body marker unique-xyz."""
ambient = "ambient-secret-value-xyz"
return f"{ambient}:{x}"
decorated = _decorate(
score,
packages=["score==1.0"],
output_nullable=True,
)
config = _get_udf_config(decorated)
text = repr(config).lower()
assert "score body marker unique-xyz" not in text
assert "ambient-secret-value-xyz" not in text
for token in (
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"on_error",
"functionversion",
"artifact",
"digest",
"geneva",
):
assert token not in text
for attr in _REMOVED_AUTHORING_KNOBS:
assert not hasattr(config, attr)
assert not isinstance(decorated, Function)
assert not isinstance(decorated, Job)
for attr in ("id", "function_id", "job", "job_id", "registration"):
assert not hasattr(decorated, attr)
@@ -1,490 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for local FunctionCapability authoring and @udf capabilities."""
from __future__ import annotations
import inspect
import pyarrow as pa
import pytest
import lancedb
from lancedb import Function, FunctionCapability, Job, udf
from lancedb._udf import _get_udf_config, _package_udf
_SECRET_REFERENCE = "secret://team/capability-redact-token-xyz"
_SECRET_ENV = "API_TOKEN"
_NETWORK_ORIGIN = "https://api.example.com"
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"authorization",
"authorized",
"value",
"plaintext",
"plaintext_secret",
"secret_value",
"job",
"job_id",
"catalog",
"retry_key",
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"geneva",
)
def _decorate(fn, **overrides):
kwargs = {
"inputs": {"x": pa.int32()},
"output": pa.int64(),
"python": "3.12",
}
kwargs.update(overrides)
return udf(**kwargs)(fn)
def _network(origin: str = _NETWORK_ORIGIN) -> FunctionCapability:
return FunctionCapability.network(origin)
def _secret(
reference: str = _SECRET_REFERENCE,
*,
environment_variable: str = _SECRET_ENV,
) -> FunctionCapability:
return FunctionCapability.secret(
reference,
environment_variable=environment_variable,
)
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_without_capabilities(x):
return x + 1
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_with_capabilities(x):
return x + 1
def test_function_capability_export_factories_projection_equality_immutability():
assert "FunctionCapability" in lancedb.__all__
assert FunctionCapability is lancedb.FunctionCapability
network = _network()
secret = _secret()
assert network.kind == "network"
assert network.origin == _NETWORK_ORIGIN
assert network.reference is None
assert network.environment_variable is None
assert secret.kind == "secret"
assert secret.reference == _SECRET_REFERENCE
assert secret.environment_variable == _SECRET_ENV
assert secret.origin is None
assert network == FunctionCapability.network(_NETWORK_ORIGIN)
assert secret == FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
)
assert network != secret
assert network != FunctionCapability.network("https://other.example.com")
assert secret != FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable="OTHER_TOKEN",
)
public_attrs = ("kind", "origin", "reference", "environment_variable")
internal_slots = ("_kind", "_origin", "_reference", "_environment_variable")
immutable_attrs = public_attrs + internal_slots
for attr in public_attrs:
with pytest.raises(AttributeError):
setattr(network, attr, None)
with pytest.raises(AttributeError):
setattr(secret, attr, None)
for attr in immutable_attrs:
# Fresh instances per attempt so a RED slot mutation cannot corrupt
# shared fixtures used by later assertions in this test.
fresh_network = _network("https://fresh-immutability.example.com")
fresh_secret = _secret(
"secret://team/fresh-immutability-token",
environment_variable="FRESH_IMMUTABILITY_TOKEN",
)
with pytest.raises(AttributeError):
setattr(fresh_network, attr, None)
with pytest.raises(AttributeError):
setattr(fresh_secret, attr, None)
with pytest.raises(AttributeError):
delattr(fresh_network, attr)
with pytest.raises(AttributeError):
delattr(fresh_secret, attr)
retained_origin = "https://config-retain.example.com"
retained_reference = "secret://team/config-retain-token"
retained_env = "CONFIG_RETAIN_TOKEN"
retained_network = FunctionCapability.network(retained_origin)
retained_secret = FunctionCapability.secret(
retained_reference,
environment_variable=retained_env,
)
expected_capabilities = (
FunctionCapability.network(retained_origin),
FunctionCapability.secret(
retained_reference,
environment_variable=retained_env,
),
)
def retain_target(x):
return x
retained = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
capabilities=[retained_network, retained_secret],
)(retain_target)
retained_config = _get_udf_config(retained)
assert retained_config.capabilities == expected_capabilities
for attr in immutable_attrs:
with pytest.raises(AttributeError):
setattr(retained_network, attr, "mutated")
with pytest.raises(AttributeError):
setattr(retained_secret, attr, "mutated")
with pytest.raises(AttributeError):
delattr(retained_network, attr)
with pytest.raises(AttributeError):
delattr(retained_secret, attr)
assert retained_config.capabilities == expected_capabilities
assert retained_config.capabilities[0] is retained_network
assert retained_config.capabilities[1] is retained_secret
assert retained_config.capabilities[0].kind == "network"
assert retained_config.capabilities[0].origin == retained_origin
assert retained_config.capabilities[0].reference is None
assert retained_config.capabilities[0].environment_variable is None
assert retained_config.capabilities[1].kind == "secret"
assert retained_config.capabilities[1].reference == retained_reference
assert retained_config.capabilities[1].environment_variable == retained_env
assert retained_config.capabilities[1].origin is None
with pytest.raises(TypeError):
FunctionCapability()
with pytest.raises(TypeError):
FunctionCapability( # type: ignore[call-arg]
kind="network",
origin=_NETWORK_ORIGIN,
)
assert not isinstance(network, Function)
assert not isinstance(secret, Function)
assert not isinstance(network, Job)
assert not isinstance(secret, Job)
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(network, attr)
assert not hasattr(secret, attr)
def test_function_capability_validation_and_secret_redaction():
with pytest.raises(TypeError):
FunctionCapability.network(None) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.network(123) # type: ignore[arg-type]
with pytest.raises(ValueError):
FunctionCapability.network("")
# Backend authorization owns URL/scheme policy; non-empty is enough here.
loose = FunctionCapability.network("example.com")
assert loose.kind == "network"
assert loose.origin == "example.com"
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[misc]
_SECRET_REFERENCE,
_SECRET_ENV,
)
with pytest.raises(TypeError):
FunctionCapability.secret(None, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(123, environment_variable=_SECRET_ENV) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=None) # type: ignore[arg-type]
with pytest.raises(TypeError):
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable=1) # type: ignore[arg-type]
with pytest.raises(ValueError) as empty_ref:
FunctionCapability.secret("", environment_variable=_SECRET_ENV)
assert _SECRET_REFERENCE not in str(empty_ref.value)
assert _SECRET_REFERENCE not in repr(empty_ref.value)
with pytest.raises(ValueError) as empty_env:
FunctionCapability.secret(_SECRET_REFERENCE, environment_variable="")
assert _SECRET_REFERENCE not in str(empty_env.value)
assert _SECRET_REFERENCE not in repr(empty_env.value)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
value="super-secret",
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
plaintext_secret="super-secret",
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
environment={_SECRET_ENV: "super-secret"},
)
with pytest.raises(TypeError):
FunctionCapability.secret( # type: ignore[call-arg]
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
headers={"Authorization": "Bearer super-secret"},
)
with pytest.raises(TypeError):
FunctionCapability.network( # type: ignore[call-arg]
_NETWORK_ORIGIN,
headers={"X-Trace": "1"},
)
secret = _secret()
assert not hasattr(secret, "value")
assert not hasattr(secret, "plaintext")
assert not hasattr(secret, "plaintext_secret")
assert not hasattr(secret, "secret_value")
secret_text = repr(secret)
assert "secret" in secret_text.lower()
assert _SECRET_ENV in secret_text
assert _SECRET_REFERENCE not in secret_text
assert "super-secret" not in secret_text
network_text = repr(_network())
assert "network" in network_text.lower()
assert _NETWORK_ORIGIN in network_text
def test_udf_capabilities_ordered_immutable_config_default_and_validation():
params = inspect.signature(udf).parameters
assert "capabilities" in params
assert params["capabilities"].kind is inspect.Parameter.KEYWORD_ONLY
assert params["capabilities"].default == ()
def identity_target(x):
"""capabilities identity marker."""
return x + 1
original = identity_target
decorated = _decorate(identity_target)
assert decorated is original
assert decorated.__name__ == "identity_target"
assert decorated.__doc__ == "capabilities identity marker."
assert decorated(2) == 3
assert _get_udf_config(decorated).capabilities == ()
first = _network("https://b.example.com")
second = _network("https://a.example.com")
third = _network("https://b.example.com")
secret = _secret()
capabilities = [first, second, third, secret]
def combine(x):
return x
with_caps = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
packages=["pkg-b==2", "pkg-a==1"],
capabilities=capabilities,
)(combine)
config = _get_udf_config(with_caps)
assert config.capabilities == (first, second, third, secret)
assert isinstance(config.capabilities, tuple)
assert config.packages == ("pkg-b==2", "pkg-a==1")
assert config.inputs == (("x", pa.int32()),)
capabilities.append(_network("https://mutated.example.com"))
capabilities[0] = _network("https://replaced.example.com")
assert config.capabilities == (first, second, third, secret)
with pytest.raises(AttributeError):
setattr(config, "capabilities", ())
def target(x):
return x
with pytest.raises(TypeError):
_decorate(target, capabilities="https://api.example.com")
with pytest.raises(TypeError):
_decorate(target, capabilities=b"https://api.example.com")
class _BadCapability:
def __repr__(self) -> str:
return "unique-bad-capability-repr-xyz"
with pytest.raises(TypeError) as bad_item:
_decorate(target, capabilities=[_BadCapability()])
assert "unique-bad-capability-repr-xyz" not in str(bad_item.value)
assert "unique-bad-capability-repr-xyz" not in repr(bad_item.value)
with pytest.raises(TypeError) as bad_mixed:
_decorate(
target,
capabilities=[_network(), "unique-bad-capability-string-xyz"],
)
assert "unique-bad-capability-string-xyz" not in str(bad_mixed.value)
assert "unique-bad-capability-string-xyz" not in repr(bad_mixed.value)
def test_udf_capabilities_rejects_function_capability_subclass_before_property_access():
marker = "unique-hostile-capability-subclass-marker-xyz"
class _HostileFunctionCapability(FunctionCapability):
@property
def kind(self) -> str:
raise RuntimeError(marker)
@property
def origin(self) -> str | None:
raise RuntimeError(marker)
@property
def reference(self) -> str | None:
raise RuntimeError(marker)
@property
def environment_variable(self) -> str | None:
raise RuntimeError(marker)
hostile = object.__new__(_HostileFunctionCapability)
assert isinstance(hostile, FunctionCapability)
assert type(hostile) is not FunctionCapability
def target(x):
return x
with pytest.raises(TypeError) as exc_info:
_decorate(target, capabilities=[hostile])
assert marker not in str(exc_info.value)
assert marker not in repr(exc_info.value)
assert _SECRET_REFERENCE not in str(exc_info.value)
assert _SECRET_REFERENCE not in repr(exc_info.value)
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
def test_package_udf_preserves_capabilities_and_redacts_secret_reference():
packaged = _package_udf(packable_with_capabilities)
config = packaged.config
assert packaged.config is _get_udf_config(packable_with_capabilities)
assert config.capabilities == (
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
)
assert config.capabilities[0].kind == "network"
assert config.capabilities[0].origin == _NETWORK_ORIGIN
assert config.capabilities[1].kind == "secret"
assert config.capabilities[1].reference == _SECRET_REFERENCE
assert config.capabilities[1].environment_variable == _SECRET_ENV
assert config.packages == ("pkg-a==1",)
assert config.python == "3.12"
assert config.output_nullable is False
nested = (
f"{packaged!r}\n{config!r}\n{config.capabilities!r}\n{config.capabilities[1]!r}"
)
assert _SECRET_REFERENCE not in nested
assert _SECRET_ENV in repr(config.capabilities[1])
def test_capabilities_are_additive_to_existing_declaration_and_packaging():
def score(x):
return x
decorated = _decorate(
score,
packages=["score==1.0"],
output_nullable=True,
)
config = _get_udf_config(decorated)
assert config.inputs == (("x", pa.int32()),)
assert config.output == pa.int64()
assert config.output_nullable is True
assert config.python == "3.12"
assert config.packages == ("score==1.0",)
assert config.capabilities == ()
assert decorated is score
assert decorated(4) == 4
packaged = _package_udf(packable_without_capabilities)
assert packaged.config is _get_udf_config(packable_without_capabilities)
assert packaged.callable_name == "packable_without_capabilities"
assert packaged.config.capabilities == ()
assert packaged.config.packages == ("pkg-a==1",)
assert packaged.config.output_nullable is False
assert packable_without_capabilities(1) == 2
params = inspect.signature(udf).parameters
for name in (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"artifact",
"digest",
"geneva",
):
assert name not in params
@@ -1,506 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for the private UDF -> FunctionDefinition bridge."""
from __future__ import annotations
import base64
import io
import json
from pathlib import Path
import pyarrow as pa
import pytest
import lancedb
from lancedb import FunctionCapability, udf
from lancedb import _lancedb as _native
from lancedb import _udf as _udf_mod
_SOURCE_MARKER = "bridge-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/bridge-redact-token-xyz"
_SECRET_ENV = "BRIDGE_API_TOKEN"
_NETWORK_ORIGIN = "https://api.bridge-example.com"
_NETWORK_ORIGIN_B = "https://other.bridge-example.com"
_FORBIDDEN_WIRE_KEYS = (
"id",
"function_id",
"FunctionId",
"catalog",
"catalog_name",
"version",
"function_version",
"FunctionVersion",
"lineage",
"user_version",
"idempotency_key",
"digest",
"artifact",
"artifact_digest",
"storage",
"storage_location",
"location",
"deterministic",
"null_policy",
"nullPolicy",
"timestamp",
"created_at",
"updated_at",
"worker",
"scheduler",
"attempt",
"attempt_id",
"replica",
"placement",
"job",
"job_id",
"retry_key",
"registration",
)
_OVERDESIGN_ATTRS = (
"id",
"function_id",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"job",
"job_id",
"catalog",
"retry_key",
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"null_handling",
)
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-b==2", "pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
FunctionCapability.network(_NETWORK_ORIGIN_B),
],
)
def packable_bridge_normalize(text, limit):
"""bridge-source-marker-unique-xyz."""
return text[:limit]
def _build_function_definition(fn: object):
return _udf_mod._build_function_definition(fn)
def _function_definition_type():
return _native._FunctionDefinition
def _new_function_definition(**kwargs):
return _native._new_function_definition(**kwargs)
def _json_bytes(definition) -> bytes:
payload = definition._to_json()
if isinstance(payload, bytes):
return payload
assert isinstance(payload, str)
return payload.encode("utf-8")
def _decode_type_ipc(encoded: str) -> pa.DataType:
raw = base64.b64decode(encoded)
reader = pa.ipc.open_file(io.BytesIO(raw))
assert reader.num_record_batches == 0
assert len(reader.schema) == 1
return reader.schema.field(0).type
def _assert_exact_object_keys(value: dict, expected: set[str], *, context: str) -> None:
assert isinstance(value, dict), f"{context} must be an object"
assert set(value) == expected, f"{context} keys must match exactly: {set(value)!r}"
def _assert_forbidden_keys_absent(value: object, *, context: str) -> None:
if isinstance(value, dict):
for key in value:
assert key not in _FORBIDDEN_WIRE_KEYS, (
f"forbidden key {key!r} at {context}: {value!r}"
)
if key == "name" and context in {
"definition",
"signature",
"signature.output",
"implementation",
}:
raise AssertionError(
f"catalog/function identity key `name` must be absent at {context}"
)
child_context = f"{context}.{key}"
if key == "parameters" and context == "signature":
child_context = "signature.parameters"
_assert_forbidden_keys_absent(value[key], context=child_context)
elif isinstance(value, list):
for idx, item in enumerate(value):
item_context = (
f"signature.parameters[{idx}]"
if context == "signature.parameters"
else f"{context}[{idx}]"
)
if context == "signature.parameters":
assert isinstance(item, dict)
assert "name" in item
for key in item:
assert key not in _FORBIDDEN_WIRE_KEYS
assert key != "catalog_name"
_assert_forbidden_keys_absent(
{k: v for k, v in item.items() if k != "name"},
context=item_context,
)
else:
_assert_forbidden_keys_absent(item, context=item_context)
def _assert_sanitized_text(*parts: object) -> None:
combined = "\n".join(str(part) for part in parts)
lowered = combined.lower()
assert _SOURCE_MARKER.lower() not in lowered
assert _SECRET_REFERENCE.lower() not in lowered
assert str(Path(__file__).resolve()).lower() not in lowered
assert Path(__file__).resolve().as_posix().lower() not in lowered
def _assert_clean_validation_error(exc_info) -> None:
_assert_sanitized_text(exc_info.value, repr(exc_info.value))
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
def _valid_builder_kwargs(**overrides):
kwargs = {
"parameters": [("text", pa.string()), ("limit", pa.int32())],
"output_type": pa.string(),
"output_nullable": True,
"module": "bridge_mod",
"callable_name": "normalize",
"source": (
"def normalize(text, limit):\n"
f" # {_SOURCE_MARKER}\n"
" return text[:limit]\n"
),
"python": "3.12",
"packages": ["pkg-b==2", "pkg-a==1"],
"capabilities": [
("network", _NETWORK_ORIGIN, None),
("secret", _SECRET_REFERENCE, _SECRET_ENV),
("network", _NETWORK_ORIGIN_B, None),
],
}
kwargs.update(overrides)
return kwargs
def test_build_function_definition_private_native_immutability_and_export_surface():
assert "_build_function_definition" not in getattr(lancedb, "__all__", [])
assert "_FunctionDefinition" not in lancedb.__all__
assert not hasattr(lancedb, "_FunctionDefinition")
assert not hasattr(lancedb, "_build_function_definition")
assert not hasattr(lancedb, "_new_function_definition")
definition = _build_function_definition(packable_bridge_normalize)
definition_type = _function_definition_type()
assert type(definition) is definition_type
assert definition_type.__module__ == "lancedb._lancedb"
assert definition_type.__name__ == "_FunctionDefinition"
with pytest.raises(TypeError):
definition_type()
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(definition, attr)
for attr in ("signature", "module", "source", "capabilities"):
with pytest.raises(AttributeError):
setattr(definition, attr, None)
def test_build_function_definition_json_wire_ordered_contract_without_identity():
definition = _build_function_definition(packable_bridge_normalize)
encoded_a = _json_bytes(definition)
encoded_b = _json_bytes(definition)
assert encoded_a == encoded_b
wire = json.loads(encoded_a.decode("utf-8"))
_assert_exact_object_keys(
wire,
{"format_version", "signature", "implementation", "capabilities"},
context="definition",
)
assert wire["format_version"] == 1
_assert_forbidden_keys_absent(wire, context="definition")
signature = wire["signature"]
_assert_exact_object_keys(signature, {"parameters", "output"}, context="signature")
parameters = signature["parameters"]
assert [parameter["name"] for parameter in parameters] == ["text", "limit"]
for parameter in parameters:
_assert_exact_object_keys(
parameter, {"name", "data_type_ipc"}, context="parameter"
)
assert isinstance(parameter["data_type_ipc"], str)
assert parameter["data_type_ipc"]
assert _decode_type_ipc(parameters[0]["data_type_ipc"]) == pa.string()
assert _decode_type_ipc(parameters[1]["data_type_ipc"]) == pa.int32()
output = signature["output"]
_assert_exact_object_keys(
output, {"data_type_ipc", "nullable"}, context="signature.output"
)
assert output["nullable"] is True
assert _decode_type_ipc(output["data_type_ipc"]) == pa.string()
implementation = wire["implementation"]
_assert_exact_object_keys(
implementation,
{"kind", "module", "callable", "source", "python", "packages"},
context="implementation",
)
assert implementation["kind"] == "python"
assert implementation["module"] == __name__
assert implementation["callable"] == "packable_bridge_normalize"
assert implementation["source"] == Path(__file__).read_text(encoding="utf-8")
assert _SOURCE_MARKER in implementation["source"]
assert implementation["python"] == "3.12"
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
capabilities = wire["capabilities"]
assert capabilities == [
{"kind": "network", "origin": _NETWORK_ORIGIN},
{
"kind": "secret",
"reference": _SECRET_REFERENCE,
"environment_variable": _SECRET_ENV,
},
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
]
for capability in capabilities:
assert "value" not in capability
assert "plaintext" not in capability
assert "plaintext_secret" not in capability
assert "secret_value" not in capability
def test_native_definition_repr_includes_safe_structure_and_redacts_sensitive_text():
definition = _build_function_definition(packable_bridge_normalize)
rendered = repr(definition)
assert "_FunctionDefinition" in rendered or "FunctionDefinition" in rendered
assert __name__ in rendered
assert "packable_bridge_normalize" in rendered
assert "3.12" in rendered
_assert_sanitized_text(rendered)
def test_new_function_definition_builder_preserves_normalized_wire():
definition = _new_function_definition(**_valid_builder_kwargs())
assert type(definition) is _function_definition_type()
encoded_a = _json_bytes(definition)
encoded_b = _json_bytes(definition)
assert encoded_a == encoded_b
wire = json.loads(encoded_a.decode("utf-8"))
assert wire["format_version"] == 1
assert [parameter["name"] for parameter in wire["signature"]["parameters"]] == [
"text",
"limit",
]
assert _decode_type_ipc(wire["signature"]["parameters"][0]["data_type_ipc"]) == (
pa.string()
)
assert _decode_type_ipc(wire["signature"]["parameters"][1]["data_type_ipc"]) == (
pa.int32()
)
assert wire["signature"]["output"]["nullable"] is True
assert _decode_type_ipc(wire["signature"]["output"]["data_type_ipc"]) == pa.string()
implementation = wire["implementation"]
assert implementation["kind"] == "python"
assert implementation["module"] == "bridge_mod"
assert implementation["callable"] == "normalize"
assert implementation["source"] == _valid_builder_kwargs()["source"]
assert _SOURCE_MARKER in implementation["source"]
assert implementation["python"] == "3.12"
assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"]
assert wire["capabilities"] == [
{"kind": "network", "origin": _NETWORK_ORIGIN},
{
"kind": "secret",
"reference": _SECRET_REFERENCE,
"environment_variable": _SECRET_ENV,
},
{"kind": "network", "origin": _NETWORK_ORIGIN_B},
]
_assert_forbidden_keys_absent(wire, context="definition")
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": [("text", pa.string()), ("text", pa.int32())]},),
({"parameters": [("", pa.string())]},),
({"module": ""},),
({"callable_name": ""},),
({"source": ""},),
({"python": ""},),
({"packages": ["pkg-a==1", ""]},),
({"packages": ["pkg-a==1", "pkg-a==1"]},),
({"capabilities": [("filesystem", _NETWORK_ORIGIN, None)]},),
({"capabilities": [("network", _NETWORK_ORIGIN, _SECRET_ENV)]},),
({"capabilities": [("secret", _SECRET_REFERENCE, None)]},),
({"capabilities": [("secret", _SECRET_REFERENCE, "")]},),
({"capabilities": [("network", "", None)]},),
({"capabilities": [("secret", "", _SECRET_ENV)]},),
],
)
def test_new_function_definition_strict_validation_rejections(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
def test_new_function_definition_validation_does_not_echo_secret_or_source_marker():
with pytest.raises(ValueError) as exc_info:
_new_function_definition(**_valid_builder_kwargs(module=""))
_assert_clean_validation_error(exc_info)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(
**_valid_builder_kwargs(packages=["pkg-a==1", "pkg-a==1"])
)
_assert_clean_validation_error(exc_info)
with pytest.raises(ValueError) as exc_info:
_new_function_definition(
**_valid_builder_kwargs(
capabilities=[("secret", _SECRET_REFERENCE, None)],
)
)
_assert_clean_validation_error(exc_info)
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": [("text", "not-a-datatype")]},),
({"parameters": [(123, pa.string())]},),
({"output_type": "not-a-datatype"},),
({"output_type": None},),
({"output_nullable": "yes"},),
({"packages": "pkg-a==1"},),
({"capabilities": "network"},),
({"capabilities": [("network", _NETWORK_ORIGIN)]},),
({"capabilities": [("network", _NETWORK_ORIGIN, None, "extra")]},),
],
)
def test_new_function_definition_wrong_pyarrow_and_shape_values_fail_closed(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises((TypeError, ValueError)) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
class _HostileRaisingIterable:
def __iter__(self):
raise RuntimeError(f"{_SECRET_REFERENCE} {_SOURCE_MARKER}")
@pytest.mark.parametrize(
("overrides",),
[
({"parameters": _HostileRaisingIterable()},),
({"packages": _HostileRaisingIterable()},),
({"capabilities": _HostileRaisingIterable()},),
(
{
"capabilities": [
("network", _NETWORK_ORIGIN, None),
_HostileRaisingIterable(),
("network", _NETWORK_ORIGIN_B, None),
]
},
),
],
)
def test_new_function_definition_hostile_iterable_iter_raises_fail_closed(overrides):
kwargs = _valid_builder_kwargs(**overrides)
with pytest.raises((TypeError, ValueError)) as exc_info:
_new_function_definition(**kwargs)
_assert_clean_validation_error(exc_info)
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_bridge_capability_exact_type(x):
return x + 1
def test_build_function_definition_rejects_forged_function_capability_subclass():
marker = f"{_SECRET_REFERENCE} {_SOURCE_MARKER}"
class _HostileFunctionCapability(FunctionCapability):
@property
def kind(self) -> str:
raise RuntimeError(marker)
@property
def origin(self) -> str | None:
raise RuntimeError(marker)
@property
def reference(self) -> str | None:
raise RuntimeError(marker)
@property
def environment_variable(self) -> str | None:
raise RuntimeError(marker)
hostile = object.__new__(_HostileFunctionCapability)
assert isinstance(hostile, FunctionCapability)
assert type(hostile) is not FunctionCapability
config_attr = _udf_mod._CONFIG_ATTR
original = getattr(packable_bridge_capability_exact_type, config_attr)
forged = _udf_mod._UdfConfig(
inputs=original.inputs,
output=original.output,
output_nullable=original.output_nullable,
python=original.python,
packages=original.packages,
capabilities=(hostile,),
)
setattr(packable_bridge_capability_exact_type, config_attr, forged)
try:
with pytest.raises((TypeError, ValueError)) as exc_info:
_build_function_definition(packable_bridge_capability_exact_type)
_assert_clean_validation_error(exc_info)
assert marker not in str(exc_info.value)
assert marker not in repr(exc_info.value)
finally:
setattr(packable_bridge_capability_exact_type, config_attr, original)
@@ -1,486 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""RED contract tests for private UDF packaging validation."""
from __future__ import annotations
import importlib
import inspect
import json
import sys
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
import pyarrow as pa
import pytest
from lancedb import Function, Job, udf
from lancedb._udf import _get_udf_config, _package_udf
_BODY_MARKER = "packaging body marker unique-xyz"
_AMBIENT_SECRET = "ambient-secret-value-xyz"
_BUILTIN_SHADOW_SECRET = "builtin-shadow-secret-xyz"
_SOURCE_MISMATCH_SECRET = "source-mismatch-secret-xyz"
_INVALID_UTF8_SECRET = "invalid-utf8-secret-xyz"
_OVERDESIGN_ATTRS = (
"user_version",
"idempotency_key",
"deterministic",
"null_handling",
"null_policy",
"on_error",
"error_policy",
"FunctionVersion",
"function_version",
"artifact",
"digest",
"geneva",
"id",
"function_id",
"job",
"job_id",
"registration",
"catalog",
"retry_key",
"source_path",
"path",
"function",
)
_PACKAGING_CONSTANT = 41
def _packaging_helper(value: int) -> int:
return value + _PACKAGING_CONSTANT
@udf(
inputs={"x": pa.int32()},
output=pa.int64(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=False,
)
def packable_add(x):
"""packaging body marker unique-xyz."""
return _packaging_helper(x) + len(json.dumps({"k": 1}))
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def packable_kwonly(x, *, y=2):
return x + y
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def packable_rebind_target(x):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def uses_injected_ambient(x):
return x + len(INJECTED_AMBIENT_GLOBAL) # noqa: F821
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def uses_shadowed_builtin_len(x):
return x + len((1, 2, 3))
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def mismatch_names(left, right):
return left + right
@udf(
inputs={"y": pa.int32(), "x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def mismatch_order(x, y):
return x + y
@udf(
inputs={"x": pa.int32(), "y": pa.int32()},
output=pa.int32(),
python="3.12",
)
def positional_only(x, /, y):
return x + y
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def varargs_fn(x, *args):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def kwargs_fn(x, **kwargs):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
async def async_fn(x):
return x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
async def async_gen_fn(x):
yield x
@udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)
def generator_fn(x):
yield x
def _assert_sanitized_text(*parts: object, secret: str = _AMBIENT_SECRET) -> None:
combined = "\n".join(str(part) for part in parts)
lowered = combined.lower()
assert _BODY_MARKER.lower() not in lowered
assert secret.lower() not in lowered
assert str(Path(__file__).resolve()).lower() not in lowered
assert Path(__file__).resolve().as_posix().lower() not in lowered
def _assert_packaging_rejection(exc_info, *, secret: str = _AMBIENT_SECRET) -> None:
_assert_sanitized_text(exc_info.value, repr(exc_info.value), secret=secret)
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
@contextmanager
def _temporary_imported_module(
directory: Path, module_name: str, source: str
) -> Iterator[tuple[Path, object]]:
path = directory / f"{module_name}.py"
path.write_text(source, encoding="utf-8")
inserted = str(directory)
sys.path.insert(0, inserted)
try:
sys.modules.pop(module_name, None)
module = importlib.import_module(module_name)
yield path, module
finally:
sys.modules.pop(module_name, None)
try:
sys.path.remove(inserted)
except ValueError:
pass
def _temp_udf_module_source(*, body: str, secret: str | None = None) -> str:
secret_line = f"_SECRET = {secret!r}\n" if secret is not None else ""
return (
"import pyarrow as pa\n"
"from lancedb import udf\n"
f"{secret_line}\n"
"@udf(\n"
' inputs={"x": pa.int32()},\n'
" output=pa.int32(),\n"
' python="3.12",\n'
")\n"
"def temp_pack_target(x):\n"
f" {body}\n"
)
def test_package_udf_success_snapshot_source_module_callable_config_and_repr():
packaged = _package_udf(packable_add)
source = Path(__file__).read_text(encoding="utf-8")
assert packaged.source == source
assert packaged.module == __name__
assert packaged.module != "__main__"
assert packaged.callable_name == "packable_add"
assert packable_add.__qualname__ == "packable_add"
assert packaged.config is _get_udf_config(packable_add)
assert packaged.config.inputs == (("x", pa.int32()),)
assert packaged.config.output == pa.int64()
assert packaged.config.output_nullable is False
assert packaged.config.python == "3.12"
assert packaged.config.packages == ("pkg-a==1",)
for attr in ("source", "module", "callable_name", "config"):
with pytest.raises(AttributeError):
setattr(packaged, attr, None)
text = repr(packaged)
_assert_sanitized_text(text)
assert _BODY_MARKER not in text
def test_package_udf_allows_source_bound_import_constant_and_helper():
packaged = _package_udf(packable_add)
assert packaged.callable_name == "packable_add"
assert "import json" in packaged.source
assert "_PACKAGING_CONSTANT" in packaged.source
assert "_packaging_helper" in packaged.source
assert packable_add(1) == _packaging_helper(1) + len(json.dumps({"k": 1}))
def test_package_udf_accepts_positional_or_keyword_and_keyword_only_defaults():
packaged = _package_udf(packable_kwonly)
assert packaged.callable_name == "packable_kwonly"
assert packaged.config.inputs == (("x", pa.int32()), ("y", pa.int32()))
assert str(inspect.signature(packable_kwonly)) == "(x, *, y=2)"
assert packable_kwonly(3) == 5
assert packable_kwonly(3, y=7) == 10
def test_package_udf_rejects_lambda_and_closure():
lam = udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda n: n + 1)
with pytest.raises(ValueError) as exc_info:
_package_udf(lam)
_assert_packaging_rejection(exc_info)
ambient = _AMBIENT_SECRET
def factory(offset):
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)
def closed(n):
return n + offset + len(ambient)
return closed
closed = factory(10)
with pytest.raises(ValueError) as exc_info:
_package_udf(closed)
_assert_packaging_rejection(exc_info)
def outer():
total = 0
@udf(
inputs={"n": pa.int32()},
output=pa.int32(),
python="3.12",
)
def nested(n):
nonlocal total
total += n
return total
return nested
with pytest.raises(ValueError) as exc_info:
_package_udf(outer())
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_signature_mismatches_and_unsupported_parameter_kinds():
for target in (
mismatch_names,
mismatch_order,
positional_only,
varargs_fn,
kwargs_fn,
):
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_async_and_generator_functions():
for target in (async_fn, async_gen_fn, generator_fn):
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_dynamic_exec_source():
namespace: dict[str, object] = {}
exec(
"def dynamic_pack_target(x):\n return x + 1\n",
namespace,
)
dynamic = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(namespace["dynamic_pack_target"])
with pytest.raises(ValueError) as exc_info:
_package_udf(dynamic)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_undecorated_and_wrong_input_types():
def plain(x):
return x
with pytest.raises(TypeError) as exc_info:
_package_udf(plain)
_assert_packaging_rejection(exc_info)
with pytest.raises(TypeError) as exc_info:
_package_udf(object())
_assert_packaging_rejection(exc_info)
with pytest.raises(TypeError) as exc_info:
_package_udf(42)
_assert_packaging_rejection(exc_info)
def test_package_udf_rejects_rebound_module_attribute():
module = sys.modules[__name__]
original = module.packable_rebind_target
replacement = udf(
inputs={"x": pa.int32()},
output=pa.int32(),
python="3.12",
)(lambda x: x)
module.packable_rebind_target = replacement
try:
with pytest.raises(ValueError) as exc_info:
_package_udf(original)
_assert_packaging_rejection(exc_info)
finally:
module.packable_rebind_target = original
def test_package_udf_rejects_injected_ambient_global():
module = sys.modules[__name__]
secret = _AMBIENT_SECRET
module.INJECTED_AMBIENT_GLOBAL = secret
try:
assert uses_injected_ambient(3) == 3 + len(secret)
with pytest.raises(ValueError) as exc_info:
_package_udf(uses_injected_ambient)
_assert_packaging_rejection(exc_info, secret=secret)
finally:
delattr(module, "INJECTED_AMBIENT_GLOBAL")
def test_package_udf_rejects_builtin_shadow_injection():
module = sys.modules[__name__]
secret = _BUILTIN_SHADOW_SECRET
assert not hasattr(module, "len")
module.len = secret
try:
with pytest.raises(ValueError) as exc_info:
_package_udf(uses_shadowed_builtin_len)
_assert_packaging_rejection(exc_info, secret=secret)
finally:
delattr(module, "len")
def test_package_udf_rejects_loaded_code_source_mismatch(tmp_path: Path):
secret = _SOURCE_MISMATCH_SECRET
module_name = "udf_pkg_source_mismatch_mod"
original = _temp_udf_module_source(body="return x + 1")
replacement = _temp_udf_module_source(
body=f"return x + 99 # {secret}",
secret=secret,
)
with _temporary_imported_module(tmp_path, module_name, original) as (
path,
module,
):
target = module.temp_pack_target
assert target(1) == 2
path.write_text(replacement, encoding="utf-8")
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
_assert_packaging_rejection(exc_info, secret=secret)
err_text = f"{exc_info.value}\n{exc_info.value!r}"
assert str(path.resolve()) not in err_text
assert path.resolve().as_posix() not in err_text
def test_package_udf_rejects_invalid_utf8_after_import(tmp_path: Path):
secret = _INVALID_UTF8_SECRET
module_name = "udf_pkg_invalid_utf8_mod"
original = _temp_udf_module_source(body="return x + 1")
with _temporary_imported_module(tmp_path, module_name, original) as (
path,
module,
):
target = module.temp_pack_target
assert target(1) == 2
path.write_bytes(secret.encode("utf-8") + b"\xff\xfe invalid-bytes")
with pytest.raises(ValueError) as exc_info:
_package_udf(target)
assert type(exc_info.value) is ValueError
assert exc_info.value.__cause__ is None
assert exc_info.value.__context__ is None
err_text = f"{exc_info.value}\n{exc_info.value!r}"
assert secret not in err_text
assert "b'" not in err_text
assert r"\xff" not in err_text
assert str(path.resolve()) not in err_text
assert path.resolve().as_posix() not in err_text
def test_package_udf_snapshot_has_no_durable_overdesign_and_is_not_function_or_job():
packaged = _package_udf(packable_add)
assert not isinstance(packaged, Function)
assert not isinstance(packaged, Job)
for attr in _OVERDESIGN_ATTRS:
assert not hasattr(packaged, attr)
text = repr(packaged).lower()
for token in (
"user_version",
"idempotency_key",
"deterministic",
"null_policy",
"on_error",
"functionversion",
"artifact",
"digest",
"geneva",
"retry_key",
):
assert token not in text
_assert_sanitized_text(text)
+1 -94
View File
@@ -12,7 +12,7 @@ import pyarrow.compute as pc
import pytest
import pytest_asyncio
from lancedb.index import BTree, FTS, IvfPq
from lancedb.index import FTS
from lancedb.table import AsyncTable, Table
@@ -99,86 +99,6 @@ async def test_async_hybrid_query_filters(table: AsyncTable):
assert result["text"].to_pylist() == ["cat", "b"]
@pytest.mark.asyncio
async def test_hybrid_query_with_stale_fixed_size_binary_prefilter(
tmpdir_factory,
):
tmp_path = str(tmpdir_factory.mktemp("stale_scalar_prefilter"))
db = await lancedb.connect_async(tmp_path)
def fixed_size_binary(value: int) -> bytes:
return value.to_bytes(16, byteorder="big")
num_rows = 1000
data = pa.table(
{
"space_id": pa.array(
[fixed_size_binary(i) for i in range(num_rows)],
type=pa.binary(16),
),
"text": ["book"] * num_rows,
"vector": pa.array(
[[float(i), float(i)] for i in range(num_rows)],
type=pa.list_(pa.float32(), 2),
),
}
)
table = await db.create_table("test", data)
await table.create_index(
"vector", config=IvfPq(num_partitions=4, num_sub_vectors=2)
)
await table.create_index("space_id", config=BTree())
await table.create_index("text", config=FTS(with_position=False))
# Advance the search indices without advancing the scalar index. This is the
# state that previously let hybrid search use an incomplete scalar prefilter.
await table.add(data)
lance_dataset = await table.to_lance()
lance_dataset.optimize.optimize_indices(index_names=["vector_idx", "text_idx"])
await table.checkout_latest()
scalar_stats = await table.index_stats("space_id_idx")
assert scalar_stats is not None
assert scalar_stats.num_indexed_rows == num_rows
assert scalar_stats.num_unindexed_rows == num_rows
for index_name in ["vector_idx", "text_idx"]:
search_stats = await table.index_stats(index_name)
assert search_stats is not None
assert search_stats.num_indexed_rows == num_rows * 2
assert search_stats.num_unindexed_rows == 0
matching_ids = [5, 10, 15, 20, 25, 30]
literals = [
f"arrow_cast(0x{fixed_size_binary(i).hex()}, 'FixedSizeBinary(16)')"
for i in matching_ids
]
predicate = f"space_id IN ({', '.join(literals)})"
expected_ids = sorted(fixed_size_binary(i) for i in matching_ids for _ in range(2))
vector_query = (
table.query().where(predicate).nearest_to([5.0, 5.0]).limit(num_rows * 2)
)
vector_results = await vector_query.to_arrow()
assert sorted(vector_results["space_id"].to_pylist()) == expected_ids
fts_query = (
table.query().where(predicate).nearest_to_text("book").limit(num_rows * 2)
)
fts_results = await fts_query.to_arrow()
assert sorted(fts_results["space_id"].to_pylist()) == expected_ids
hybrid_results = await (
table.query()
.where(predicate)
.nearest_to([5.0, 5.0])
.nearest_to_text("book")
.limit(num_rows * 2)
.to_arrow()
)
assert sorted(hybrid_results["space_id"].to_pylist()) == expected_ids
@pytest.mark.asyncio
async def test_async_hybrid_query_default_limit(table: AsyncTable):
# add 10 new rows
@@ -203,19 +123,6 @@ async def test_async_hybrid_query_default_limit(table: AsyncTable):
assert texts.count("a") == 1
def test_hybrid_query_minimum_nprobes_zero_raises(sync_table: Table):
# minimum_nprobes(0) must raise the same validation error a plain vector
# query raises, not silently no-op because 0 is falsy.
with pytest.raises(ValueError, match="minimum_nprobes must be greater than 0"):
(
sync_table.search(query_type="hybrid")
.vector([0.0, 0.4])
.text("dog")
.minimum_nprobes(0)
.to_arrow()
)
def test_hybrid_query_distance_range(sync_table: Table):
reranker = RRFReranker(return_score="all")
result = (
-33
View File
@@ -1,33 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import re
import shutil
import subprocess
import sys
import lancedb._lancedb as _lancedb
import pytest
@pytest.mark.skipif(sys.platform != "linux", reason="ldd is Linux-specific")
def test_native_extension_does_not_link_openssl():
"""OpenSSL-linked wheels abort when imported on RHEL hosts in FIPS mode."""
ldd = shutil.which("ldd")
if ldd is None:
pytest.skip("ldd is not installed")
result = subprocess.run(
[ldd, _lancedb.__file__],
check=True,
capture_output=True,
text=True,
)
openssl_libraries = re.findall(
r"^\s*(lib(?:crypto|ssl)\S*)\s+=>", result.stdout, flags=re.MULTILINE
)
assert not openssl_libraries, (
"the LanceDB native extension must use rustls instead of linking OpenSSL: "
f"{openssl_libraries}"
)
-34
View File
@@ -84,15 +84,6 @@ async def binary_table(db_async):
)
@pytest.mark.asyncio
async def test_create_index_async_returns_done_job(some_table: AsyncTable):
job = await some_table.create_index_async("id", config=BTree())
assert job.id is None
await job.wait()
assert len(await some_table.list_indices()) == 1
await job.cancel()
@pytest.mark.asyncio
async def test_create_scalar_index(some_table: AsyncTable):
# Can create
@@ -372,31 +363,6 @@ async def test_create_vector_index(some_table: AsyncTable):
assert stats.num_indices == 1
@pytest.mark.asyncio
async def test_create_ivf_index_reports_unsplittable_partitions(db_async):
dim = 8
num_partitions = 300 # More than 256 selects hierarchical k-means.
base_vectors = [[float(row == column) for column in range(dim)] for row in range(5)]
vectors = pa.array(base_vectors * 200, pa.list_(pa.float32(), dim))
table = await db_async.create_table(
"unsplittable_partitions",
pa.table({"vector": vectors}),
)
error_pattern = (
rf"Cannot create {num_partitions} IVF partitions: k-means could only form"
)
with pytest.raises(RuntimeError, match=error_pattern):
await table.create_index(
"vector",
config=IvfFlat(
distance_type="dot",
num_partitions=num_partitions,
max_iterations=10,
),
)
@pytest.mark.asyncio
async def test_create_4bit_ivfpq_index(some_table: AsyncTable):
# Can create
+4 -11
View File
@@ -83,9 +83,7 @@ def test_lsm_write_spec_repr():
assert s.spec_type == "bucket"
assert s.column == "id"
assert s.num_buckets == 4
# A fresh spec defers its maintained set to install time.
assert s.maintained_indexes is None
assert s.with_maintained_indexes([]).maintained_indexes == []
assert s.maintained_indexes == []
assert "bucket" in repr(s)
assert "id" in repr(s)
assert "4" in repr(s)
@@ -171,23 +169,18 @@ def test_get_lsm_write_spec(tmp_path):
table.unset_lsm_write_spec()
assert table.get_lsm_write_spec() is None
# Identity round-trips (column recovered from the schema). Leaving the
# maintained set to be inferred picks up the index on the table, so the
# spec reads back naming it rather than as "infer".
# Identity round-trips (column recovered from the schema).
table.set_lsm_write_spec(LsmWriteSpec.identity("id"))
spec = table.get_lsm_write_spec()
assert spec.spec_type == "identity"
assert spec.column == "id"
assert spec.maintained_indexes == [idx_name]
table.unset_lsm_write_spec()
# Unsharded round-trips (no routing column). Opting out is distinct from
# the inferred default.
table.set_lsm_write_spec(LsmWriteSpec.unsharded().with_maintained_indexes([]))
# Unsharded round-trips (no routing column).
table.set_lsm_write_spec(LsmWriteSpec.unsharded())
spec = table.get_lsm_write_spec()
assert spec.spec_type == "unsharded"
assert spec.column is None
assert spec.maintained_indexes == []
@pytest.mark.asyncio
+2 -2
View File
@@ -544,7 +544,7 @@ def test_lsm_read_fts_unmaintained_index_errors(tmp_path):
table.create_index("text", config=FTS())
# No maintained indexes: the active memtable FTS arm cannot serve un-compacted
# docs, so the search would silently omit them — reject instead.
table.set_lsm_write_spec(LsmWriteSpec.unsharded().with_maintained_indexes([]))
table.set_lsm_write_spec(LsmWriteSpec.unsharded())
with pytest.raises(Exception, match="maintained"):
table.search("fox", query_type="fts", fts_columns="text").to_arrow()
@@ -631,7 +631,7 @@ def test_lsm_read_vector_unmaintained_index_errors(tmp_path):
)
# Spec with NO maintained indexes: the base vector index's catch-up is untracked,
# so the scanner rejects rather than risk dropping compacted-but-unindexed rows.
table.set_lsm_write_spec(LsmWriteSpec.unsharded().with_maintained_indexes([]))
table.set_lsm_write_spec(LsmWriteSpec.unsharded())
with pytest.raises(Exception, match="maintained"):
table.search([1.0] * VECTOR_DIM).to_arrow()
@@ -18,7 +18,6 @@ Tests verify:
"""
import copy
import os
import shutil
import sys
import tempfile
@@ -240,7 +239,7 @@ def create_tracking_namespace(
dir_props = {f"storage.{k}": v for k, v in storage_options_with_refresh.items()}
if os.path.isabs(bucket_name) or bucket_name.startswith("file://"):
if bucket_name.startswith("/") or bucket_name.startswith("file://"):
dir_props["root"] = f"{bucket_name}/namespace_root"
else:
dir_props["root"] = f"s3://{bucket_name}/namespace_root"
@@ -768,70 +767,3 @@ def test_namespace_with_schema_only(s3_bucket: str, use_custom: bool):
# Verify data was added
assert table.count_rows() == 2
@pytest.mark.parametrize("use_custom", [False, True], ids=["DirectoryNS", "CustomNS"])
def test_namespace_exists(use_custom: bool):
"""
Test namespace_exists returns True for existing and False for non-existent.
"""
temp_dir = tempfile.mkdtemp()
try:
ns_client, _ = create_tracking_namespace(
bucket_name=temp_dir,
storage_options={},
credential_expires_in_seconds=3600,
use_custom=use_custom,
)
db = LanceNamespaceDBConnection(ns_client)
namespace_name = f"test_ns_{uuid.uuid4().hex[:8]}"
db.create_namespace([namespace_name])
# Existing namespace should return True
assert db.namespace_exists(namespace_id=[namespace_name]) is True
# Non-existent namespace should return False
assert db.namespace_exists(namespace_id=["nonexistent_ns"]) is False
finally:
shutil.rmtree(temp_dir, ignore_errors=True)
@pytest.mark.parametrize("use_custom", [False, True], ids=["DirectoryNS", "CustomNS"])
def test_table_exists(use_custom: bool):
"""
Test table_exists returns True for existing table and False for non-existent.
"""
temp_dir = tempfile.mkdtemp()
try:
ns_client, _ = create_tracking_namespace(
bucket_name=temp_dir,
storage_options={},
credential_expires_in_seconds=3600,
use_custom=use_custom,
)
db = LanceNamespaceDBConnection(ns_client)
namespace_name = f"test_ns_{uuid.uuid4().hex[:8]}"
db.create_namespace([namespace_name])
table_name = f"test_table_{uuid.uuid4().hex}"
namespace_path = [namespace_name]
schema = pa.schema(
[
pa.field("id", pa.int64()),
pa.field("vector", pa.list_(pa.float32(), 2)),
pa.field("text", pa.string()),
]
)
db.create_table(table_name, schema=schema, namespace_path=namespace_path)
# Existing table should return True
table_id = namespace_path + [table_name]
assert db.table_exists(table_id=table_id) is True
# Non-existent table should return False
assert db.table_exists(table_id=namespace_path + ["nonexistent_table"]) is False
finally:
shutil.rmtree(temp_dir, ignore_errors=True)
@@ -1,42 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import importlib
import re
import sys
from pathlib import Path
import pytest
def test_pyo3_abi_matches_minimum_supported_python():
project_dir = Path(__file__).parents[2]
pyproject = (project_dir / "pyproject.toml").read_text()
cargo_manifest = (project_dir / "Cargo.toml").read_text()
minimum_python = re.search(
r'^requires-python\s*=\s*">=(\d+)\.(\d+)"$', pyproject, re.MULTILINE
)
assert minimum_python is not None
major, minor = minimum_python.groups()
expected_abi = f"abi3-py{major}{minor}"
configured_abis = re.findall(r'"(abi3-py\d+)"', cargo_manifest)
assert configured_abis == [expected_abi, expected_abi], (
"the pyo3 runtime and build ABI features must both match requires-python"
)
@pytest.mark.skipif(sys.platform != "win32", reason="Windows wheel regression test")
def test_windows_wheel_tag_and_native_import():
project_dir = Path(__file__).parents[2]
wheels = list((project_dir.parent / "target" / "wheels").glob("lancedb-*.whl"))
if not wheels:
pytest.skip("no wheel artifact is available in this development environment")
assert len(wheels) == 1
assert wheels[0].name.endswith("-cp310-abi3-win_amd64.whl")
native_module = importlib.import_module("lancedb._lancedb")
assert Path(native_module.__file__).suffix == ".pyd"
-20
View File
@@ -6,7 +6,6 @@ import math
import pytest
from lancedb import DBConnection, Table, connect
from lancedb.background_loop import LOOP
from lancedb.permutation import Permutation, Permutations, permutation_builder
@@ -32,25 +31,6 @@ def test_split_random_ratios(mem_db):
assert 65 <= split_1_count <= 75 # ~70% ± tolerance
def test_execute_does_not_reenter_background_loop(tmp_path, monkeypatch):
import threading
db = connect(tmp_path)
tbl = db.create_table("test_table", pa.table({"x": range(10)}))
original_run = LOOP.run
def fail_on_reentry(future):
assert threading.current_thread() is not LOOP.thread
return original_run(future)
monkeypatch.setattr(LOOP, "run", fail_on_reentry)
permutation_tbl = permutation_builder(tbl).execute()
assert permutation_tbl.count_rows() == 10
assert permutation_tbl._conn.read_consistency_interval is None
def test_split_random_counts(mem_db):
"""Test random splitting with absolute counts."""
tbl = mem_db.create_table(
-11
View File
@@ -415,17 +415,6 @@ def test_nullable_vector():
assert schema == pa.schema([pa.field("vec", pa.list_(pa.float32(), 16), True)])
def test_bare_vector_raises_clear_error():
namespace = {
"__name__": "test_model_without_pyarrow",
"LanceModel": LanceModel,
"Vector": Vector,
}
with pytest.raises(TypeError, match=r"Vector must be parameterized.*Vector\(128\)"):
exec("class TestModel(LanceModel):\n vector: Vector", namespace)
def test_fixed_size_list_field():
class TestModel(pydantic.BaseModel):
vec: Vector(16)
-9
View File
@@ -570,15 +570,6 @@ def test_query_builder(table):
assert all(np.array(rs[0]["vector"]) == [1, 2])
def test_query_multiple_vectors(table):
results = table.search([np.array([1, 2]), np.array([4, 5])]).limit(1).to_list()
assert len(results) == 2
results_by_query = {result["query_index"]: result for result in results}
assert results_by_query[0]["id"] == 1
assert results_by_query[1]["id"] == 2
def test_with_row_id(table: lancedb.table.Table):
rs = table.search().with_row_id(True).to_arrow()
assert "_rowid" in rs.column_names
+1 -674
View File
@@ -35,12 +35,6 @@ def make_mock_http_handler(handler):
return MockLanceDBHandler
@pytest.mark.parametrize("db_name", ["a" * 64, "invalid..database"])
def test_connect_rejects_invalid_cloud_dns_hostname(db_name):
with pytest.raises(ValueError, match="DNS labels must contain 1 to 63 bytes"):
lancedb.connect(f"db://{db_name}", api_key="fake")
@contextlib.contextmanager
def mock_lancedb_connection(handler):
with http.server.HTTPServer(
@@ -818,121 +812,6 @@ def test_table_create_indices():
table.drop_index("custom_fts_idx")
def test_remote_create_index_async_returns_job():
from lancedb.index import BTree
describe_calls = []
def handler(request):
content_len = int(request.headers.get("Content-Length", 0))
body = request.rfile.read(content_len) if content_len > 0 else b""
if request.path == "/v1/table/test/create_index/":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b'{"job_id": "job-1"}')
elif request.path == "/v1/jobs/describe":
assert json.loads(body)["job_id"] == "job-1"
describe_calls.append(1)
state = "IN_PROGRESS" if len(describe_calls) == 1 else "DONE"
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(dict(job_id="job-1", job_state=state)).encode()
)
elif request.path == "/v1/jobs/cancel":
assert json.loads(body)["job_id"] == "job-1"
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b"{}")
elif request.path == "/v1/table/test/create/?mode=create":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b"{}")
elif request.path == "/v1/table/test/describe/":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
dict(
version=1,
schema=dict(
fields=[
dict(name="id", type={"type": "int64"}, nullable=False),
]
),
)
).encode()
)
else:
request.send_response(404)
request.end_headers()
with mock_lancedb_connection(handler) as db:
table = db.create_table("test", [{"id": 1}])
job = table.create_index_async("id", config=BTree())
assert job.id == "job-1"
job.wait(timeout=timedelta(seconds=30))
assert len(describe_calls) == 2
job.cancel()
def test_remote_job_wait_raises_on_failure():
from lancedb.exceptions import JobFailedError
from lancedb.index import BTree
def handler(request):
content_len = int(request.headers.get("Content-Length", 0))
body = request.rfile.read(content_len) if content_len > 0 else b""
if request.path == "/v1/table/test/create_index/":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b'{"job_id": "job-2"}')
elif request.path == "/v1/jobs/describe":
assert json.loads(body)["job_id"] == "job-2"
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(dict(job_id="job-2", job_state="FAILED")).encode()
)
elif request.path == "/v1/table/test/create/?mode=create":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b"{}")
elif request.path == "/v1/table/test/describe/":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
dict(
version=1,
schema=dict(
fields=[
dict(name="id", type={"type": "int64"}, nullable=False),
]
),
)
).encode()
)
else:
request.send_response(404)
request.end_headers()
with mock_lancedb_connection(handler) as db:
table = db.create_table("test", [{"id": 1}])
job = table.create_index_async("id", config=BTree())
with pytest.raises(JobFailedError, match="job-2"):
job.wait()
def test_remote_create_index_new_api():
received_requests = []
@@ -1141,7 +1020,7 @@ def query_test_table(query_handler, *, server_version=Version("0.1.0")):
request.send_header("Content-Type", "application/json")
request.send_header("phalanx-version", str(server_version))
request.end_headers()
request.wfile.write(b'{"version": 1, "schema": {"fields": []}}')
request.wfile.write(b"{}")
elif request.path == "/v1/table/test/query/":
content_len = int(request.headers.get("Content-Length"))
body = request.rfile.read(content_len)
@@ -1979,555 +1858,3 @@ def test_inherited_remote_table_reopens_after_fork():
finally:
server.shutdown()
server_thread.join()
BLOB_DESCRIBE_RESPONSE = {
"table": "test",
"version": 1,
"schema": {
"fields": [
{"name": "id", "type": {"type": "int64"}, "nullable": False},
{
"name": "image",
"type": {
"type": "struct",
"fields": [
{
"name": "data",
"type": {"type": "large_binary"},
"nullable": True,
},
{"name": "uri", "type": {"type": "string"}, "nullable": True},
],
},
"nullable": True,
"metadata": {
"ARROW:extension:name": "lance.blob.v2",
"ARROW:extension:metadata": "",
},
},
]
},
}
def blob_query_response_table():
image_field = pa.field(
"image",
pa.struct(
[
pa.field("kind", pa.uint8(), nullable=False),
pa.field("position", pa.uint64(), nullable=False),
pa.field("size", pa.uint64(), nullable=False),
pa.field("blob_id", pa.uint32(), nullable=False),
pa.field("blob_uri", pa.string(), nullable=False),
]
),
metadata={"lance-encoding:blob": "true"},
)
images = pa.StructArray.from_arrays(
[
pa.array([1, 0, 0], type=pa.uint8()),
pa.array([0, 0, 0], type=pa.uint64()),
pa.array([5, 0, 5], type=pa.uint64()),
pa.array([1, 0, 2], type=pa.uint32()),
pa.array(["", "", ""], type=pa.string()),
],
fields=image_field.type,
mask=pa.array([False, True, False]),
)
return pa.Table.from_arrays(
[
pa.array([1, 2, 3], type=pa.int64()),
images,
pa.array([10, 20, 30], type=pa.uint64()),
],
schema=pa.schema(
[
pa.field("id", pa.int64(), nullable=False),
image_field,
pa.field("_rowid", pa.uint64()),
]
),
)
@contextlib.contextmanager
def blob_remote_table(*, server_version=Version("0.5.0")):
def handler(request):
if request.path == "/v1/table/test/describe/":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.send_header("phalanx-version", str(server_version))
request.end_headers()
request.wfile.write(json.dumps(BLOB_DESCRIBE_RESPONSE).encode())
elif request.path.startswith("/v1/table/test/blob/image/"):
path = request.path.partition("?")[0]
row_id = int(path.split("/")[-2])
payload = {10: b"alpha", 20: None, 30: b"gamma"}[row_id]
if payload is None:
request.send_response(204)
request.end_headers()
return
byte_range = request.headers["Range"].removeprefix("bytes=")
start_text, end_text = byte_range.split("-", maxsplit=1)
start = int(start_text)
end = int(end_text) if end_text else len(payload) - 1
chunk = payload[start : end + 1]
request.send_response(206)
request.send_header("Content-Range", f"bytes {start}-{end}/{len(payload)}")
request.send_header("Content-Length", str(len(chunk)))
request.end_headers()
request.wfile.write(chunk)
elif request.path == "/v1/table/test/query/":
content_len = int(request.headers.get("Content-Length", 0))
body = json.loads(request.rfile.read(content_len))
assert body["columns"] == ["id", "image"]
assert body["with_row_id"] is True
response_table = blob_query_response_table()
request.send_response(200)
request.send_header("Content-Type", "application/vnd.apache.arrow.file")
request.end_headers()
with pa.ipc.new_file(request.wfile, response_table.schema) as writer:
writer.write_table(response_table)
elif request.path == "/v1/table/test/fetch_blobs/":
content_len = int(request.headers.get("Content-Length", 0))
body = json.loads(request.rfile.read(content_len))
assert body["column"] == "image"
assert body["row_ids"] == [10, 20, 30]
response_table = pa.table(
{"image": pa.array([b"alpha", None, b"gamma"], type=pa.large_binary())}
)
request.send_response(200)
request.send_header("Content-Type", "application/vnd.apache.arrow.stream")
request.end_headers()
with pa.ipc.new_stream(request.wfile, response_table.schema) as writer:
writer.write_table(response_table)
else:
request.send_response(404)
request.end_headers()
with mock_lancedb_connection(handler) as db:
yield db.open_table("test")
def test_remote_blob_columns_and_fetch():
with blob_remote_table() as table:
assert table.blob_columns() == ["image"]
blobs = table.fetch_blobs("image", [10, 20, 30])
assert blobs.to_pylist() == [b"alpha", None, b"gamma"]
def test_remote_blob_files_are_lazy_seekable_handles():
with blob_remote_table() as table:
files = table.fetch_blob_files("image", [10, 20, 30])
assert len(files) == 3
alpha, null_row, gamma = files
assert null_row is None
assert alpha is not None
assert gamma is not None
assert alpha.size() == 5
assert alpha.read_range(1, 3) == b"lph"
gamma.seek(2)
assert gamma.read() == b"mma"
def test_remote_blob_fetch_accepts_query_table():
hits = pa.table({"_rowid": pa.array([10, 20, 30], type=pa.uint64())})
with blob_remote_table() as table:
blobs = table.fetch_blobs("image", hits)
assert blobs.to_pylist() == [b"alpha", None, b"gamma"]
def test_remote_blob_query_stashes_row_ids_for_fetch():
with blob_remote_table() as table:
hits = table.search().select(["id", "image"]).limit(3).to_arrow()
assert "_rowid" not in hits.column_names
assert "_lance_row_id" in hits.schema.field("image").type.names
blobs = table.fetch_blobs("image", hits)
assert blobs.to_pylist() == [b"alpha", None, b"gamma"]
def test_remote_blob_query_survives_a_server_that_ignores_the_row_id_request():
def handler(request):
if request.path == "/v1/table/test/describe/":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.send_header("phalanx-version", "0.5.0")
request.end_headers()
request.wfile.write(json.dumps(BLOB_DESCRIBE_RESPONSE).encode())
elif request.path == "/v1/table/test/query/":
content_len = int(request.headers.get("Content-Length", 0))
assert json.loads(request.rfile.read(content_len))["with_row_id"] is True
response_table = blob_query_response_table().drop_columns(["_rowid"])
request.send_response(200)
request.send_header("Content-Type", "application/vnd.apache.arrow.file")
request.end_headers()
with pa.ipc.new_file(request.wfile, response_table.schema) as writer:
writer.write_table(response_table)
else:
request.send_response(404)
request.end_headers()
with mock_lancedb_connection(handler) as db:
table = db.open_table("test")
hits = table.search().select(["id", "image"]).limit(3).to_arrow()
assert hits.column_names == ["id", "image"]
assert "_lance_row_id" not in hits.schema.field("image").type.names
with pytest.raises(ValueError, match="pass a list of row ids"):
table.fetch_blobs("image", hits)
def test_remote_blob_byte_apis_not_supported_on_old_server():
with blob_remote_table(server_version=Version("0.1.0")) as table:
assert table.blob_columns() == ["image"]
with pytest.raises(NotImplementedError, match="not supported"):
table.fetch_blobs("image", [1])
with pytest.raises(NotImplementedError, match="not supported"):
table.fetch_blob_files("image", [1])
def test_remote_connection_jobs_surface():
from lancedb.exceptions import JobFailedError
schema = pa.schema([("state", pa.string())])
batch = pa.record_batch([pa.array(["created", "done"])], schema=schema)
sink = pa.BufferOutputStream()
with pa.ipc.new_stream(sink, schema) as writer:
writer.write_batch(batch)
events_body = sink.getvalue().to_pybytes()
def handler(request):
content_len = int(request.headers.get("Content-Length", 0))
body = request.rfile.read(content_len) if content_len > 0 else b""
payload = json.loads(body) if body else {}
if request.path == "/v1/jobs/list":
if payload.get("page_token") is None:
rsp = dict(
jobs=[
dict(
job_id="job-1",
table="t1",
job_type="create_index",
state="in_progress",
created_at_millis=1000,
)
],
page_token="next",
)
else:
assert payload["page_token"] == "next"
rsp = dict(
jobs=[
dict(
job_id="job-2",
table="t2",
job_type="create_index",
state="succeeded",
created_at_millis=2000,
)
]
)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(rsp).encode())
elif request.path == "/v1/jobs/describe":
if payload["job_id"] != "job-1":
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
dict(
job_id="job-1",
job_type="create_index",
job_state="FAILED",
creation_ms=1000,
spec=dict(column="vec"),
failure=dict(
phase="execute", message="worker died", retryable=True
),
)
).encode()
)
elif request.path == "/v1/jobs/cancel":
if payload["job_id"] != "job-1":
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b'{"job_id": "job-1"}')
elif request.path == "/v1/jobs/query_events":
assert payload["job_id"] == "job-1"
request.send_response(200)
request.send_header("Content-Type", "application/vnd.apache.arrow.stream")
request.end_headers()
request.wfile.write(events_body)
else:
request.send_response(404)
request.end_headers()
with mock_lancedb_connection(handler) as db:
jobs = db.list_jobs()
assert [job.job_id for job in jobs] == ["job-1", "job-2"]
assert jobs[0].state == "running"
assert jobs[0].table == "t1"
assert jobs[1].state == "finished"
description = db.get_job("job-1")
assert description.job_type == "create_index"
assert description.state == "failed"
assert json.loads(description.spec_json) == {"column": "vec"}
assert description.failure.message == "worker died"
assert description.failure.retryable is True
assert db.get_job("missing") is None
assert db.cancel_job("job-1") is True
assert db.cancel_job("missing") is False
batches = db.job_history("job-1")
assert len(batches) == 1
assert batches[0].num_rows == 2
assert batches[0].column("state").to_pylist() == ["created", "done"]
job = db.job("job-1")
assert job.id == "job-1"
assert job.status() == "failed"
with pytest.raises(JobFailedError, match="worker died"):
job.wait(timeout=timedelta(seconds=5))
# Pinned Rust-canonical schema-only type IPC (base64). PyArrow's schema-only
# FileWriter bytes are not byte-identical to the Arrow Rust FileWriter used by
# the strict Function decoder, so these fixtures are derived from Rust serde.
_FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID = "fn.exact.python-job-result"
_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT = object()
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL = object()
def _first_class_function_job_result_function_wire():
int32_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_INT32_TYPE_IPC_B64
utf8_ipc = _FIRST_CLASS_FUNCTION_JOB_RESULT_UTF8_TYPE_IPC_B64
return {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "x", "data_type_ipc": int32_ipc},
{"name": "label", "data_type_ipc": utf8_ipc},
],
"output": {
"data_type_ipc": int32_ipc,
"nullable": True,
},
},
},
}
def _first_class_function_job_result_none_wire():
return {"kind": "none", "format_version": 1}
def _first_class_function_job_result_describe_body(
job_id, job_type, result=_FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT
):
body = {
"job_id": job_id,
"job_state": "DONE",
"job_type": job_type,
"creation_ms": 1,
"spec": {},
}
if result is _FIRST_CLASS_FUNCTION_JOB_RESULT_NULL:
body["result"] = None
elif result is not _FIRST_CLASS_FUNCTION_JOB_RESULT_ABSENT:
body["result"] = result
return body
def _first_class_function_job_result_describe_handler(bodies_by_job_id):
def handler(request):
content_len = int(request.headers.get("Content-Length", 0))
body = request.rfile.read(content_len) if content_len > 0 else b""
payload = json.loads(body) if body else {}
if request.path != "/v1/jobs/describe":
request.send_response(404)
request.end_headers()
return
job_id = payload["job_id"]
if job_id not in bodies_by_job_id:
request.send_response(404)
request.end_headers()
return
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps(bodies_by_job_id[job_id]).encode())
return handler
def _assert_exact_first_class_function_job_result(function):
assert isinstance(function, lancedb.Function)
assert function is not None
assert not isinstance(function, dict)
assert function.id == _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID
assert function.parameters == (("x", pa.int32()), ("label", pa.utf8()))
assert function.output_type == pa.int32()
assert function.output_nullable is True
text = repr(function)
assert "Function" in text
assert _FIRST_CLASS_FUNCTION_JOB_RESULT_FUNCTION_ID in text
for token in ("definition", "source", "packages", "artifact", "digest", "secret"):
assert token not in text.lower()
def test_first_class_function_job_result_sync_wait_returns_exact_function():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
)
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
result = db.job("job-register").wait()
_assert_exact_first_class_function_job_result(result)
timed_out = db.job("job-register").wait(timeout=timedelta(seconds=5))
_assert_exact_first_class_function_job_result(timed_out)
with pytest.raises(TypeError):
lancedb.Function()
with pytest.raises(AttributeError):
result.id = "mutated"
with pytest.raises(AttributeError):
result.parameters = ()
with pytest.raises(AttributeError):
result.output_type = pa.int64()
with pytest.raises(AttributeError):
result.output_nullable = False
@pytest.mark.asyncio
async def test_first_class_function_job_result_async_wait_returns_exact_function():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
)
}
async with mock_lancedb_connection_async(
_first_class_function_job_result_describe_handler(bodies)
) as db:
result = await db.job("job-register").wait()
_assert_exact_first_class_function_job_result(result)
timed_out = await db.job("job-register").wait(timeout=timedelta(seconds=5))
_assert_exact_first_class_function_job_result(timed_out)
def test_first_class_function_job_result_no_result_wait_returns_none():
bodies = {
"job-index-absent": _first_class_function_job_result_describe_body(
"job-index-absent", "create_index"
),
"job-index-explicit": _first_class_function_job_result_describe_body(
"job-index-explicit",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
assert db.job("job-index-absent").wait() is None
assert db.job("job-index-explicit").wait(timeout=timedelta(seconds=5)) is None
@pytest.mark.asyncio
async def test_first_class_function_job_result_async_no_result_wait_returns_none():
bodies = {
"job-index-absent": _first_class_function_job_result_describe_body(
"job-index-absent", "create_index"
),
"job-index-explicit": _first_class_function_job_result_describe_body(
"job-index-explicit",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
async with mock_lancedb_connection_async(
_first_class_function_job_result_describe_handler(bodies)
) as db:
assert await db.job("job-index-absent").wait() is None
assert (
await db.job("job-index-explicit").wait(timeout=timedelta(seconds=5))
is None
)
def test_first_class_function_job_result_get_job_result_projection():
bodies = {
"job-register": _first_class_function_job_result_describe_body(
"job-register",
"register_function",
_first_class_function_job_result_function_wire(),
),
"job-absent": _first_class_function_job_result_describe_body(
"job-absent", "create_index"
),
"job-null": _first_class_function_job_result_describe_body(
"job-null",
"create_index",
_FIRST_CLASS_FUNCTION_JOB_RESULT_NULL,
),
"job-explicit-none": _first_class_function_job_result_describe_body(
"job-explicit-none",
"create_index",
_first_class_function_job_result_none_wire(),
),
}
with mock_lancedb_connection(
_first_class_function_job_result_describe_handler(bodies)
) as db:
register_description = db.get_job("job-register")
_assert_exact_first_class_function_job_result(register_description.result)
assert db.get_job("job-absent").result is None
assert db.get_job("job-null").result is None
assert db.get_job("job-explicit-none").result is None
+7 -334
View File
@@ -2,14 +2,10 @@
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import ctypes
import gc
import os
import sys
import threading
import warnings
import weakref
from concurrent.futures import ThreadPoolExecutor
from datetime import date, datetime, timedelta
from time import sleep
from typing import List
@@ -102,30 +98,6 @@ def test_basic(mem_db: DBConnection):
assert table.to_arrow() == expected_data
def test_search_preserves_nulls_from_sliced_arrow_table(mem_db: DBConnection):
data = pa.table(
{
"id": [0, 1, 2, 3, 4],
"score_cn": [None, 22, None, 5, 8],
"score_mt": [None, 42, None, 5, 8],
"vector": [
[20, 19, -1, -1],
[41, 38, 22, 42],
[10, 10, -1, -1],
[5, 5, 5, 5],
[8, 8, 8, 8],
],
}
).slice(1)
table = mem_db.create_table("sliced_nullable", data=data)
result = table.search([41, 38, 22, 42]).limit(1).to_arrow()
assert result["id"].to_pylist() == [1]
assert result["score_cn"].to_pylist() == [22]
assert result["score_mt"].to_pylist() == [42]
def test_table_to_pandas_default_matches_arrow(tmp_db: DBConnection):
pd = pytest.importorskip("pandas")
data = pa.table({"id": [1, 2], "text": ["one", "two"]})
@@ -462,38 +434,6 @@ def test_add(mem_db: DBConnection):
_add(table, schema)
def test_add_releases_arrow_buffers_without_gc(mem_db: DBConnection):
"""Regression test for https://github.com/lancedb/lancedb/issues/2512."""
schema = pa.schema([pa.field("x", pa.int64())])
table = mem_db.create_table("test_add_releases_arrow_buffers", schema=schema)
class BufferOwner:
def __init__(self, size: int):
self.memory = ctypes.create_string_buffer(size)
owner_refs = []
gc_was_enabled = gc.isenabled()
gc.disable()
try:
for _ in range(3):
size = 8 * 1024
owner = BufferOwner(size)
arrow_buffer = pa.foreign_buffer(
ctypes.addressof(owner.memory), size, owner
)
array = pa.Array.from_buffers(pa.int64(), 1024, [None, arrow_buffer])
batch = pa.RecordBatch.from_arrays([array], schema=schema)
owner_refs.append(weakref.ref(owner))
table.add(batch)
del batch, array, arrow_buffer, owner
assert all(owner_ref() is None for owner_ref in owner_refs)
finally:
if gc_was_enabled:
gc.enable()
def test_add_write_parallelism(mem_db: DBConnection):
schema = pa.schema([pa.field("id", pa.int64())])
table = mem_db.create_table("test", schema=schema)
@@ -929,7 +869,6 @@ def test_polars(mem_db: DBConnection):
# enter table to polars dataframe
result = table.to_polars()
assert isinstance(result, pl.LazyFrame)
assert np.allclose(result.collect()["vector"].to_list(), data["vector"])
# make sure filtering isn't broken
@@ -1463,15 +1402,6 @@ async def test_async_open_table_with_branch_version(tmp_path):
assert await pinned.count_rows() == 4 # writable again
def test_create_index_async_returns_done_job(mem_db: DBConnection):
table = mem_db.create_table("job_test", [{"id": i} for i in range(10)])
job = table.create_index_async("id", config=BTree())
assert job.id is None
job.wait()
assert len(table.list_indices()) == 1
job.cancel()
@patch("lancedb.table.AsyncTable.create_index")
def test_create_index_method(mock_create_index, mem_db: DBConnection):
table = mem_db.create_table(
@@ -1846,27 +1776,6 @@ def test_add_with_empty_fixed_size_list_drops_bad_rows(mem_db: DBConnection):
assert np.allclose(data["embedding"].to_pylist()[0], np.array([0.1] * 16))
def test_add_nullable_fixed_size_list_with_none(mem_db: DBConnection):
"""Regression test for issue #2340."""
table = mem_db.create_table(
"test_nullable_fixed_size_list",
schema=pa.schema(
[
pa.field("id", pa.string()),
pa.field("feature", pa.list_(pa.float32(), 256)),
pa.field("tags", pa.list_(pa.string())),
]
),
)
table.add([{"id": "1", "feature": None, "tags": ["tag1", "tag2"]}])
result = table.to_arrow()
assert result.to_pylist() == [
{"id": "1", "feature": None, "tags": ["tag1", "tag2"]}
]
def test_add_nullable_struct_with_none(mem_db: DBConnection):
"""Regression test for issue #2654: a nullable struct column whose
first batch contains only None values must not crash in
@@ -1906,33 +1815,6 @@ def test_add_nullable_struct_with_none(mem_db: DBConnection):
assert result.column("data").to_pylist() == [{"x": 1.0}, None]
def test_read_mostly_null_list_v2_2_page_boundary(tmp_path):
# Regression test for #3194. This row/value count crosses a v2.2 structural
# encoding page boundary where Lance 3.0.0 sliced repetition/definition
# levels by row offset and decoded child arrays at different lengths.
num_rows = 64_885
num_values = 217
list_type = pa.list_(pa.float32())
source = pa.table(
{
"id": np.arange(num_rows, dtype=np.int64),
"coords": pa.array(
[[1.0, 2.0, 3.0, 4.0]] * num_values + [None] * (num_rows - num_values),
type=list_type,
),
}
)
db = lancedb.connect(
tmp_path,
storage_options={"new_table_data_storage_version": "2.2"},
)
table = db.create_table("test_sparse_nullable_list", data=source)
result = table.search().select(["id", "coords"]).limit(num_rows).to_arrow()
assert result.equals(source)
def test_add_with_integer_embeddings_preserves_casting(mem_db: DBConnection):
class Schema(LanceModel):
text: str
@@ -2218,45 +2100,6 @@ def test_merge(tmp_db: DBConnection, tmp_path):
table.merge(other_dataset, left_on="id")
@pytest.mark.parametrize("storage_version", ["legacy", "stable"])
def test_search_after_merge(tmp_path, storage_version):
pytest.importorskip("lance")
pd = pytest.importorskip("pandas")
db = lancedb.connect(
tmp_path,
storage_options={"new_table_data_storage_version": storage_version},
)
rng = np.random.default_rng(42)
row_count = 512
vectors = rng.standard_normal((row_count, 8)).astype(np.float32)
table = db.create_table(
"search_after_merge",
data=pd.DataFrame(
{
"id": [str(i) for i in range(row_count)],
"vector": list(vectors),
}
),
)
table.create_index("vector", config=IvfPq(num_partitions=1, num_sub_vectors=2))
links = pd.DataFrame(
{
"id": [str(i) for i in range(row_count // 2)],
"link": [f"https://example.com/{i}" for i in range(row_count // 2)],
}
)
table.merge(links, left_on="id")
query = table.search(vectors[-1]).refine_factor(50).limit(10)
assert "ANN" in query.explain_plan(verbose=True)
result = query.to_arrow()
links_by_id = dict(zip(result["id"].to_pylist(), result["link"].to_pylist()))
assert links_by_id[str(row_count - 1)] is None
def test_delete(mem_db: DBConnection):
table = mem_db.create_table(
"my_table",
@@ -2272,27 +2115,6 @@ def test_delete(mem_db: DBConnection):
assert table.to_arrow()["id"].to_pylist() == [1]
def test_concurrent_deletes_are_thread_safe(mem_db: DBConnection):
num_workers = 8
table = mem_db.create_table(
"my_table", data=[{"id": row_id} for row_id in range(num_workers)]
)
barrier = threading.Barrier(num_workers)
def delete(row_id: int):
barrier.wait()
return table.delete(f"id = {row_id}")
with ThreadPoolExecutor(max_workers=num_workers) as pool:
results = list(pool.map(delete, range(num_workers)))
assert all(result.num_deleted_rows == 1 for result in results)
assert sorted(result.version for result in results) == list(
range(2, num_workers + 2)
)
assert table.count_rows() == 0
def test_delete_expr(mem_db: DBConnection):
table = mem_db.create_table(
"my_table",
@@ -2343,20 +2165,6 @@ def test_update(mem_db: DBConnection):
assert np.allclose(v, np.array([[1.2, 1.9], [1.1, 1.1]]))
def test_update_with_arrow_scalar(mem_db: DBConnection):
schema = pa.schema({"id": pa.int64(), "vector": pa.list_(pa.float32(), 4)})
table = mem_db.create_table("my_table", schema=schema)
table.add([{"id": 1, "vector": [1.0, 2.0, 3.0, 4.0]}])
value = table.search().select(["vector"]).limit(1).to_arrow()["vector"][0]
assert isinstance(value, pa.FixedSizeListScalar)
result = table.update(where="id == 1", values={"vector": value})
assert result.rows_updated == 1
assert table.to_arrow()["vector"].to_pylist() == [[1.0, 2.0, 3.0, 4.0]]
def test_update_types(mem_db: DBConnection):
table = mem_db.create_table(
"my_table",
@@ -2524,55 +2332,6 @@ def test_merge_insert(mem_db: DBConnection):
)
def test_merge_insert_nullable_pandas_into_pydantic_schema(mem_db: DBConnection):
# Regression test for https://github.com/lancedb/lancedb/issues/2366
pd = pytest.importorskip("pandas")
class Document(LanceModel):
id: int
title: str
content: str
table = mem_db.create_table("documents", schema=Document)
table.add(
pd.DataFrame(
{
"title": ["Old title", "Unchanged"],
"id": [2, 3],
"content": ["Old content", "Keep this"],
}
)
)
# Pandas produces nullable Arrow fields, in an order that differs from the
# non-nullable Pydantic schema. This is valid as long as the data has no nulls.
new_data = pd.DataFrame(
{
"title": ["Inserted", "Updated"],
"id": [1, 2],
"content": ["New row", "New content"],
}
)
result = (
table.merge_insert("id")
.when_matched_update_all()
.when_not_matched_insert_all()
.execute(new_data)
)
assert result.num_inserted_rows == 1
assert result.num_updated_rows == 1
expected = pa.Table.from_pylist(
[
{"id": 1, "title": "Inserted", "content": "New row"},
{"id": 2, "title": "Updated", "content": "New content"},
{"id": 3, "title": "Unchanged", "content": "Keep this"},
],
schema=Document.to_arrow_schema(),
)
assert table.to_arrow().sort_by("id") == expected
def test_merge_insert_by_source_delete_expr(mem_db: DBConnection):
table = mem_db.create_table(
"my_table",
@@ -2596,29 +2355,6 @@ def test_merge_insert_by_source_delete_expr(mem_db: DBConnection):
assert table.to_arrow().sort_by("a") == expected
def test_merge_insert_by_source_delete_reconfigure(mem_db: DBConnection):
# Calling when_not_matched_by_source_delete() again with no condition must
# widen the delete to unconditional, not keep the earlier condition around.
table = mem_db.create_table(
"my_table",
data=pa.table({"a": [1, 2, 3], "b": ["a", "b", "c"]}),
)
new_data = pa.table({"a": [2, 4], "b": ["x", "z"]})
merge_insert_res = (
table.merge_insert("a")
.when_matched_update_all()
.when_not_matched_insert_all()
.when_not_matched_by_source_delete("a > 2")
.when_not_matched_by_source_delete()
.execute(new_data)
)
assert merge_insert_res.num_deleted_rows == 2
expected = pa.table({"a": [2, 4], "b": ["x", "z"]})
assert table.to_arrow().sort_by("a") == expected
@pytest.mark.asyncio
async def test_merge_insert_by_source_delete_expr_async(
mem_db_async: AsyncConnection,
@@ -2673,36 +2409,6 @@ def test_merge_insert_subschema(mem_db: DBConnection, data_format):
assert table.to_arrow().sort_by("id") == expected
def test_repeated_partial_merge_insert_with_scalar_index(mem_db: DBConnection):
def make_batch(start: int) -> pa.Table:
return pa.table(
{
"id": [f"id-{i:04}" for i in range(start, start + 100)],
"category": ["A"] * 100,
"value_a": [float(i) for i in range(start, start + 100)],
"value_b": [float(i) / 10 for i in range(100)],
}
)
table = mem_db.create_table("my_table", data=make_batch(0))
table.add(make_batch(100))
table.add(make_batch(200))
table.create_index("id", config=BTree())
ids = [f"id-{i:04}" for i in range(100, 200)]
for value in (999.0, 888.0):
result = (
table.merge_insert("id")
.when_matched_update_all()
.execute(pa.table({"id": ids, "value_a": [value] * 100}))
)
assert result.num_updated_rows == 100
actual = table.to_arrow().sort_by("id")
assert actual.num_rows == 300
assert actual["value_a"].to_pylist()[100:200] == [888.0] * 100
@pytest.mark.asyncio
async def test_merge_insert_async(mem_db_async: AsyncConnection):
data = pa.table({"a": [1, 2, 3], "b": ["a", "b", "c"]})
@@ -2799,40 +2505,15 @@ def test_create_with_embedding_function(mem_db: DBConnection):
assert actual == expected
def test_create_f16_table_from_arrow_data(mem_db: DBConnection):
dimension = 32
num_rows = 512
values = pa.array(
np.random.default_rng(42)
.standard_normal(num_rows * dimension)
.astype(np.float16)
)
df = pa.table(
{
"text": [f"s-{i}" for i in range(num_rows)],
"vector": pa.FixedSizeListArray.from_arrays(values, dimension),
}
)
table = mem_db.create_table("f16_tbl", data=df)
assert table.schema.field("vector").type == pa.list_(pa.float16(), dimension)
table.create_index(num_partitions=2, num_sub_vectors=2)
query = df["vector"][2].as_py()
expected = table.search(query).limit(2).to_arrow()
assert "s-2" in expected["text"].to_pylist()
def test_create_f16_table(mem_db: DBConnection):
class MyTable(LanceModel):
text: str
vector: Vector(32, value_type=pa.float16())
rng = np.random.default_rng(42)
df = pa.table(
{
"text": [f"s-{i}" for i in range(512)],
"vector": [rng.standard_normal(32).astype(np.float16) for _ in range(512)],
"vector": [np.random.randn(32).astype(np.float16) for _ in range(512)],
}
)
table = mem_db.create_table(
@@ -3406,6 +3087,9 @@ def test_consistency(tmp_path, consistency_interval):
db2 = lancedb.connect(tmp_path, read_consistency_interval=consistency_interval)
table2 = db2.open_table("my_table")
if consistency_interval is not None:
assert "read_consistency_interval=datetime.timedelta(" in repr(db2)
assert "read_consistency_interval=datetime.timedelta(" in repr(table2)
assert table2.version == table.version
table.add([{"id": 1}])
@@ -3713,8 +3397,7 @@ def test_stats(mem_db: DBConnection):
stats = table.stats()
print(f"{stats=}")
assert stats == {
# Full on-disk size of the data file, footer and metadata included.
"total_bytes": 633,
"total_bytes": 60,
"num_rows": 2,
"num_indices": 0,
"fragment_stats": {
@@ -3732,13 +3415,6 @@ def test_stats(mem_db: DBConnection):
},
}
# Index files count toward total_bytes too (only deletion files and
# manifests are excluded).
table.create_index("id", config=BTree())
stats_with_index = table.stats()
assert stats_with_index["num_indices"] == 1
assert stats_with_index["total_bytes"] > stats["total_bytes"]
def test_create_table_empty_list_with_schema(mem_db: DBConnection):
"""Test creating table with empty list data and schema
@@ -3762,8 +3438,8 @@ def test_create_table_empty_list_no_schema_error(mem_db: DBConnection):
mem_db.create_table("test_empty_no_schema", data=[])
def test_create_table_without_data_with_vector_schema(tmp_path):
"""Test exact scenario from issue #1968.
def test_add_table_with_empty_embeddings(tmp_path):
"""Test exact scenario from issue #1968
Regression test for issue #1968:
https://github.com/lancedb/lancedb/issues/1968
@@ -3775,9 +3451,6 @@ def test_create_table_without_data_with_vector_schema(tmp_path):
embedding: Vector(16)
table = db.create_table("test", schema=MySchema)
assert table.count_rows() == 0
assert table.schema == MySchema.to_arrow_schema()
table.add(
[{"text": "bar", "embedding": [0.1] * 16}],
on_bad_vectors="drop",
@@ -75,22 +75,6 @@ class TestVoyageAIModelRegistration:
with pytest.raises(ValueError, match="not supported"):
func.ndims()
def test_voyage3_source_embeddings_use_text_api(self, mock_voyageai_client):
"""Regression test for text table data being sent to the multimodal API."""
mock_voyageai_client.tokenize.return_value = [["hello", "world"]]
mock_voyageai_client.embed.return_value.embeddings = [[0.1] * 1024]
registry = get_registry()
func = registry.get("voyageai").create(name="voyage-3")
embeddings = func.compute_source_embeddings("hello world")
assert embeddings == [[0.1] * 1024]
mock_voyageai_client.embed.assert_called_once_with(
texts=["hello world"], model="voyage-3", input_type="document"
)
mock_voyageai_client.multimodal_embed.assert_not_called()
@pytest.mark.parametrize(
"model_name",
[
-15
View File
@@ -1,15 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
from typing import assert_type
import lancedb
from lancedb import AsyncConnection, DBConnection
def check_connect_type() -> None:
assert_type(lancedb.connect("memory://"), DBConnection)
async def check_connect_async_type() -> None:
assert_type(await lancedb.connect_async("memory://"), AsyncConnection)

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