Compare commits

..
Author SHA1 Message Date
Lance Release 7b46f31cd8 Bump version: 0.38.0-beta.8 → 0.38.0-beta.9 2026-08-25 10:31:36 +00:00
Xuanwo c988e4848d refactor: simplify Function binding identity (#4046)
Function applications and bindings currently encode `group_id` and
binding `revision` even though `binding_id` already owns the complete
immutable binding lifecycle and `outputs` already defines the atomic
multi-output set.

Make `binding_id` the sole binding identity, remove the redundant fields
from the Rust and Python client contracts, and describe multi-output
declarations directly. This intentionally replaces the removed wire
fields without a compatibility path.
2026-08-25 18:30:09 +08:00
Lance Release 2fea7cd48d Bump version: 0.38.0-beta.7 → 0.38.0-beta.8 2026-08-25 06:26:35 +00:00
Wyatt Alt 0e65123bd8 fix: package ordinary @udf bodies and emit only the V1 type grammar (#4044)
Registering a real (embedding) Function failed on the client for three
reasons:

- `_package_source` treated `inspect.getclosurevars().unbound` as
  "unresolved globals"; CPython puts attribute names there, so any body
with `np.linalg.norm(...)` or `body.split()` was rejected. Module-scope
  references now come from Python's own scope analysis (`symtable`) over
  the function source, recursively, and each is resolved the way the
  interpreter would: the function's globals first (a module global may
  shadow a builtin), then builtins. Free variables of nested scopes stay
  lexical; postponed annotations are not runtime loads. A genuinely
  missing global still fails.

- `_canonical_arrow_type` emitted spellings the server's frozen grammar
  rejects (`fixed_size_list<T>[n]`, `timestamp[us]`, `struct<...>`,
  zero-sized lists). It now emits exactly the grammar, with the server's
  `fixed_size_list<item, size>` form, and the Rust declaration planner
  parses that form too.

A shared golden
(`tests/fixtures/first_class_functions/v1/arrow_types.json`)
enumerates every grammar type, nested forms and rejected spellings; the
Python emitter and Rust parser are tested against it, and the same file
is under test in sophon. Packaging tests execute the shipped artifact in
a fresh namespace.


Contract changes (hence `breaking-change`):

- `@udf` now rejects namespace acquisition structurally
(`globals()`/`eval`/... by name, plus `import
sys`/`builtins`/`importlib`/`inspect` inside the body), requires the
function's captured `__builtins__` to be the standard mapping itself
(identity, so neither lookups nor implicit hooks such as `__import__`
can differ), rejects module globals that are namespace-bearing modules
(`builtins`, `sys`, ...), and treats the function's own name as
recursion only when the module binds it to the function or to the exact
`UdfDefinition` the decorator produced; it resolves module globals
through the function's real namespace (a module global may shadow a
builtin) and ships importable classes/functions as imports.
- List outputs must declare a non-nullable, metadata-free child named
`item` (`pa.list_(pa.field("item", t, nullable=False))`); that is what
the grammar means, and pyarrow's default nullable child was being
silently collapsed into it.

Contract, stated in the `udf` docstring: the artifact is a snapshot of
the function source plus exactly the module names it references.
Reaching the module namespace by another route is rejected where a
static packager can see it and is otherwise unsupported; there is no
dynamic-access detection beyond that.
2026-08-25 14:15:07 +08:00
Jack Ye 6ed3074d4c feat: pin the base table version for data loader reads (#3982)
A permutation stores `_rowid`s, which are row addresses unless stable
row ids are enabled. Nothing in the data loader pinned a table version,
so a compaction between building a permutation and reading it can
resolve those ids to different rows.

The exposure differs by backend but exists on both:

- Remote never pins. `prepare_query_bodies` stamps `"version":
current_version()` on every request, but `current_version()` is `None`
unless `checkout` was called, so every request means "latest".
- Native pins implicitly by holding an `Arc<Dataset>` under
`ConsistencyMode::Lazy`, but `StreamingDataset.__setstate__` reopens the
table in each DataLoader worker, so each worker pins to whatever is
latest at fork time.

## Changes

`Table::at_version` returns an independent handle pinned to a version
without mutating the receiver. `checkout` cannot serve this: on remote
the version cell is an `Arc<RwLock<Option<u64>>>` shared across clones,
so pinning through it would silently pin the caller's table too.

`PermutationBuilder::build` pins for the whole build and records the
version in the permutation table's schema metadata, alongside the
existing split names. `PermutationReader` pins the base table to that
version before any take.

Because the reader pins on construction, the Python worker fork is
covered without touching the pickle format — `Permutation.__setstate__`
drops the reader and `_ensure_open` rebuilds it, which re-pins.

## Behaviour change

A permutation is now bound to the version it was built against, so rows
appended to the base table afterwards are not visible through an
existing permutation. That is the intended semantics — the permutation
only addresses rows that existed when it was built — but it is a change
worth flagging.

Permutations written before this carry no version key and read exactly
as they did before.
2026-08-24 17:16:42 -07:00
lancedb-gatefixer[bot] c1a8c3f089 fix(node): validate inferred types across records (#3786)
## Summary

- compare inferred Arrow types by their semantic representation across
records
- throw the schema inference error when a later record has an
incompatible type
- cover compatible and incompatible multi-record inference across
supported Arrow versions

## Root cause

Schema inference compared newly allocated Arrow DataType objects by
identity, so equivalent inferred types did not compare equal. The
mismatch path also constructed an Error without throwing it, which
silently accepted incompatible values.

## Validation

- pnpm test __test__/arrow.test.ts --runInBand (176 tests passed)
- pnpm lint
- pnpm build
- pnpm run docs

Fixes #3781

<!-- lance-gatekeeper-fix:v1 agent=00ec4f61a3fd82694fa4fb9bb2b37aa8
generation=1 -->

---------

Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com>
2026-08-24 16:41:05 -07:00
Will JonesandClaude Opus 5 fce45ba9fc feat(nodejs): add listTables, deprecate tableNames (#4041)
`table_names` is being replaced by `list_tables` across the SDKs, but
TypeScript only had `tableNames`. This PR adds `listTables`, which
returns a page of table names together with the token that resumes after
it, and marks `tableNames` and `TableNamesOptions` deprecated in favor
of it.

It binds the `Connection::list_tables` that already exists, so nothing
in the Rust API changes and nothing existing breaks. `pageToken` is
documented as opaque rather than as a table name, since what resumes a
listing is the database's to decide — that keeps callers off a detail
that is going to change.

Stacked on #4040, which fixes a table being dropped at every page
boundary. The page-walking test here needs that fix to pass. Review the
last commit only until #4040 lands.

## Example

```ts
const names = [];
let pageToken = undefined;
do {
  const page = await conn.listTables({ pageToken, limit: 100 });
  names.push(...page.tables);
  pageToken = page.pageToken;
} while (pageToken);
```

A namespace can be listed by passing its path first, mirroring
`tableNames`:

```ts
const page = await conn.listTables(["analytics"], { limit: 100 });
```

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-08-24 16:25:14 -07:00
lancedb-gatefixer[bot]andXuanwo 5013c176dd fix: pin remote snapshots during permutation construction (#4022)
Fixes #4015

<!-- lance-gatekeeper-fix:v1 agent=1c6407eec60f0af2336ae2dee811e1be
generation=1 -->

## Root cause

PermutationBuilder issued count_rows and the projected row-ID scan
through an unpinned RemoteTable handle. Each request could independently
resolve latest, so a concurrent table update could make the count and
scanned rows come from different snapshots.

## Fix

- add a backend hook for obtaining an independent handle pinned to the
currently selected version
- resolve latest once for remote tables while preserving an explicit
checkout and leaving the caller handle unchanged
- build the count, filtered projection, and scan from that pinned handle
while retaining native-table behavior
- add a remote mock regression that advances latest between count and
scan and covers explicit checkout preservation

## Validation

- cargo test --quiet --features remote -p lancedb --lib
test_remote_permutation_builder_pins_snapshot
- cargo test --quiet --features remote -p lancedb --lib
dataloader::permutation::builder::tests
- cargo check --quiet --features remote --tests --examples
- cargo clippy --quiet --features remote --tests --examples
- cargo fmt --all -- --check

Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com>
Co-authored-by: Xuanwo <github@xuanwo.io>
2026-08-25 07:14:33 +08:00
Lance Release 71f85a8d9f Bump version: 0.38.0-beta.6 → 0.38.0-beta.7 2026-08-24 21:09:12 +00:00
Wyatt Alt c72f5b2960 feat: bind a materialized view refresh to the view incarnation (#4043)
A caller that queues a refresh and executes it later can only tell the
view it captured from a drop-and-recreate by comparing the definition
and
version. A recreated view with the same definition and an equal or
higher
version passes that check, and the check runs before the refresh reloads
the view, so it never sees the state it commits against.

This mints an `mv.incarnation` token in the schema metadata at each
physical creation of a view table.
`RefreshMaterializedViewBuilder::expect_incarnation` carries the
captured
token into the refresh, which compares it against the latest stored
manifest before planning and again immediately before each commit
(publish, fragment swap, watermark stamp), refusing to land in a
different
incarnation. A view with no token -- declared before tokens existed, or
its metadata replaced wholesale -- is refused under a bound refresh with
its own wording and is minted one by its next unbound refresh. The token
is exposed through `MaterializedView::incarnation`; refreshes without an
expectation are unchanged.

This is best effort: the token is not part of lance's commit condition,
so a recreation landing between the final pre-commit read and the commit
itself is not caught. Closing that window needs a base-manifest
precondition in lance's commit path.
2026-08-24 14:06:40 -07:00
Will JonesandClaude Opus 5 93f47b8aab fix(remote): stop table_names inventing a page token for a namespace (#4039)
`table_names` paginates using a `start_after` table name. This works for
the `/v1/table` endpoint, which guarantees table-order. But the
`/v1/namespace/{id}/table/list` does not. We change that caller to
instead collect all table names, sort, and apply the pagination locally.

We are deprecating this API, so this is just an interim fix. For good
performance, users should move to the `list_tables` API instead, which
uses opaque tokens that don't rely on lexical sorting.

---------

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-08-24 13:35:31 -07:00
lancedb-gatefixer[bot]andXuanwo 105fd73bc6 fix(python): commit streaming worker checkpoints on consumption (#4023)
## Summary

- add `StreamingDataLoader`, which transports worker snapshots with
prefetched batches and commits them to the parent dataset only when the
trainer receives each batch
- preserve exact non-uniform per-split progress and resume lagging
splits without replaying already-consumed rows
- reject stale parent checkpoints after a standard multi-process
`DataLoader` has started, with guidance to use the consumer-aware loader
- document the new public loader and merge non-uniform state across
ranks

## Root cause

PyTorch runs `StreamingDataset.__iter__` in private worker-process
copies, while callers invoke `state_dict()` on the parent dataset.
Sharing producer counters would still be incorrect because DataLoader
prefetch can advance workers beyond batches returned to the trainer.

## Validation

- `uv run --extra tests pytest python/tests/test_elastic_dataloader.py
-q` (154 passed)
- focused non-uniform merge regression (1 passed)
- `uv run --project python --extra tests --extra dev ruff format .`
- `uv run --project python --extra tests --extra dev ruff check .`
- `cd docs && PYTHONPATH=. ../python/.venv/bin/mkdocs build`

Fixes #3967

<!-- lance-gatekeeper-fix:v1 agent=572be272619660b97e87fd5c85188341
generation=1 -->

---------

Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com>
Co-authored-by: Xuanwo <github@xuanwo.io>
2026-08-25 04:22:22 +08:00
Will JonesandClaude Opus 5 94d484f539 fix(listing): don't drop a table at a page boundary (#4040)
Listing tables a page at a time against a local database silently
skipped one table at every page boundary. `ListingDatabase::list_tables`
returned the first name of the *next* page as that page's token, but
resuming from a token drops every name at or before it — so the table
the token named was never handed to the caller. Walking `[a, b, c, d,
e]` with a limit of 2 returned `[a, b, d, e]`.

This PR returns the last name of the page as the token instead, which is
what resuming after the token expects.

This is reachable from Python today through
`db.list_tables(page_token=...)` on a local connection; it also affects
`len(db)` and `name in db`, which walk the pages. Remote and
namespace-backed connections page on the server and were never affected.

## Example

```python
db = lancedb.connect(tmp_path)
for name in ["a", "b", "c", "d", "e"]:
    db.create_table(name, [{"id": 1}])

names, token = [], None
while True:
    page = db.list_tables(page_token=token, limit=2)
    names += page.tables
    token = page.page_token
    if not token:
        break

# before: ['a', 'b', 'd', 'e']
# after:  ['a', 'b', 'c', 'd', 'e']
```

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-08-24 12:00:48 -07:00
65 changed files with 4566 additions and 647 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
[tool.bumpversion]
current_version = "0.38.0-beta.6"
current_version = "0.38.0-beta.9"
parse = """(?x)
(?P<major>0|[1-9]\\d*)\\.
(?P<minor>0|[1-9]\\d*)\\.
Generated
+45 -45
View File
@@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
[[package]]
name = "fsst"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"rand 0.9.5",
@@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a"
[[package]]
name = "lance"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arc-swap",
"arrow",
@@ -4888,8 +4888,8 @@ dependencies = [
[[package]]
name = "lance-arrow"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4911,7 +4911,7 @@ dependencies = [
[[package]]
name = "lance-arrow-scalar"
version = "58.0.0"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4925,7 +4925,7 @@ dependencies = [
[[package]]
name = "lance-arrow-stats"
version = "58.0.0"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -4934,8 +4934,8 @@ dependencies = [
[[package]]
name = "lance-bitpacking"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrayref",
"crunchy",
@@ -4945,8 +4945,8 @@ dependencies = [
[[package]]
name = "lance-core"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -4983,8 +4983,8 @@ dependencies = [
[[package]]
name = "lance-datafusion"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow",
"arrow-array",
@@ -5013,8 +5013,8 @@ dependencies = [
[[package]]
name = "lance-datagen"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow",
"arrow-array",
@@ -5031,8 +5031,8 @@ dependencies = [
[[package]]
name = "lance-derive"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"proc-macro2",
"quote",
@@ -5041,8 +5041,8 @@ dependencies = [
[[package]]
name = "lance-encoding"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-arith",
"arrow-array",
@@ -5075,8 +5075,8 @@ dependencies = [
[[package]]
name = "lance-file"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-arith",
"arrow-array",
@@ -5107,8 +5107,8 @@ dependencies = [
[[package]]
name = "lance-index"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arc-swap",
"arrow",
@@ -5172,8 +5172,8 @@ dependencies = [
[[package]]
name = "lance-index-core"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -5195,8 +5195,8 @@ dependencies = [
[[package]]
name = "lance-io"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow",
"arrow-array",
@@ -5236,8 +5236,8 @@ dependencies = [
[[package]]
name = "lance-linalg"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -5251,8 +5251,8 @@ dependencies = [
[[package]]
name = "lance-namespace"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow",
"async-trait",
@@ -5264,8 +5264,8 @@ dependencies = [
[[package]]
name = "lance-namespace-impls"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow",
"arrow-ipc",
@@ -5318,8 +5318,8 @@ dependencies = [
[[package]]
name = "lance-select"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -5333,8 +5333,8 @@ dependencies = [
[[package]]
name = "lance-table"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow",
"arrow-array",
@@ -5374,8 +5374,8 @@ dependencies = [
[[package]]
name = "lance-testing"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"arrow-array",
"arrow-schema",
@@ -5388,8 +5388,8 @@ dependencies = [
[[package]]
name = "lance-tokenizer"
version = "11.0.0-rc.1"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-rc.1#308c00db4f4dae6b2ea4a1d1d2ceb1c3ed88a159"
version = "11.0.0-beta.22"
source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20"
dependencies = [
"frostem",
"icu_segmenter",
@@ -5402,7 +5402,7 @@ dependencies = [
[[package]]
name = "lancedb"
version = "0.38.0-beta.6"
version = "0.38.0-beta.8"
dependencies = [
"ahash",
"anyhow",
@@ -5490,7 +5490,7 @@ dependencies = [
[[package]]
name = "lancedb-nodejs"
version = "0.38.0-beta.6"
version = "0.38.0-beta.8"
dependencies = [
"arrow-array",
"arrow-buffer",
@@ -5515,7 +5515,7 @@ dependencies = [
[[package]]
name = "lancedb-python"
version = "0.38.0-beta.6"
version = "0.38.0-beta.8"
dependencies = [
"arrow",
"async-trait",
+14 -14
View File
@@ -13,20 +13,20 @@ categories = ["database-implementations"]
rust-version = "1.91.0"
[workspace.dependencies]
lance = { "version" = "=11.0.0-rc.1", default-features = false, "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-core = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-datagen = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-file = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-io = { "version" = "=11.0.0-rc.1", default-features = false, "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-index = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-linalg = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace-impls = { "version" = "=11.0.0-rc.1", default-features = false, "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-table = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-testing = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-datafusion = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-encoding = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance-arrow = { "version" = "=11.0.0-rc.1", "tag" = "v11.0.0-rc.1", "git" = "https://github.com/lance-format/lance.git" }
lance = { "version" = "=11.0.0-beta.22", default-features = false, "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-core = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-datagen = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-file = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-io = { "version" = "=11.0.0-beta.22", default-features = false, "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-index = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-linalg = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-namespace-impls = { "version" = "=11.0.0-beta.22", default-features = false, "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-table = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-testing = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-datafusion = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-encoding = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lance-arrow = { "version" = "=11.0.0-beta.22", "tag" = "v11.0.0-beta.22", "git" = "https://github.com/lance-format/lance.git" }
lancedb = { path = "rust/lancedb", default-features = false }
ahash = "0.8"
# Note that this one does not include pyarrow
+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.38.0-beta.6</version>
<version>0.38.0-beta.9</version>
</dependency>
```
+73 -1
View File
@@ -584,6 +584,70 @@ Child namespace names and
***
### listTables()
#### listTables(options)
```ts
abstract listTables(options?): Promise<ListTablesResponse>
```
List a page of the tables in this database.
To retrieve the tables after the page, pass the `pageToken` the response
carries back in. A page can be shorter than `limit` without being the last
one, so walk until a response carries no page token:
```ts
const names = [];
let pageToken = undefined;
do {
const page = await conn.listTables({ pageToken, limit: 100 });
names.push(...page.tables);
pageToken = page.pageToken;
} while (pageToken);
```
##### Parameters
* **options?**: `Partial`&lt;[`ListTablesOptions`](../interfaces/ListTablesOptions.md)&gt;
Pagination options
(`pageToken`, `limit`).
##### Returns
`Promise`&lt;[`ListTablesResponse`](../interfaces/ListTablesResponse.md)&gt;
A page of table names and an
optional token for the tables after it.
#### listTables(namespacePath, options)
```ts
abstract listTables(namespacePath?, options?): Promise<ListTablesResponse>
```
List a page of the tables in this database.
##### Parameters
* **namespacePath?**: `string`[]
The namespace path to list tables from
(defaults to root namespace)
* **options?**: `Partial`&lt;[`ListTablesOptions`](../interfaces/ListTablesOptions.md)&gt;
Pagination options
(`pageToken`, `limit`).
##### Returns
`Promise`&lt;[`ListTablesResponse`](../interfaces/ListTablesResponse.md)&gt;
A page of table names and an
optional token for the tables after it.
***
### openMaterializedView()
```ts
@@ -660,7 +724,7 @@ a "not supported" error.
***
### tableNames()
### ~~tableNames()~~
#### tableNames(options)
@@ -682,6 +746,10 @@ Tables will be returned in lexicographical order.
`Promise`&lt;`string`[]&gt;
##### Deprecated
Use [Connection.listTables](Connection.md#listtables) instead.
#### tableNames(namespacePath, options)
```ts
@@ -704,3 +772,7 @@ Tables will be returned in lexicographical order.
##### Returns
`Promise`&lt;`string`[]&gt;
##### Deprecated
Use [Connection.listTables](Connection.md#listtables) instead.
+2
View File
@@ -100,6 +100,8 @@
- [JobInfo](interfaces/JobInfo.md)
- [ListNamespacesOptions](interfaces/ListNamespacesOptions.md)
- [ListNamespacesResponse](interfaces/ListNamespacesResponse.md)
- [ListTablesOptions](interfaces/ListTablesOptions.md)
- [ListTablesResponse](interfaces/ListTablesResponse.md)
- [LsmStats](interfaces/LsmStats.md)
- [LsmWriteSpec](interfaces/LsmWriteSpec.md)
- [MaterializedViewDefinition](interfaces/MaterializedViewDefinition.md)
@@ -0,0 +1,34 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / ListTablesOptions
# Interface: ListTablesOptions
## Properties
### limit?
```ts
optional limit: number;
```
An upper bound on how many tables to return.
A page may hold fewer than this and still not be the last one, so keep
going while the response carries a page token rather than while pages are
full.
***
### pageToken?
```ts
optional pageToken: string;
```
Token from a previous response, to resume listing where it left off.
The token is opaque: it carries whatever the database needs to resume, and
callers should not construct or interpret one.
@@ -0,0 +1,23 @@
[**@lancedb/lancedb**](../README.md) • **Docs**
***
[@lancedb/lancedb](../globals.md) / ListTablesResponse
# Interface: ListTablesResponse
## Properties
### pageToken?
```ts
optional pageToken: string;
```
***
### tables
```ts
tables: string[];
```
+8 -3
View File
@@ -4,11 +4,16 @@
[@lancedb/lancedb](../globals.md) / TableNamesOptions
# Interface: TableNamesOptions
# Interface: ~~TableNamesOptions~~
## Deprecated
Use [ListTablesOptions](ListTablesOptions.md) with [Connection.listTables](../classes/Connection.md#listtables)
instead.
## Properties
### limit?
### ~~limit?~~
```ts
optional limit: number;
@@ -18,7 +23,7 @@ An optional limit to the number of results to return.
***
### startAfter?
### ~~startAfter?~~
```ts
optional startAfter: string;
+2
View File
@@ -261,6 +261,8 @@ instead of being materialized with the rest of the row.
::: lancedb.streaming.StreamingDataset
::: lancedb.streaming.StreamingDataLoader
::: lancedb.permutation.permutation_builder
::: lancedb.permutation.PermutationBuilder
+1 -1
View File
@@ -8,7 +8,7 @@
<parent>
<groupId>com.lancedb</groupId>
<artifactId>lancedb-parent</artifactId>
<version>0.38.0-beta.6</version>
<version>0.38.0-beta.9</version>
<relativePath>../pom.xml</relativePath>
</parent>
+2 -2
View File
@@ -6,7 +6,7 @@
<groupId>com.lancedb</groupId>
<artifactId>lancedb-parent</artifactId>
<version>0.38.0-beta.6</version>
<version>0.38.0-beta.9</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-rc.1</lance-core.version>
<lance-core.version>11.0.0-beta.22</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.38.0-beta.6"
version = "0.38.0-beta.9"
publish = false
license.workspace = true
description.workspace = true
+131
View File
@@ -515,6 +515,137 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
);
});
it("will allow matching inferred types across records", function () {
expect(() =>
makeArrowTable([{ value: 1 }, { value: 2 }]),
).not.toThrow();
});
it("will reject mismatched inferred types across records", function () {
expect(() => makeArrowTable([{ value: 1 }, { value: "two" }])).toThrow(
"Failed to infer schema for data. Previously inferred type Float64 but found Utf8 for field value at row 1. Consider providing an explicit schema.",
);
});
it("will ignore generated dictionary IDs when comparing inferred types", function () {
const table = makeArrowTable([{ str: "a" }, { str: "b" }], {
dictionaryEncodeStrings: true,
});
expect(table.getChild("str")?.toJSON()).toEqual(["a", "b"]);
});
it("will preserve null values without treating them as type mismatches", function () {
for (const records of [
[{ vector: [1, 2, 3] }, { vector: null }],
[{ vector: null }, { vector: [1, 2, 3] }],
]) {
const table = makeArrowTable(records);
expect(table.numRows).toBe(2);
expect(table.getChild("vector")?.nullCount).toBe(1);
}
});
it("will preserve empty variable-size lists", function () {
for (const records of [
[{ items: [1] }, { items: [] }],
[{ items: [] }, { items: [1] }],
]) {
const table = makeArrowTable(records);
expect(
table
.getChild("items")
?.toJSON()
.map((value) => value.toJSON()),
).toEqual(records.map((record) => record.items));
}
});
it("will propagate deferred evidence through nested lists", function () {
for (const records of [
[{ items: [1] }, { items: [null] }],
[{ items: [null] }, { items: [1] }],
[{ items: [null, 1] }, { items: [2, null] }],
]) {
const table = makeArrowTable(records);
expect(
table
.getChild("items")
?.toJSON()
.map((value) => value.toJSON()),
).toEqual(records.map((record) => record.items));
}
const nestedRecords = [{ items: [[1]] }, { items: [[null]] }];
const nestedTable = makeArrowTable(nestedRecords);
expect(
nestedTable
.getChild("items")
?.toJSON()
.map((value) =>
value
.toJSON()
.map((nestedValue: { toJSON: () => unknown[] }) =>
nestedValue.toJSON(),
),
),
).toEqual(nestedRecords.map((record) => record.items));
});
it("will reject incompatible deferred evidence within a list", function () {
for (const items of [
[[], 1],
[1, []],
[[null], 1],
[1, [null]],
]) {
expect(() => makeArrowTable([{ items }])).toThrow(
"Failed to infer data type for field items at row 0.",
);
}
});
it("will reject empty fixed-size lists", function () {
expect(() =>
makeArrowTable([{ vector: [1, 2, 3] }, { vector: [] }]),
).toThrow(
"Failed to infer schema for data. Previously inferred type FixedSizeList[3]<Float32> but found List[0] for field vector at row 1.",
);
});
it("will reject inferred leaf and branch shape changes", function () {
expect(() =>
makeArrowTable([{ value: 1 }, { value: { nested: 2 } }]),
).toThrow(
"Failed to infer schema for data. Previously inferred type Float64 but found Struct for field value at row 1.",
);
expect(() =>
makeArrowTable([{ value: { nested: 1 } }, { value: 2 }]),
).toThrow(
"Failed to infer schema for data. Previously inferred type Struct but found Float64 for field value at row 1.",
);
});
it("will allow null values around inferred struct values", function () {
for (const { records, nullIndex } of [
{
records: [{ value: null }, { value: { nested: 2 } }],
nullIndex: 0,
},
{
records: [{ value: { nested: 1 } }, { value: null }],
nullIndex: 1,
},
]) {
const table = makeArrowTable(records);
const values = table.getChild("value");
expect(values?.nullCount).toBe(1);
expect(values?.get(nullIndex)).toBeNull();
}
});
it("will allow a schema to be provided", async function () {
await checkTableCreation(
async (records, _, schema) =>
+68 -1
View File
@@ -4,7 +4,13 @@
import { readdirSync } from "fs";
import { Field, Float64, Schema } from "apache-arrow";
import * as tmp from "tmp";
import { Connection, Table, connect, connectNamespace } from "../lancedb";
import {
Connection,
ListTablesResponse,
Table,
connect,
connectNamespace,
} from "../lancedb";
import { LocalTable } from "../lancedb/table";
describe("when connecting", () => {
@@ -47,6 +53,7 @@ describe("given a connection", () => {
await db.close();
expect(db.isOpen()).toBe(false);
await expect(db.tableNames()).rejects.toThrow("Connection is closed");
await expect(db.listTables()).rejects.toThrow("Connection is closed");
await expect(db.renameTable("a", "b")).rejects.toThrow(
"Connection is closed",
);
@@ -129,6 +136,66 @@ describe("given a connection", () => {
expect(tables).toEqual(["b", "c"]);
});
it("should respect limit and page token when listing tables", async () => {
const db = await connect(tmpDir.name);
await db.createTable("b", [{ id: 1 }]);
await db.createTable("a", [{ id: 1 }]);
await db.createTable("c", [{ id: 1 }]);
const all = await db.listTables();
expect(all.tables).toEqual(["a", "b", "c"]);
expect(all.pageToken).toBeUndefined();
const first = await db.listTables({ limit: 1 });
expect(first.tables).toEqual(["a"]);
expect(first.pageToken).toBeDefined();
const second = await db.listTables({
limit: 1,
pageToken: first.pageToken,
});
expect(second.tables).toEqual(["b"]);
});
it("should visit every table exactly once when walking pages", async () => {
const db = await connect(tmpDir.name);
const created = ["a", "b", "c", "d", "e"];
for (const name of created) {
await db.createTable(name, [{ id: 1 }]);
}
const seen: string[] = [];
let pageToken: string | undefined = undefined;
do {
const page: ListTablesResponse = await db.listTables({
limit: 2,
pageToken,
});
seen.push(...page.tables);
pageToken = page.pageToken;
} while (pageToken);
expect(seen).toEqual(created);
});
it("should list tables in a namespace", async () => {
const db = await connect(tmpDir.name, {
// biome-ignore lint/style/useNamingConvention: opaque backend property key, must match Rust
namespaceClientProperties: { manifest_enabled: "true" },
});
await db.createNamespace(["child"]);
await db.createTable("nested", [{ id: 1 }], ["child"]);
await expect(db.listTables(["child"])).resolves.toEqual(
expect.objectContaining({ tables: ["nested"] }),
);
await expect(db.listTables()).resolves.toEqual(
expect.objectContaining({ tables: [] }),
);
});
it("should create tables in v2 mode", async () => {
const db = await connect(tmpDir.name);
const data = [...Array(10000).keys()].map((i) => ({ id: i }));
+33 -307
View File
@@ -5,7 +5,6 @@ import {
Data as ArrowData,
Table as ArrowTable,
Binary,
Bool,
BufferType,
DataType,
DateUnit,
@@ -18,12 +17,7 @@ import {
FixedSizeList,
Float,
Float32,
Float64,
Int,
Int8,
Int16,
Int32,
Int64,
LargeBinary,
List,
Null,
@@ -36,17 +30,16 @@ import {
Struct,
Timestamp,
Type,
Uint8,
Uint16,
Uint32,
Utf8,
Vector,
makeVector as arrowMakeVector,
util as arrowUtil,
vectorFromArray as badVectorFromArray,
makeBuilder,
makeData,
} from "apache-arrow";
import { Buffers } from "apache-arrow/data";
import { typedArrayToArrowType } from "./arrow_type";
import { type EmbeddingFunction } from "./embedding/embedding_function";
import {
EmbeddingFunctionConfig,
@@ -59,14 +52,7 @@ import {
sanitizeTable,
sanitizeType,
} from "./sanitize";
/**
* Check if a field name indicates a vector column.
*/
function nameSuggestsVectorColumn(fieldName: string): boolean {
const nameLower = fieldName.toLowerCase();
return nameLower.includes("vector") || nameLower.includes("embedding");
}
import { inferSchema } from "./schema";
export * from "apache-arrow";
export type SchemaLike =
@@ -459,110 +445,6 @@ export function makeArrowTable(
return new ArrowTable(inferredSchema, finalColumns);
}
function inferSchema(
data: Array<Record<string, unknown>>,
schema: Schema | undefined,
opts: MakeArrowTableOptions,
): Schema {
// We will collect all fields we see in the data.
const pathTree = new PathTree<DataType>();
for (const [rowI, row] of data.entries()) {
for (const [path, value] of rowPathsAndValues(row)) {
if (!pathTree.has(path)) {
// First time seeing this field.
if (schema !== undefined) {
const field = getFieldForPath(schema, path);
if (field === undefined) {
throw new Error(
`Found field not in schema: ${path.join(".")} at row ${rowI}`,
);
} else {
pathTree.set(path, field.type);
}
} else {
const inferredType = inferType(value, path, opts);
if (inferredType === undefined) {
throw new Error(`Failed to infer data type for field ${path.join(
".",
)} at row ${rowI}. \
Consider providing an explicit schema.`);
}
pathTree.set(path, inferredType);
}
} else if (schema === undefined) {
const currentType = pathTree.get(path);
const newType = inferType(value, path, opts);
if (currentType !== newType) {
new Error(`Failed to infer schema for data. Previously inferred type \
${currentType} but found ${newType} at row ${rowI}. Consider \
providing an explicit schema.`);
}
}
}
}
if (schema === undefined) {
function fieldsFromPathTree(pathTree: PathTree<DataType>): Field[] {
const fields = [];
for (const [name, value] of pathTree.map.entries()) {
if (value instanceof PathTree) {
const children = fieldsFromPathTree(value);
fields.push(new Field(name, new Struct(children), true));
} else {
fields.push(new Field(name, value, true));
}
}
return fields;
}
const fields = fieldsFromPathTree(pathTree);
return new Schema(fields);
} else {
function takeMatchingFields(
fields: Field[],
pathTree: PathTree<DataType>,
): Field[] {
const outFields = [];
for (const field of fields) {
if (pathTree.map.has(field.name)) {
const value = pathTree.get([field.name]);
if (value instanceof PathTree) {
const struct = field.type as Struct;
const children = takeMatchingFields(struct.children, value);
outFields.push(
new Field(field.name, new Struct(children), field.nullable),
);
} else {
outFields.push(
new Field(field.name, value as DataType, field.nullable),
);
}
}
}
return outFields;
}
const fields = takeMatchingFields(schema.fields, pathTree);
return new Schema(fields);
}
}
function* rowPathsAndValues(
row: Record<string, unknown>,
basePath: string[] = [],
): Generator<[string[], unknown]> {
for (const [key, value] of Object.entries(row)) {
if (isObject(value)) {
yield* rowPathsAndValues(value, [...basePath, key]);
} else {
// Skip undefined values - they should be treated the same as missing fields
// for embedding function purposes
if (value !== undefined) {
yield [[...basePath, key], value];
}
}
}
}
function isObject(value: unknown): value is Record<string, unknown> {
return (
typeof value === "object" &&
@@ -577,146 +459,19 @@ function isObject(value: unknown): value is Record<string, unknown> {
);
}
function getFieldForPath(schema: Schema, path: string[]): Field | undefined {
let current: Field | Schema = schema;
function valueAtPath(datum: Record<string, unknown>, path: string[]): unknown {
let current: unknown = datum;
for (const key of path) {
if (current instanceof Schema) {
const field: Field | undefined = current.fields.find(
(f) => f.name === key,
);
if (field === undefined) {
return undefined;
}
current = field;
} else if (current instanceof Field && DataType.isStruct(current.type)) {
const struct: Struct = current.type;
const field = struct.children.find((f) => f.name === key);
if (field === undefined) {
return undefined;
}
current = field;
if (current == null) {
return null;
}
if (isObject(current) && (Object.hasOwn(current, key) || key in current)) {
current = current[key];
} else {
return undefined;
}
}
if (current instanceof Field) {
return current;
} else {
return undefined;
}
}
/**
* Try to infer which Arrow type to use for a given value.
*
* May return undefined if the type cannot be inferred.
*/
function inferType(
value: unknown,
path: string[],
opts: MakeArrowTableOptions,
): DataType | undefined {
if (typeof value === "bigint") {
return new Int64();
} else if (typeof value === "number") {
// Even if it's an integer, it's safer to assume Float64. Users can
// always provide an explicit schema or use BigInt if they mean integer.
return new Float64();
} else if (typeof value === "string") {
if (opts.dictionaryEncodeStrings) {
return new Dictionary(new Utf8(), new Int32());
} else {
return new Utf8();
}
} else if (typeof value === "boolean") {
return new Bool();
} else if (value instanceof Buffer) {
return new Binary();
} else if (ArrayBuffer.isView(value) && !(value instanceof DataView)) {
const info = typedArrayToArrowType(value);
if (info !== undefined) {
const child = new Field("item", info.elementType, true);
return new FixedSizeList(info.length, child);
}
return undefined;
} else if (Array.isArray(value)) {
if (value.length === 0) {
return undefined; // Without any values we can't infer the type
}
if (path.length === 1 && Object.hasOwn(opts.vectorColumns, path[0])) {
const floatType = sanitizeType(opts.vectorColumns[path[0]].type);
return new FixedSizeList(
value.length,
new Field("item", floatType, true),
);
}
const valueType = inferType(value[0], path, opts);
if (valueType === undefined) {
return undefined;
}
// Try to automatically detect embedding columns.
if (nameSuggestsVectorColumn(path[path.length - 1])) {
// Check if value is a Uint8Array for integer vector type determination
if (value instanceof Uint8Array) {
// For integer vectors, we default to Uint8 (matching Python implementation)
const child = new Field("item", new Uint8(), true);
return new FixedSizeList(value.length, child);
} else {
// For float vectors, we default to Float32
const child = new Field("item", new Float32(), true);
return new FixedSizeList(value.length, child);
}
} else {
const child = new Field("item", valueType, true);
return new List(child);
}
} else {
// TODO: timestamp
return undefined;
}
}
class PathTree<V> {
map: Map<string, V | PathTree<V>>;
constructor(entries?: [string[], V][]) {
this.map = new Map();
if (entries !== undefined) {
for (const [path, value] of entries) {
this.set(path, value);
}
}
}
has(path: string[]): boolean {
let ref: PathTree<V> = this;
for (const part of path) {
if (!(ref instanceof PathTree) || !ref.map.has(part)) {
return false;
}
ref = ref.map.get(part) as PathTree<V>;
}
return true;
}
get(path: string[]): V | undefined {
let ref: PathTree<V> = this;
for (const part of path) {
if (!(ref instanceof PathTree) || !ref.map.has(part)) {
return undefined;
}
ref = ref.map.get(part) as PathTree<V>;
}
return ref as V;
}
set(path: string[], value: V): void {
let ref: PathTree<V> = this;
for (const part of path.slice(0, path.length - 1)) {
if (!ref.map.has(part)) {
ref.map.set(part, new PathTree<V>());
}
ref = ref.map.get(part) as PathTree<V>;
}
ref.map.set(path[path.length - 1], value);
}
return current;
}
function transposeData(
@@ -724,37 +479,26 @@ function transposeData(
field: Field,
path: string[] = [],
): Vector {
const valuesPath = [...path, field.name];
const values = data.map((datum) => valueAtPath(datum, valuesPath));
if (field.type instanceof Struct) {
const childFields = field.type.children;
const fullPath = [...path, field.name];
const childVectors = childFields.map((child) => {
return transposeData(data, child, fullPath);
return transposeData(data, child, valuesPath);
});
const nullCount = values.filter((value) => value === null).length;
const structData = makeData({
type: field.type,
length: values.length,
nullCount,
nullBitmap:
nullCount > 0
? arrowUtil.packBools(values.map((value) => value !== null))
: undefined,
children: childVectors as unknown as ArrowData<DataType>[],
});
return arrowMakeVector(structData);
} else {
const valuesPath = [...path, field.name];
const values = data.map((datum) => {
let current: unknown = datum;
for (const key of valuesPath) {
if (current == null) {
return null;
}
if (
isObject(current) &&
(Object.hasOwn(current, key) || key in current)
) {
current = current[key];
} else {
return null;
}
}
return current;
});
return makeVector(values, field.type, undefined, field.nullable);
}
}
@@ -797,32 +541,6 @@ function makeListVector(lists: unknown[][]): Vector<unknown> {
return listBuilder.finish().toVector();
}
/**
* Map a JS TypedArray instance to the corresponding Arrow element DataType
* and its length. Returns undefined if the value is not a recognized TypedArray.
*/
function typedArrayToArrowType(
value: ArrayBufferView,
): { elementType: DataType; length: number } | undefined {
if (value instanceof Float32Array)
return { elementType: new Float32(), length: value.length };
if (value instanceof Float64Array)
return { elementType: new Float64(), length: value.length };
if (value instanceof Uint8Array)
return { elementType: new Uint8(), length: value.length };
if (value instanceof Uint16Array)
return { elementType: new Uint16(), length: value.length };
if (value instanceof Uint32Array)
return { elementType: new Uint32(), length: value.length };
if (value instanceof Int8Array)
return { elementType: new Int8(), length: value.length };
if (value instanceof Int16Array)
return { elementType: new Int16(), length: value.length };
if (value instanceof Int32Array)
return { elementType: new Int32(), length: value.length };
return undefined;
}
/** Helper function to convert an Array of JS values to an Arrow Vector */
function makeVector(
values: unknown[],
@@ -1462,8 +1180,12 @@ export function ensureNestedFieldsExist(
completeRow[field.name] = row[field.name];
}
} else {
// Field is missing from the data - set to null
completeRow[field.name] = null;
// Keep a missing struct valid while filling each of its children with
// null. This is distinct from an explicitly null struct value.
completeRow[field.name] =
field.type.constructor.name === "Struct"
? ensureStructFieldsExist({}, field.type as Struct)
: null;
}
}
@@ -1498,8 +1220,12 @@ function ensureStructFieldsExist(
completeStruct[childField.name] = data[childField.name];
}
} else {
// Field is missing - set to null
completeStruct[childField.name] = null;
// Keep a missing struct valid while filling each of its children with
// null. This is distinct from an explicitly null struct value.
completeStruct[childField.name] =
childField.type.constructor.name === "Struct"
? ensureStructFieldsExist({}, childField.type as Struct)
: null;
}
}
+40
View File
@@ -0,0 +1,40 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
import {
type DataType,
Float32,
Float64,
Int8,
Int16,
Int32,
Uint8,
Uint16,
Uint32,
} from "apache-arrow";
/**
* Map a JS TypedArray instance to the corresponding Arrow element type and
* length. Returns undefined when the view is not a supported TypedArray.
*/
export function typedArrayToArrowType(
value: ArrayBufferView,
): { elementType: DataType; length: number } | undefined {
if (value instanceof Float32Array)
return { elementType: new Float32(), length: value.length };
if (value instanceof Float64Array)
return { elementType: new Float64(), length: value.length };
if (value instanceof Uint8Array)
return { elementType: new Uint8(), length: value.length };
if (value instanceof Uint16Array)
return { elementType: new Uint16(), length: value.length };
if (value instanceof Uint32Array)
return { elementType: new Uint32(), length: value.length };
if (value instanceof Int8Array)
return { elementType: new Int8(), length: value.length };
if (value instanceof Int16Array)
return { elementType: new Int16(), length: value.length };
if (value instanceof Int32Array)
return { elementType: new Int32(), length: value.length };
return undefined;
}
+85
View File
@@ -31,12 +31,14 @@ import type {
JobDescription,
JobInfo,
ListNamespacesResponse,
ListTablesResponse,
} from "./native";
export type {
CreateNamespaceResponse,
DescribeNamespaceResponse,
DropNamespaceResponse,
ListNamespacesResponse,
ListTablesResponse,
};
import { sanitizeTable } from "./sanitize";
import { LocalTable, Table } from "./table";
@@ -134,6 +136,10 @@ export interface OpenTableOptions {
indexCacheSize?: number;
}
/**
* @deprecated Use {@link ListTablesOptions} with {@link Connection.listTables}
* instead.
*/
export interface TableNamesOptions {
/**
* If present, only return names that come lexicographically after the
@@ -147,6 +153,24 @@ export interface TableNamesOptions {
limit?: number;
}
export interface ListTablesOptions {
/**
* Token from a previous response, to resume listing where it left off.
*
* The token is opaque: it carries whatever the database needs to resume, and
* callers should not construct or interpret one.
*/
pageToken?: string;
/**
* An upper bound on how many tables to return.
*
* A page may hold fewer than this and still not be the last one, so keep
* going while the response carries a page token rather than while pages are
* full.
*/
limit?: number;
}
export interface ListNamespacesOptions {
/** Token from a previous response for pagination. */
pageToken?: string;
@@ -231,6 +255,7 @@ export abstract class Connection {
* @param {Partial<TableNamesOptions>} options - options to control the
* paging / start point (backwards compatibility)
*
* @deprecated Use {@link Connection.listTables} instead.
*/
abstract tableNames(options?: Partial<TableNamesOptions>): Promise<string[]>;
/**
@@ -241,12 +266,53 @@ export abstract class Connection {
* @param {Partial<TableNamesOptions>} options - options to control the
* paging / start point
*
* @deprecated Use {@link Connection.listTables} instead.
*/
abstract tableNames(
namespacePath?: string[],
options?: Partial<TableNamesOptions>,
): Promise<string[]>;
/**
* List a page of the tables in this database.
*
* To retrieve the tables after the page, pass the `pageToken` the response
* carries back in. A page can be shorter than `limit` without being the last
* one, so walk until a response carries no page token:
*
* ```ts
* const names = [];
* let pageToken = undefined;
* do {
* const page = await conn.listTables({ pageToken, limit: 100 });
* names.push(...page.tables);
* pageToken = page.pageToken;
* } while (pageToken);
* ```
*
* @param {Partial<ListTablesOptions>} options - Pagination options
* (`pageToken`, `limit`).
* @returns {Promise<ListTablesResponse>} A page of table names and an
* optional token for the tables after it.
*/
abstract listTables(
options?: Partial<ListTablesOptions>,
): Promise<ListTablesResponse>;
/**
* List a page of the tables in this database.
*
* @param {string[]} namespacePath - The namespace path to list tables from
* (defaults to root namespace)
* @param {Partial<ListTablesOptions>} options - Pagination options
* (`pageToken`, `limit`).
* @returns {Promise<ListTablesResponse>} A page of table names and an
* optional token for the tables after it.
*/
abstract listTables(
namespacePath?: string[],
options?: Partial<ListTablesOptions>,
): Promise<ListTablesResponse>;
/**
* Open a table in the database.
* @param {string} name - The name of the table
@@ -601,6 +667,25 @@ export class LocalConnection extends Connection {
return await this.inner.listMaterializedViews();
}
async listTables(
namespacePathOrOptions?: string[] | Partial<ListTablesOptions>,
options?: Partial<ListTablesOptions>,
): Promise<ListTablesResponse> {
// Detect if first argument is namespacePath array or options object
const namespacePath = Array.isArray(namespacePathOrOptions)
? namespacePathOrOptions
: undefined;
const listTablesOptions = Array.isArray(namespacePathOrOptions)
? options
: namespacePathOrOptions;
return this.inner.listTables(
namespacePath ?? [],
listTablesOptions?.pageToken,
listTablesOptions?.limit,
);
}
async openTable(
name: string,
namespacePath?: string[],
+2
View File
@@ -81,11 +81,13 @@ export {
Connection,
CreateTableOptions,
TableNamesOptions,
ListTablesOptions,
OpenTableOptions,
ListNamespacesOptions,
CreateNamespaceOptions,
DropNamespaceOptions,
ListNamespacesResponse,
ListTablesResponse,
CreateNamespaceResponse,
DropNamespaceResponse,
DescribeNamespaceResponse,
+566
View File
@@ -0,0 +1,566 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
import {
Binary,
Bool,
DataType,
Dictionary,
Field,
FixedSizeList,
Float32,
Float64,
Int32,
Int64,
List,
Schema,
Struct,
Utf8,
util as arrowUtil,
} from "apache-arrow";
import { typedArrayToArrowType } from "./arrow_type";
import { sanitizeType } from "./sanitize";
type InferenceOptions = {
dictionaryEncodeStrings: boolean;
vectorColumns: Record<string, { type: unknown }>;
};
/**
* Infer the Arrow schema represented by a set of records.
*
* This is the intentionally small interface to schema inference. The stateful
* details of combining partial type evidence are encapsulated below so callers
* only need to provide records, an optional schema, and inference options.
*/
export function inferSchema(
data: Array<Record<string, unknown>>,
schema: Schema | undefined,
options: InferenceOptions,
): Schema {
return new SchemaInferrer(schema, options).infer(data);
}
class SchemaInferrer {
private readonly fields = new FieldTree();
constructor(
private readonly providedSchema: Schema | undefined,
private readonly options: InferenceOptions,
) {}
infer(data: Array<Record<string, unknown>>): Schema {
for (const [row, record] of data.entries()) {
for (const [path, value] of recordPathsAndValues(record)) {
this.observe(path, value, row);
}
}
return this.providedSchema === undefined
? new Schema(fieldsFromTree(this.fields))
: new Schema(matchingFields(this.providedSchema.fields, this.fields));
}
private observe(path: string[], value: unknown, row: number): void {
const current = this.fields.get(path);
if (current === undefined) {
this.addField(path, value, row);
} else if (this.providedSchema === undefined) {
this.updateInferredField(path, value, row, current);
}
}
private addField(path: string[], value: unknown, row: number): void {
if (this.providedSchema !== undefined) {
this.addSchemaField(this.providedSchema, path, row);
return;
}
const evidence =
this.inferType(value, path) ?? DeferredTypeEvidence.from(value, row);
if (evidence === undefined) {
throw typeInferenceError(path, row);
}
const conflict = this.fields.set(
path,
evidence,
(existing) =>
existing instanceof DeferredTypeEvidence && existing.isOnlyNulls(),
);
if (conflict !== undefined) {
throw branchConflictError(conflict, row, "Struct");
}
}
private addSchemaField(schema: Schema, path: string[], row: number): void {
const field = fieldAtPath(schema, path);
if (field === undefined) {
throw new Error(
`Found field not in schema: ${path.join(".")} at row ${row}`,
);
}
const conflict = this.fields.set(path, field.type);
if (conflict !== undefined) {
throw branchConflictError(conflict, row, "Struct");
}
}
private updateInferredField(
path: string[],
value: unknown,
row: number,
current: FieldNode,
): void {
const newType = this.inferType(value, path);
const deferred = DeferredTypeEvidence.from(value, row);
if (current instanceof FieldTree) {
if (deferred?.isOnlyNulls()) {
return;
}
throw schemaInferenceError(
path,
row,
"Struct",
describeEvidence(newType ?? deferred),
);
}
if (current instanceof DeferredTypeEvidence) {
this.resolveDeferredField(path, row, current, newType, deferred);
return;
}
if (newType !== undefined) {
if (!inferredTypesEqual(current, newType)) {
throw schemaInferenceError(
path,
row,
describeEvidence(current),
describeEvidence(newType),
);
}
return;
}
if (deferred === undefined || !deferred.matches(current)) {
throw schemaInferenceError(
path,
row,
describeEvidence(current),
describeEvidence(deferred),
);
}
}
private resolveDeferredField(
path: string[],
row: number,
current: DeferredTypeEvidence,
newType: DataType | undefined,
deferred: DeferredTypeEvidence | undefined,
): void {
if (newType !== undefined) {
if (!current.matches(newType)) {
throw schemaInferenceError(
path,
row,
current.describe(),
describeEvidence(newType),
);
}
this.fields.set(path, newType);
return;
}
if (deferred !== undefined) {
this.fields.set(path, current.merge(deferred));
return;
}
throw schemaInferenceError(
path,
row,
current.describe(),
describeEvidence(newType),
);
}
private inferType(value: unknown, path: string[]): DataType | undefined {
if (typeof value === "bigint") {
return new Int64();
}
if (typeof value === "number") {
return new Float64();
}
if (typeof value === "string") {
return this.options.dictionaryEncodeStrings
? new Dictionary(new Utf8(), new Int32())
: new Utf8();
}
if (typeof value === "boolean") {
return new Bool();
}
if (value instanceof Buffer) {
return new Binary();
}
if (ArrayBuffer.isView(value) && !(value instanceof DataView)) {
const typedArray = typedArrayToArrowType(value);
return typedArray === undefined
? undefined
: new FixedSizeList(
typedArray.length,
new Field("item", typedArray.elementType, true),
);
}
if (!Array.isArray(value) || value.length === 0) {
return undefined;
}
const configuredVector =
path.length === 1 ? this.options.vectorColumns[path[0]] : undefined;
if (configuredVector !== undefined) {
return new FixedSizeList(
value.length,
new Field("item", sanitizeType(configuredVector.type), true),
);
}
const itemType = this.inferArrayItemType(value, path);
if (itemType === undefined) {
return undefined;
}
return nameSuggestsVectorColumn(path[path.length - 1])
? new FixedSizeList(value.length, new Field("item", new Float32(), true))
: new List(new Field("item", itemType, true));
}
private inferArrayItemType(
values: unknown[],
path: string[],
): DataType | undefined {
let itemType: DataType | undefined;
const deferredItems: unknown[] = [];
for (const value of values) {
const candidate = this.inferType(value, path);
if (candidate === undefined) {
if (!isDeferredValue(value)) {
return undefined;
}
deferredItems.push(value);
} else if (itemType === undefined) {
itemType = candidate;
} else if (!inferredTypesEqual(itemType, candidate)) {
return undefined;
}
}
if (itemType === undefined) {
return undefined;
}
return deferredItems.every((value) =>
deferredValueMatchesType(value, itemType),
)
? itemType
: undefined;
}
}
/** Nulls and empty/all-null lists that do not determine a type by themselves. */
class DeferredTypeEvidence {
private constructor(
private readonly values: Array<{ value: unknown; row: number }>,
) {}
static from(value: unknown, row: number): DeferredTypeEvidence | undefined {
return isDeferredValue(value)
? new DeferredTypeEvidence([{ value, row }])
: undefined;
}
isOnlyNulls(): boolean {
return this.values.every(({ value }) => value == null);
}
matches(type: DataType): boolean {
return this.values.every(({ value }) =>
deferredValueMatchesType(value, type),
);
}
merge(other: DeferredTypeEvidence): DeferredTypeEvidence {
return new DeferredTypeEvidence([...this.values, ...other.values]);
}
describe(): string {
const list = this.values.find(({ value }) => Array.isArray(value));
return list === undefined
? "null"
: `List[${(list.value as unknown[]).length}]`;
}
firstRow(): number {
return this.values[0].row;
}
}
type FieldNode = DataType | DeferredTypeEvidence | FieldTree;
type LeafNode = Exclude<FieldNode, FieldTree>;
type FieldConflict = { path: string[]; value: FieldNode };
/** Nested field state, kept separate from Arrow's eventual Struct types. */
class FieldTree {
private readonly children = new Map<string, FieldNode>();
get(path: string[]): FieldNode | undefined {
let current: FieldNode = this;
for (const part of path) {
if (!(current instanceof FieldTree)) {
return undefined;
}
const child = current.children.get(part);
if (child === undefined) {
return undefined;
}
current = child;
}
return current;
}
set(
path: string[],
value: LeafNode,
canReplaceLeaf: (value: LeafNode) => boolean = () => false,
): FieldConflict | undefined {
let branch: FieldTree = this;
for (const [index, part] of path.slice(0, -1).entries()) {
const child = branch.children.get(part);
if (child === undefined || (isLeaf(child) && canReplaceLeaf(child))) {
const nextBranch = new FieldTree();
branch.children.set(part, nextBranch);
branch = nextBranch;
} else if (child instanceof FieldTree) {
branch = child;
} else {
return { path: path.slice(0, index + 1), value: child };
}
}
const name = path[path.length - 1];
const current = branch.children.get(name);
if (current instanceof FieldTree) {
return { path, value: current };
}
branch.children.set(name, value);
return undefined;
}
entries(): IterableIterator<[string, FieldNode]> {
return this.children.entries();
}
has(name: string): boolean {
return this.children.has(name);
}
}
function isLeaf(value: FieldNode): value is LeafNode {
return !(value instanceof FieldTree);
}
function fieldsFromTree(tree: FieldTree, path: string[] = []): Field[] {
const fields: Field[] = [];
for (const [name, value] of tree.entries()) {
if (value instanceof FieldTree) {
fields.push(
new Field(
name,
new Struct(fieldsFromTree(value, [...path, name])),
true,
),
);
} else if (value instanceof DeferredTypeEvidence) {
throw typeInferenceError([...path, name], value.firstRow());
} else {
fields.push(new Field(name, value, true));
}
}
return fields;
}
function matchingFields(fields: Field[], tree: FieldTree): Field[] {
const matches: Field[] = [];
for (const field of fields) {
if (!tree.has(field.name)) {
continue;
}
const value = tree.get([field.name]);
if (value instanceof FieldTree) {
const struct = field.type as Struct;
matches.push(
new Field(
field.name,
new Struct(matchingFields(struct.children, value)),
field.nullable,
),
);
} else {
matches.push(new Field(field.name, value as DataType, field.nullable));
}
}
return matches;
}
function* recordPathsAndValues(
record: Record<string, unknown>,
path: string[] = [],
): Generator<[string[], unknown]> {
for (const [name, value] of Object.entries(record)) {
if (isRecord(value)) {
yield* recordPathsAndValues(value, [...path, name]);
} else if (value !== undefined) {
yield [[...path, name], value];
}
}
}
function isRecord(value: unknown): value is Record<string, unknown> {
return (
typeof value === "object" &&
value !== null &&
!Array.isArray(value) &&
!(value instanceof RegExp) &&
!(value instanceof Date) &&
!(value instanceof Set) &&
!(value instanceof Map) &&
!(value instanceof Buffer) &&
!ArrayBuffer.isView(value)
);
}
function fieldAtPath(schema: Schema, path: string[]): Field | undefined {
let fields = schema.fields;
let field: Field | undefined;
for (const [index, name] of path.entries()) {
field = fields.find((candidate) => candidate.name === name);
if (field === undefined || index === path.length - 1) {
return field;
}
if (!DataType.isStruct(field.type)) {
return undefined;
}
fields = field.type.children;
}
return field;
}
function isDeferredValue(value: unknown): boolean {
return (
value == null || (Array.isArray(value) && value.every(isDeferredValue))
);
}
function deferredValueMatchesType(value: unknown, type: DataType): boolean {
if (value == null) {
return true;
}
if (!Array.isArray(value)) {
return false;
}
if (DataType.isList(type)) {
return value.every((item) =>
deferredValueMatchesType(item, type.valueType),
);
}
if (DataType.isFixedSizeList(type)) {
return (
value.length === type.listSize &&
value.every((item) => deferredValueMatchesType(item, type.valueType))
);
}
return false;
}
function inferredTypesEqual(current: DataType, candidate: DataType): boolean {
if (DataType.isDictionary(current)) {
return (
DataType.isDictionary(candidate) &&
current.isOrdered === candidate.isOrdered &&
inferredTypesEqual(current.indices, candidate.indices) &&
inferredTypesEqual(current.dictionary, candidate.dictionary)
);
}
if (DataType.isList(current)) {
return (
DataType.isList(candidate) &&
current.valueField.name === candidate.valueField.name &&
current.valueField.nullable === candidate.valueField.nullable &&
inferredTypesEqual(current.valueType, candidate.valueType)
);
}
if (DataType.isFixedSizeList(current)) {
return (
DataType.isFixedSizeList(candidate) &&
current.listSize === candidate.listSize &&
current.valueField.name === candidate.valueField.name &&
current.valueField.nullable === candidate.valueField.nullable &&
inferredTypesEqual(current.valueType, candidate.valueType)
);
}
return arrowUtil.compareTypes(current, candidate);
}
function describeEvidence(
evidence: DataType | DeferredTypeEvidence | undefined,
): string {
if (evidence === undefined) {
return "an unsupported value";
}
return evidence instanceof DeferredTypeEvidence
? evidence.describe()
: evidence.toString();
}
function branchConflictError(
conflict: FieldConflict,
row: number,
candidate: string,
): Error {
return schemaInferenceError(
conflict.path,
row,
conflict.value instanceof FieldTree
? "Struct"
: describeEvidence(conflict.value),
candidate,
);
}
function schemaInferenceError(
path: string[],
row: number,
currentType: string,
newType: string,
): Error {
return new Error(
`Failed to infer schema for data. Previously inferred type ${currentType} ` +
`but found ${newType} for field ${path.join(".")} at row ${row}. ` +
"Consider providing an explicit schema.",
);
}
function typeInferenceError(path: string[], row: number): Error {
return new Error(
`Failed to infer data type for field ${path.join(".")} at row ${row}. ` +
"Consider providing an explicit schema.",
);
}
function nameSuggestsVectorColumn(name: string): boolean {
const normalized = name.toLowerCase();
return normalized.includes("vector") || normalized.includes("embedding");
}
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-darwin-arm64",
"version": "0.38.0-beta.6",
"version": "0.38.0-beta.9",
"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.38.0-beta.6",
"version": "0.38.0-beta.9",
"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.38.0-beta.6",
"version": "0.38.0-beta.9",
"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.38.0-beta.6",
"version": "0.38.0-beta.9",
"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.38.0-beta.6",
"version": "0.38.0-beta.9",
"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.38.0-beta.6",
"version": "0.38.0-beta.9",
"os": [
"win32"
],
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@lancedb/lancedb-win32-x64-msvc",
"version": "0.38.0-beta.6",
"version": "0.38.0-beta.9",
"os": ["win32"],
"cpu": ["x64"],
"main": "lancedb.win32-x64-msvc.node",
+2 -2
View File
@@ -1,12 +1,12 @@
{
"name": "@lancedb/lancedb",
"version": "0.38.0-beta.6",
"version": "0.38.0-beta.8",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "@lancedb/lancedb",
"version": "0.38.0-beta.6",
"version": "0.38.0-beta.8",
"cpu": [
"x64",
"arm64"
+1 -1
View File
@@ -11,7 +11,7 @@
"ann"
],
"private": false,
"version": "0.38.0-beta.6",
"version": "0.38.0-beta.9",
"main": "dist/index.js",
"exports": {
".": "./dist/index.js",
+34
View File
@@ -17,6 +17,7 @@ use lancedb::connection::{ConnectBuilder, Connection as LanceDBConnection, conne
use lance_namespace::models::{
CreateNamespaceRequest, DescribeNamespaceRequest, DropNamespaceRequest, ListNamespacesRequest,
ListTablesRequest,
};
use lancedb::ipc::{ipc_file_to_batches, ipc_file_to_schema};
@@ -36,6 +37,12 @@ pub struct ListNamespacesResponse {
pub page_token: Option<String>,
}
#[napi(object)]
pub struct ListTablesResponse {
pub tables: Vec<String>,
pub page_token: Option<String>,
}
#[napi(object)]
pub struct CreateNamespaceResponse {
pub properties: Option<HashMap<String, String>>,
@@ -206,6 +213,33 @@ impl Connection {
op.execute().await.default_error()
}
/// List a page of tables in the database.
#[napi(catch_unwind)]
pub async fn list_tables(
&self,
namespace_path: Option<Vec<String>>,
page_token: Option<String>,
limit: Option<u32>,
) -> napi::Result<ListTablesResponse> {
let request = ListTablesRequest {
// The root namespace is an empty path, not an absent one: a namespace-backed
// database rejects a request that names no namespace.
id: Some(namespace_path.unwrap_or_default()),
page_token,
limit: limit.map(|limit| i32::try_from(limit).unwrap_or(i32::MAX)),
..Default::default()
};
let response = self
.get_inner()?
.list_tables(request)
.await
.default_error()?;
Ok(ListTablesResponse {
tables: response.tables,
page_token: response.page_token,
})
}
/// Create table from a Apache Arrow IPC (file) buffer.
///
/// Parameters:
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "lancedb-python"
version = "0.38.0-beta.6"
version = "0.38.0-beta.9"
publish = false
edition.workspace = true
description = "Python bindings for LanceDB"
+192 -67
View File
@@ -12,17 +12,19 @@ expression-backed refresh job.
from __future__ import annotations
import ast
import builtins
import base64
import functools
import hashlib
import importlib
import inspect
import symtable
import json
import math
import re
import sys
import textwrap
import types
import uuid
from collections.abc import Mapping
from datetime import date, datetime
from typing import (
@@ -273,7 +275,7 @@ class FunctionVersion(_RemoteValue):
Every input must be a direct [lancedb.col][lancedb.expr.col]
reference. The returned application is immutable and retains a
named-struct output as one sibling group, so every row's sibling values
named-struct output as one binding, so every row's sibling values
come from one logical Function evaluation. Map result fields to table
columns with
[FunctionApplication.rename][lancedb.functions.FunctionApplication.rename],
@@ -323,7 +325,6 @@ class FunctionVersion(_RemoteValue):
function=FunctionVersionRef(name=self.name, version=self.version),
inputs=tuple(bindings),
output=self.signature.output,
group_id=f"fg_{uuid.uuid4().hex}",
)
@@ -367,7 +368,7 @@ class ApplicationInput(_OpenRemoteValue):
class FunctionApplication(_OpenRemoteValue):
"""Immutable pre-declaration application of an exact Function version.
A named-struct output remains one grouped application through table
A named-struct output remains one application through table
declaration and execution.
[FunctionApplication.rename][lancedb.functions.FunctionApplication.rename]
records the result-field to table-column mapping without splitting sibling
@@ -377,7 +378,6 @@ class FunctionApplication(_OpenRemoteValue):
function: FunctionVersionRef
inputs: tuple[ApplicationInput, ...]
output: FunctionOutput
group_id: str
columns: Mapping[str, str] = Field(default_factory=dict)
def _known_dict(self) -> dict[str, Any]:
@@ -449,12 +449,10 @@ class OutputMapping(_RemoteValue):
class FunctionBinding(_RemoteValue):
"""Immutable grouped binding persisted by the Enterprise table service."""
"""Immutable Function binding persisted by the Enterprise table service."""
binding_id: str
revision: _UInt64
function: FunctionVersionRef
group_id: str
inputs: tuple[InputBinding, ...]
outputs: tuple[OutputMapping, ...]
input_schema: Optional[Mapping[str, Any]] = None
@@ -489,59 +487,58 @@ _FUNCTION_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_.-]*$")
_SECRET_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
_GRAMMAR_PRIMITIVES = (
(pa.bool_(), "bool"),
(pa.int8(), "int8"),
(pa.int16(), "int16"),
(pa.int32(), "int32"),
(pa.int64(), "int64"),
(pa.uint8(), "uint8"),
(pa.uint16(), "uint16"),
(pa.uint32(), "uint32"),
(pa.uint64(), "uint64"),
(pa.float16(), "float16"),
(pa.float32(), "float32"),
(pa.float64(), "float64"),
(pa.string(), "utf8"),
(pa.binary(), "binary"),
(pa.date32(), "date32"),
(pa.date64(), "date64"),
)
def _canonical_arrow_type(data_type: pa.DataType) -> str:
primitive_types = (
(pa.bool_(), "bool"),
(pa.int8(), "int8"),
(pa.int16(), "int16"),
(pa.int32(), "int32"),
(pa.int64(), "int64"),
(pa.uint8(), "uint8"),
(pa.uint16(), "uint16"),
(pa.uint32(), "uint32"),
(pa.uint64(), "uint64"),
(pa.float16(), "float16"),
(pa.float32(), "float32"),
(pa.float64(), "float64"),
(pa.string(), "utf8"),
(pa.large_utf8(), "large_utf8"),
(pa.binary(), "binary"),
(pa.large_binary(), "large_binary"),
(pa.date32(), "date32"),
(pa.date64(), "date64"),
)
for candidate, name in primitive_types:
"""The server's V1 Function type grammar. Anything outside it is rejected
here rather than at registration."""
for candidate, name in _GRAMMAR_PRIMITIVES:
if data_type == candidate:
return name
if pa.types.is_fixed_size_binary(data_type):
return f"fixed_size_binary[{data_type.byte_width}]"
if pa.types.is_list(data_type):
return f"list<{_canonical_arrow_type(data_type.value_type)}>"
if pa.types.is_large_list(data_type):
return f"large_list<{_canonical_arrow_type(data_type.value_type)}>"
if pa.types.is_fixed_size_list(data_type):
if pa.types.is_list(data_type) or pa.types.is_large_list(data_type):
prefix = "list" if pa.types.is_list(data_type) else "large_list"
return f"{prefix}<{_canonical_list_item(data_type)}>"
if pa.types.is_fixed_size_list(data_type) and data_type.list_size > 0:
return (
f"fixed_size_list<{_canonical_arrow_type(data_type.value_type)}>"
f"[{data_type.list_size}]"
f"fixed_size_list<{_canonical_list_item(data_type)}, {data_type.list_size}>"
)
if pa.types.is_struct(data_type):
fields = ",".join(
f"{field.name}:{_canonical_arrow_type(field.type)}" for field in data_type
)
return f"struct<{fields}>"
if pa.types.is_timestamp(data_type):
timezone = f",tz={data_type.tz}" if data_type.tz is not None else ""
return f"timestamp[{data_type.unit}{timezone}]"
if pa.types.is_time32(data_type) or pa.types.is_time64(data_type):
return f"time[{data_type.unit}]"
if pa.types.is_duration(data_type):
return f"duration[{data_type.unit}]"
if pa.types.is_decimal(data_type):
bit_width = data_type.bit_width
return f"decimal{bit_width}({data_type.precision},{data_type.scale})"
raise TypeError(f"unsupported Arrow type for Function signature: {data_type}")
def _canonical_list_item(data_type: pa.DataType) -> str:
"""The grammar names only the item type; it always means a non-nullable
child called `item`, so any other child metadata cannot be represented."""
child = data_type.value_field
if child.name != "item" or child.nullable or child.metadata:
raise TypeError(
"unsupported Arrow type for Function signature: list items must be a "
f"non-nullable field named 'item', got {child}"
)
return _canonical_arrow_type(child.type)
def _list_of(item: pa.DataType) -> pa.DataType:
return pa.list_(pa.field("item", item, nullable=False))
def _annotation_type(annotation: Any) -> tuple[pa.DataType, bool]:
nullable = False
origin = get_origin(annotation)
@@ -589,7 +586,7 @@ def _annotation_type(annotation: Any) -> tuple[pa.DataType, bool]:
value_type, value_nullable = _annotation_type(arguments[0])
if value_nullable:
raise TypeError("nullable Function list elements are not supported")
return pa.list_(value_type), nullable
return _list_of(value_type), nullable
raise TypeError(f"unsupported Function annotation: {annotation!r}")
@@ -736,6 +733,104 @@ def _literal_source(value: Any) -> str:
)
_DYNAMIC_NAMESPACE_ACCESS = frozenset(
{"globals", "locals", "vars", "eval", "exec", "compile", "__import__"}
)
# Modules that hand out namespaces (`sys.modules`, `builtins`, importers,
# introspection). The artifact's module namespace holds only the names it was
# packaged with, so reaching around it cannot be represented.
_NAMESPACE_MODULES = frozenset(
{"sys", "builtins", "importlib", "inspect", "gc", "ctypes", "types"}
)
def _namespace_acquisition(
definition: ast.FunctionDef, references: set[str]
) -> list[str]:
found = set(references & _DYNAMIC_NAMESPACE_ACCESS)
for node in ast.walk(definition):
if isinstance(node, ast.Import):
found.update(
alias.name
for alias in node.names
if alias.name.split(".")[0] in _NAMESPACE_MODULES
)
elif isinstance(node, ast.ImportFrom) and node.module:
if node.module.split(".")[0] in _NAMESPACE_MODULES:
found.add(node.module)
return sorted(found)
def _module_references(module_source: str) -> set[str]:
"""Names any scope in `module_source` binds or loads at module scope.
Python's own scope analysis on the exact text that ships: free variables
belong to an enclosing scope inside the function, and postponed
annotations are not runtime loads."""
def visit(table: symtable.SymbolTable, found: set[str]) -> None:
for symbol in table.get_symbols():
if symbol.is_global() and (
symbol.is_referenced() or symbol.is_declared_global()
):
found.add(symbol.get_name())
for child in table.get_children():
visit(child, found)
found: set[str] = set()
for table in symtable.symtable(module_source, "<udf>", "exec").get_children():
visit(table, found)
return found
def _global_source(name: str, value: Any) -> str:
"""One module-level line that rebinds `name` to `value` in the artifact:
an import for modules and importable classes/functions, a literal otherwise."""
if isinstance(value, types.ModuleType):
if value.__name__.split(".")[0] in _NAMESPACE_MODULES:
raise ValueError(
f"@udf cannot package dynamic namespace access: {value.__name__!r}"
)
try:
imported = importlib.import_module(value.__name__)
except ImportError:
imported = None
if imported is not value:
raise TypeError(
f"Function source references module {name!r} that does not import "
f"as {value.__name__!r}"
)
return f"import {value.__name__} as {name}"
module_name = getattr(value, "__module__", None)
qualname = getattr(value, "__qualname__", None)
if (
isinstance(module_name, str)
and isinstance(qualname, str)
and module_name != "__main__"
and "." not in qualname
and "<" not in qualname
):
try:
imported = getattr(importlib.import_module(module_name), qualname)
except (ImportError, AttributeError):
imported = None
if imported is value:
return f"from {module_name} import {qualname} as {name}"
return f"{name} = {_literal_source(value)}"
def _is_recursive_reference(function: Callable[..., Any], name: str) -> bool:
"""`name` inside the body means the function itself unless the module has
since bound it to something else."""
if name != function.__name__:
return False
bound = function.__globals__.get(name, function)
if bound is function:
return True
# The decorator's own result is the one wrapper known to call `function`
# unchanged; any other binding may behave differently from a self-call.
return type(bound) is UdfDefinition and bound._function is function
def _package_source(function: Callable[..., Any]) -> bytes:
if not inspect.isfunction(function) or inspect.iscoroutinefunction(function):
raise TypeError("@udf requires a synchronous Python function")
@@ -760,23 +855,46 @@ def _package_source(function: Callable[..., Any]) -> bytes:
closure = inspect.getclosurevars(function)
if closure.nonlocals:
raise ValueError("@udf cannot package functions that capture closure values")
if closure.unbound:
raise ValueError(
f"@udf source contains unresolved global names: {sorted(closure.unbound)!r}"
)
globals_source = []
for name, value in sorted(closure.globals.items()):
if isinstance(value, types.ModuleType):
globals_source.append(f"import {value.__name__} as {name}")
else:
globals_source.append(f"{name} = {_literal_source(value)}")
function_source = ast.unparse(definition)
parts = ["from __future__ import annotations"]
module_header = "from __future__ import annotations"
references = _module_references(f"{module_header}\n\n{function_source}\n")
dynamic = _namespace_acquisition(definition, references)
if dynamic:
raise ValueError(f"@udf cannot package dynamic namespace access: {dynamic!r}")
# Resolve every module-scope reference the way the interpreter would: the
# function's own globals first (a module global may shadow a builtin, and
# nested scopes are not visible to getclosurevars), then its builtins.
# The artifact runs under the standard builtins; only the exact mapping is
# provably equivalent (a subclass or copy can change lookups and hooks).
if function.__builtins__ is not vars(builtins):
raise ValueError("@udf cannot package a non-standard builtins environment")
globals_source = []
unresolved = []
for name in sorted(references):
if name == function.__name__:
if not _is_recursive_reference(function, name):
raise ValueError(
f"@udf cannot package {name!r}: the module binds that name to "
"another value, which the artifact's own definition would shadow"
)
continue
if name in function.__globals__:
globals_source.append(_global_source(name, function.__globals__[name]))
elif hasattr(builtins, name):
pass
else:
unresolved.append(name)
if unresolved:
raise ValueError(
f"@udf source contains unresolved global names: {unresolved!r}"
)
parts = [module_header]
if globals_source:
parts.extend(["", *globals_source])
parts.extend(["", function_source, ""])
return "\n".join(parts).encode("utf-8")
packaged = "\n".join(parts)
return packaged.encode("utf-8")
class UdfDefinition:
@@ -923,6 +1041,13 @@ def udf(
python_version : str, optional
Remote Python major/minor version. Defaults to the client version.
The packaged artifact is a snapshot: the function source plus exactly
the module-level names it references (modules as imports, importable
classes and functions as imports, literals inline). Code that reaches the
module namespace another way -- ``globals()``/``eval``, ``sys.modules``,
``builtins`` -- is rejected where it can be seen and otherwise
unsupported; closures and a non-standard ``__builtins__`` are rejected.
Returns
-------
UdfDefinition
+21 -5
View File
@@ -391,6 +391,15 @@ def _table_to_pickle_state(table: Table) -> dict[str, Any]:
}
def _drop_base_version(permutation_data: pa.Table) -> pa.Table:
"""Strip the recorded base version so the reader leaves the base table unpinned."""
metadata = dict(permutation_data.schema.metadata or {})
if metadata.pop(b"base_version", None) is None:
return permutation_data
metadata.pop(b"base_branch", None)
return permutation_data.replace_schema_metadata(metadata)
def _table_from_pickle_state(state: dict[str, Any]) -> Table:
from . import connect
@@ -679,11 +688,15 @@ class Permutation:
from . import connect
connection_factory = state["connection_factory"]
rebuilt_base = False
if connection_factory is not None:
base_table = connection_factory(state["base_table_name"])
elif "base_table_state" in state:
base_table = _table_from_pickle_state(state["base_table_state"])
base_state = state["base_table_state"]
rebuilt_base = base_state["kind"] == "memory"
base_table = _table_from_pickle_state(base_state)
elif "base_table_data" in state:
rebuilt_base = True
# In-memory base table inlined into the pickle; rebuild the same
# way we rebuild the in-memory permutation table.
mem_db = connect("memory://")
@@ -701,11 +714,14 @@ class Permutation:
)
permutation_table: Optional[Table] = None
if state["permutation_data"] is not None:
permutation_data = state["permutation_data"]
if permutation_data is not None:
if rebuilt_base:
# The base table was materialized from Arrow, so it is a fresh
# single-version dataset and the recorded pin cannot resolve on it.
permutation_data = _drop_base_version(permutation_data)
mem_db = connect("memory://")
permutation_table = mem_db.create_table(
"permutation", state["permutation_data"]
)
permutation_table = mem_db.create_table("permutation", permutation_data)
self.base_table = base_table
self.permutation_table = permutation_table
+565 -33
View File
@@ -19,6 +19,7 @@ above.
"""
import ctypes
import heapq
import logging
import os
import random
@@ -29,17 +30,18 @@ from collections import deque
from concurrent.futures import ThreadPoolExecutor
from copy import deepcopy
from multiprocessing import RawArray
from typing import Any, Callable, cast, Iterator, Literal, Optional, Union
from typing import Any, Callable, cast, Iterator, Literal, NamedTuple, Optional, Union
import pyarrow as pa
import pyarrow.compute as pc
import torch
from torch.utils.data import IterableDataset, get_worker_info
from torch.utils.data import DataLoader, IterableDataset, get_worker_info
from .permutation import (
Permutation,
Transforms,
permutation_builder,
_drop_base_version,
_table_from_pickle_state,
_table_to_pickle_state,
)
@@ -55,6 +57,155 @@ DEFAULT_READ_BATCH_SIZE = 64
DEFAULT_PREFETCH_BATCHES = 4
class _WorkerSample(NamedTuple):
data: Any
dataset: "StreamingDataset"
class _WorkerBatch(NamedTuple):
data: Any
state: dict
class _ConsumerIteratorLease(NamedTuple):
owner_token: int
owner_thread: int
class _CheckpointCollate:
"""Attach the worker's post-fetch state to a collated batch."""
def __init__(self, collate_fn: Callable):
self._collate_fn = collate_fn
def __call__(self, samples):
try:
if isinstance(samples, list):
if not samples:
return _WorkerBatch(self._collate_fn(samples), {})
worker_samples = samples
data = self._collate_fn([sample.data for sample in worker_samples])
dataset = worker_samples[-1].dataset
else:
data = self._collate_fn(samples.data)
dataset = samples.dataset
except StopIteration as exc:
raise RuntimeError(
"collate_fn raised StopIteration before returning a batch"
) from exc
return _WorkerBatch(data, dataset._checkpoint_snapshot())
class _StreamingDatasetAdapter(IterableDataset):
"""Yield private sample wrappers for :class:`StreamingDataLoader`."""
def __init__(self, dataset: "StreamingDataset"):
super().__init__()
self.dataset = dataset
def __iter__(self):
for sample in self.dataset._iter(consumer_checkpoint_transport=True):
yield _WorkerSample(sample, self.dataset)
def __getattr__(self, name):
dataset = self.__dict__.get("dataset")
if dataset is None:
raise AttributeError(name)
return getattr(dataset, name)
class _ConsumerCommitIterator:
def __init__(
self,
iterator,
dataset: "StreamingDataset",
*,
owner_token: int,
require_uniform: bool,
):
self._iterator = iterator
self._dataset = dataset
self._owner_token = owner_token
self._require_uniform = require_uniform
self._released = False
self._terminal = False
def __iter__(self):
return self
def __next__(self):
if self._terminal:
raise StopIteration
try:
batch = next(self._iterator)
except StopIteration:
self._terminal = True
self._release()
raise
except BaseException as exc:
self._dataset._invalidate_checkpoint(
f"a DataLoader batch failed before it was returned: {exc}"
)
raise
try:
if not isinstance(batch, _WorkerBatch):
raise RuntimeError(
"StreamingDataLoader did not receive worker checkpoint metadata"
)
self._dataset._commit_worker_state(
batch.state, require_uniform=self._require_uniform
)
return batch.data
except BaseException as exc:
self._dataset._invalidate_checkpoint(
f"a DataLoader batch failed before it was returned: {exc}"
)
raise
def _release(self) -> None:
if self.__dict__.get("_released", True):
return
self._released = True
dataset = self.__dict__.get("_dataset")
if dataset is not None:
dataset._release_consumer_iterator(self._owner_token)
def _shutdown_workers(self):
if self.__dict__.get("_released", True):
return None
self._terminal = True
iterator = self.__dict__.get("_iterator")
shutdown = getattr(iterator, "_shutdown_workers", None)
try:
if shutdown is not None:
shutdown()
else:
fetcher = getattr(iterator, "_dataset_fetcher", None)
dataset_iterator = getattr(fetcher, "dataset_iter", None)
close = getattr(dataset_iterator, "close", None)
if close is None:
raise RuntimeError(
"StreamingDataLoader could not close its inner iterator"
)
close()
except BaseException as exc:
self._dataset._invalidate_checkpoint(
f"a DataLoader iterator could not be shut down safely: {exc}"
)
raise
else:
self._release()
def __del__(self):
try:
self._shutdown_workers()
except BaseException:
pass
def __getattr__(self, name):
return getattr(self._iterator, name)
class StreamingDataset(IterableDataset):
"""An elastic, resumable PyTorch IterableDataset backed by a LanceDB table.
@@ -384,6 +535,22 @@ class StreamingDataset(IterableDataset):
# rows_skipped]
self._worker_stats: RawArray = RawArray(ctypes.c_int64, 8)
# A standard multi-process DataLoader cannot report which prefetched
# batches were actually returned to its consumer. Workers set this
# shared flag so state_dict() can reject a stale parent checkpoint
# unless StreamingDataLoader installed the consumer-commit transport.
self._untracked_worker_iteration: RawArray = RawArray(ctypes.c_int64, 1)
# Parent-side checkpoint lifecycle. A failed DataLoader task creates
# a permanent hole in that iterator's delivery stream, while a
# multi-worker checkpoint is safe to restore only after all splits
# reach the same logical step boundary.
self._checkpoint_invalid_reason: Optional[str] = None
self._consumer_checkpoint_requires_uniform = False
self._consumer_iterator_lock = threading.Lock()
self._consumer_iterator_generation = 0
self._consumer_iterator_lease: Optional[_ConsumerIteratorLease] = None
# 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.
@@ -396,6 +563,10 @@ class StreamingDataset(IterableDataset):
# step boundaries all splits have consumed this many samples, so a
# single scalar captures the topology-independent checkpoint state.
self._resume_offset: int = 0
# Exact yielded-sample counts for splits this process has advanced.
# Missing entries use _resume_offset, which remains the lower-bound
# checkpoint inherited from an earlier uniform/global state.
self._resume_samples: dict[int, int] = {}
# 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
@@ -521,11 +692,45 @@ class StreamingDataset(IterableDataset):
return self._rank_splits[start : start + splits_per_worker]
def __iter__(self) -> Iterator[dict[str, Any]]:
return self._iter()
def _iter(
self, *, consumer_checkpoint_transport: bool = False
) -> Iterator[dict[str, Any]]:
owner_token = None
previous_lease = self._consumer_iterator_lease
if consumer_checkpoint_transport:
if not self._consumer_iterator_active:
raise RuntimeError(
"StreamingDataLoader worker transport requires an active "
"parent iterator reservation"
)
else:
try:
owner_token = self._acquire_consumer_iterator()
except BaseException:
self._release_consumer_iterator_after_failed_acquire(previous_lease)
raise
try:
yield from self._iter_owned(
consumer_checkpoint_transport=consumer_checkpoint_transport
)
finally:
if owner_token is not None:
self._release_consumer_iterator(owner_token)
def _iter_owned(
self, *, consumer_checkpoint_transport: bool
) -> Iterator[dict[str, Any]]:
if self._raw_batches_ref is not None:
raise RuntimeError(
"StreamingDataset does not support concurrent iteration. "
"Only one active iterator per dataset instance is allowed."
)
real_worker = get_worker_info() is not None
if real_worker and not consumer_checkpoint_transport:
self._untracked_worker_iteration[0] = 1
my_splits = self._resolve_my_splits()
if not my_splits:
return
@@ -533,6 +738,7 @@ class StreamingDataset(IterableDataset):
# Set identity transform on each Permutation so __getitems__ returns
# the raw RecordBatch. Stage 2 applies the real transform.
permutations: list[Permutation] = []
initial_samples: list[int] = []
initial_positions: list[int] = []
for split_idx in my_splits:
perm = Permutation.from_tables(
@@ -541,21 +747,22 @@ class StreamingDataset(IterableDataset):
if self._columns is not None:
perm = perm.select_columns(self._columns)
perm = perm.with_transform(Transforms.arrow2arrow)
sample_count = self._resume_samples.get(split_idx, self._resume_offset)
# Both modes resume from absolute permutation positions. Packing
# stores them separately because it also checkpoints partial blocks.
start_pos = (
self._pack_consumed[split_idx]
if self._pack_sequences is not None
else self._resume_positions.get(split_idx, self._resume_offset)
else self._resume_positions.get(split_idx, sample_count)
)
if start_pos > 0:
perm = perm.with_skip(start_pos)
initial_samples.append(sample_count)
initial_positions.append(start_pos)
permutations.append(perm)
n = len(permutations)
split_sizes = [perm.num_rows for perm in permutations]
initial_offset = self._resume_offset
local_consumed = [0] * n
# Permutation position each split has consumed through (absolute,
# i.e. counted from the start of the unskipped split). Runs ahead of
@@ -853,6 +1060,27 @@ class StreamingDataset(IterableDataset):
for i in range(n):
_fill_io(i)
def _yield_row(i: int):
pos, row = cooked[i].popleft()
# Surface any completed prefetched failure before the
# current row becomes durable checkpoint progress.
_advance(i)
local_consumed[i] += 1
pos_consumed[i] = pos + 1
split_idx = my_splits[i]
self._resume_samples[split_idx] = (
initial_samples[i] + local_consumed[i]
)
self._resume_positions[split_idx] = pos_consumed[i]
return row
def _update_progress_stats() -> None:
if not real_worker:
self._resume_offset = min(
initial_samples[j] + local_consumed[j] for j in range(n)
)
_update_stats()
if self._pack_sequences is not None:
first_count = pack_blocks_emitted[my_splits[0]]
if any(
@@ -878,12 +1106,38 @@ class StreamingDataset(IterableDataset):
tokens.extend([pad_id] * (pack_len - len(tokens)))
block = _emit_block(i)
pack_blocks_emitted[my_splits[i]] += 1
# Checkpoint state must advance before yielding so
# StreamingDataLoader can attach the exact state to
# the batch it transports to the parent process.
_commit_pack_state()
if i == n - 1:
_commit_pack_state()
_update_stats()
yield block
return
# A checkpoint taken between round-robin split turns has
# non-uniform counts. Resume lagging splits first so the
# exact canonical sequence continues without replaying
# already-consumed rows.
if len(set(initial_samples)) > 1:
catch_up_to = max(initial_samples)
pending = [
(initial_samples[i], my_splits[i], i)
for i in range(n)
if initial_samples[i] < catch_up_to
]
heapq.heapify(pending)
while pending:
consumed, _, i = heapq.heappop(pending)
_ensure_cooked(i)
if not cooked[i]:
return
row = _yield_row(i)
if consumed + 1 < catch_up_to:
heapq.heappush(pending, (consumed + 1, my_splits[i], i))
_update_progress_stats()
yield row
while True:
# A cycle only runs if every split can still produce a
# row. Without skips all splits exhaust simultaneously
@@ -904,20 +1158,14 @@ class StreamingDataset(IterableDataset):
break
for i in range(n):
pos, row = cooked[i].popleft()
local_consumed[i] += 1
pos_consumed[i] = pos + 1
_advance(i)
row = _yield_row(i)
# After the last split in each cycle: update the
# global offset and refresh the shared-memory stats
# so the main process can observe pipeline depth
# 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]
_update_stats()
_update_progress_stats()
yield row
finally:
@@ -1064,6 +1312,7 @@ class StreamingDataset(IterableDataset):
"_local_consumed_ref",
):
state[key] = None
state["_consumer_iterator_lock"] = None
return state
def __setstate__(self, state):
@@ -1074,19 +1323,31 @@ class StreamingDataset(IterableDataset):
table_state = state.pop("_table")
perm_name, perm_data = state.pop("_perm_table")
self.__dict__.update(state)
self._consumer_iterator_lock = threading.Lock()
if self._connection_factory is not None:
self._table = self._connection_factory(table_name)
else:
self._table = _table_from_pickle_state(table_state)
if table_state["kind"] == "memory":
# Rebuilt from Arrow, so the recorded pin cannot resolve on it.
perm_data = _drop_base_version(perm_data)
self._perm_table = _connect("memory://").create_table(perm_name, perm_data)
def state_dict(self) -> dict:
"""Snapshot the dataset's consumption state.
When using DataLoader workers, construct a
[StreamingDataLoader][lancedb.streaming.StreamingDataLoader]. It
commits worker state only when a prefetched batch is returned to the
trainer. A standard multi-process ``DataLoader`` cannot expose that
boundary, so calling this method after one has started raises
``RuntimeError`` instead of returning stale producer state.
In row mode, the returned dict is topology-independent at global step
boundaries. ``positions_consumed_per_split`` records how far each
split's permutation has advanced, which can differ from the sample
count when ``on_transform_error`` skips rows. Combine state dicts from
count when ``on_transform_error`` skips rows. ``StreamingDataLoader``
combines worker state in its parent process. Combine state dicts from
every rank with
[merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts]
before resuming on a different topology.
@@ -1095,6 +1356,43 @@ class StreamingDataset(IterableDataset):
for every logical split. When packing is sharded, merge every rank
state with ``merge_state_dicts`` before loading it.
"""
if self._untracked_worker_iteration[0] and get_worker_info() is None:
raise RuntimeError(
"StreamingDataset cannot checkpoint a standard DataLoader with "
"num_workers > 0 because prefetched worker progress is not "
"consumer-committed. Use StreamingDataLoader instead."
)
if self._checkpoint_invalid_reason is not None:
raise RuntimeError(
"StreamingDataset checkpointing is invalid because "
f"{self._checkpoint_invalid_reason}. Load the last valid "
"checkpoint into a fresh dataset before continuing."
)
state = self._checkpoint_snapshot()
if self._pack_sequences is not None:
rank_blocks = [
state["blocks_emitted_per_split"][split] for split in self._rank_splits
]
if len(set(rank_blocks)) > 1:
raise RuntimeError(
"Packed StreamingDataset checkpointing is only safe at a "
"complete logical step boundary, when every split assigned "
"to this rank has emitted the same block count. Consume more "
"batches before calling state_dict()."
)
elif self._consumer_checkpoint_requires_uniform:
samples = state["samples_consumed_per_split"]
rank_samples = [samples[split] for split in self._rank_splits]
if len(set(rank_samples)) > 1:
raise RuntimeError(
"StreamingDataLoader checkpointing with multiple workers is "
"only safe at a complete logical step boundary, when every "
"split assigned to this rank has the same consumed-sample "
"count. Consume more batches before calling state_dict()."
)
return state
def _checkpoint_snapshot(self) -> dict:
if self._pack_sequences is not None:
return {
"shuffle_seed": self._shuffle_seed,
@@ -1108,18 +1406,141 @@ class StreamingDataset(IterableDataset):
"blocks_emitted_per_split": list(self._pack_blocks_emitted),
"pack_buffers": deepcopy(self._pack_buffers),
}
samples = [
self._resume_samples.get(split, self._resume_offset)
for split in range(self._num_splits)
]
positions = [
self._resume_positions.get(split, self._resume_offset)
self._resume_positions.get(split, samples[split])
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,
"samples_consumed_per_split": samples,
"positions_consumed_per_split": positions,
}
def _invalidate_checkpoint(self, reason: str) -> None:
if self._checkpoint_invalid_reason is None:
self._checkpoint_invalid_reason = reason
@property
def _consumer_iterator_active(self) -> bool:
return self._consumer_iterator_lease is not None
@property
def _consumer_iterator_owner(self) -> Optional[int]:
lease = self._consumer_iterator_lease
return lease.owner_token if lease is not None else None
@property
def _consumer_iterator_owner_thread(self) -> Optional[int]:
lease = self._consumer_iterator_lease
return lease.owner_thread if lease is not None else None
def _acquire_consumer_iterator(self) -> int:
"""Reserve this parent dataset for one checkpoint-aware iterator."""
with self._consumer_iterator_lock:
if self._consumer_iterator_active or self._raw_batches_ref is not None:
raise RuntimeError(
"StreamingDataset does not support concurrent iteration. "
"Only one active iterator per dataset instance is allowed."
)
owner_thread = threading.get_ident()
owner_token = self._consumer_iterator_generation + 1
lease = _ConsumerIteratorLease(owner_token, owner_thread)
self._consumer_iterator_generation = owner_token
self._consumer_iterator_lease = lease
return owner_token
def _release_consumer_iterator(self, owner_token: int) -> None:
with self._consumer_iterator_lock:
lease = self._consumer_iterator_lease
if lease is not None and lease.owner_token == owner_token:
self._consumer_iterator_lease = None
def _release_consumer_iterator_after_failed_acquire(
self, previous_lease: Optional[_ConsumerIteratorLease]
) -> None:
"""Clean up when an interrupted acquire set a lease but did not return it."""
owner_thread = threading.current_thread().ident
with self._consumer_iterator_lock:
lease = self._consumer_iterator_lease
if (
lease is not None
and lease is not previous_lease
and lease.owner_thread == owner_thread
):
self._consumer_iterator_lease = None
def _commit_worker_state(self, state: dict, *, require_uniform: bool) -> None:
"""Merge one trainer-consumed worker batch into parent state."""
for key, expected in (
("shuffle_seed", self._shuffle_seed),
("num_splits", self._num_splits),
("epoch", self._epoch),
):
if state.get(key) != expected:
raise ValueError(
f"{key} mismatch in worker checkpoint: "
f"{state.get(key)} != {expected}"
)
packed = "pack_buffers" in state
if packed != (self._pack_sequences is not None):
raise ValueError("worker checkpoint mode does not match the dataset")
if packed:
for key in ("pack_sequences", "eos_id", "pad_id", "blocks_per_epoch"):
expected = getattr(self, f"_{key}")
if state.get(key) != expected:
raise ValueError(
f"{key} mismatch in worker checkpoint: "
f"{state.get(key)} != {expected}"
)
samples = state["samples_consumed_per_split"]
emitted = state["blocks_emitted_per_split"]
if len(samples) != self._num_splits or len(emitted) != self._num_splits:
raise ValueError(
"packed worker checkpoint must contain one entry per split"
)
buffers = state["pack_buffers"]
for split, (count, blocks) in enumerate(zip(samples, emitted)):
incoming = (int(blocks), int(count))
current = (
self._pack_blocks_emitted[split],
self._pack_consumed[split],
)
if incoming > current:
self._pack_blocks_emitted[split] = incoming[0]
self._pack_consumed[split] = incoming[1]
buffer = buffers.get(split, buffers.get(str(split)))
if buffer is None:
self._pack_buffers.pop(split, None)
else:
self._pack_buffers[split] = {
"tokens": list(buffer["tokens"]),
"starts": list(buffer["starts"]),
}
self._consumer_checkpoint_requires_uniform |= require_uniform
return
samples = state["samples_consumed_per_split"]
positions = state.get("positions_consumed_per_split", samples)
for split, count in enumerate(samples):
current = self._resume_samples.get(split, self._resume_offset)
self._resume_samples[split] = max(current, int(count))
for split, position in enumerate(positions):
current = self._resume_positions.get(
split, self._resume_samples.get(split, self._resume_offset)
)
self._resume_positions[split] = max(current, int(position))
self._resume_offset = min(
self._resume_samples.get(split, self._resume_offset)
for split in range(self._num_splits)
)
self._consumer_checkpoint_requires_uniform |= require_uniform
def load_state_dict(self, state: dict) -> None:
"""Resume from a previously snapshotted state.
@@ -1139,6 +1560,7 @@ class StreamingDataset(IterableDataset):
f"shuffle_seed mismatch: checkpoint has {state['shuffle_seed']}, "
f"current dataset has {self._shuffle_seed}"
)
self._consumer_checkpoint_requires_uniform = False
if "pack_buffers" in state or self._pack_sequences is not None:
for key in (
@@ -1165,14 +1587,17 @@ class StreamingDataset(IterableDataset):
return
consumed = state["samples_consumed_per_split"]
# All entries are equal at step boundaries; use the first.
if isinstance(consumed, list):
self._resume_offset = consumed[0] if consumed else 0
self._resume_offset = min(consumed) if consumed else 0
self._resume_samples = {
split: int(count) for split, count in enumerate(consumed)
}
else:
self._resume_offset = int(consumed)
self._resume_samples = {}
# 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.
# the per-split sample count (the .get default in __iter__) is exact.
positions = state.get("positions_consumed_per_split")
if positions is None:
self._resume_positions = {}
@@ -1185,10 +1610,11 @@ class StreamingDataset(IterableDataset):
def merge_state_dicts(states: list[dict]) -> dict:
"""Merge state dicts saved by different ranks into one exact state.
For row mode, the elementwise maximum of permutation positions recovers
splits advanced by different ranks after transform failures. For packed
mode, the state that emitted the most blocks for each logical split
supplies that split's permutation position and partial token buffer. Packed
In row mode, each rank records exact consumer-committed progress for
its own splits and lower bounds for the rest, so elementwise maxima
recover both sample counts and permutation positions. In packed mode,
the state that emitted the most blocks for each logical split supplies
that split's permutation position and partial token buffer. Packed
states must cover every rank at the same global step.
Raises ``ValueError`` if the states are empty, were not produced by
@@ -1299,17 +1725,13 @@ class StreamingDataset(IterableDataset):
merged["pack_buffers"] = merged_buffers
return merged
for state in states[1:]:
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)
merged["samples_consumed_per_split"] = [
max(per_split)
for per_split in zip(
*(state["samples_consumed_per_split"] for state in states)
)
]
all_positions = [
state.get(
"positions_consumed_per_split", state["samples_consumed_per_split"]
@@ -1320,3 +1742,113 @@ class StreamingDataset(IterableDataset):
max(per_split) for per_split in zip(*all_positions)
]
return merged
class StreamingDataLoader(DataLoader):
"""A PyTorch DataLoader with consumer-committed dataset checkpoints.
PyTorch workers prefetch batches ahead of the trainer, so worker-local
producer progress is not a safe checkpoint. This loader carries a state
snapshot alongside every internal batch and applies it to the parent
[StreamingDataset][lancedb.streaming.StreamingDataset] only when that batch
is returned by ``next()``.
The trainer receives the same collated batch it would receive from a
standard ``torch.utils.data.DataLoader``.
With more than one worker, row-mode ``state_dict()`` is available only at
complete logical step boundaries, when every split assigned to the rank has
the same consumed-sample count. Packed checkpoints require equal emitted-block
counts across the rank's splits for any worker count. ``persistent_workers=True``
is not supported because prefetched worker copies cannot be restored from
parent-committed state. If batch collation raises, checkpointing remains
invalid for that dataset instance; restore the last valid checkpoint into a
fresh dataset before continuing.
Only one active iterator may own a dataset at a time, including when worker
processes are used. Exhausting or explicitly shutting down the iterator
releases that ownership. ``drop_last=True`` is not supported because worker
replicas discard incomplete tails independently, which cannot produce a
topology-independent checkpoint.
Parameters are the same as ``torch.utils.data.DataLoader`` except that
``dataset`` must be a
[StreamingDataset][lancedb.streaming.StreamingDataset].
Subclasses that override ``StreamingDataset.__iter__`` are not supported
because the custom iterator cannot provide the exact per-yield checkpoint
snapshots required by this loader.
Examples
--------
>>> # dataset = StreamingDataset(table, num_splits=2)
>>> # loader = StreamingDataLoader(dataset, batch_size=8, num_workers=2)
>>> # batch = next(iter(loader))
>>> # checkpoint = dataset.state_dict()
"""
def __init__(self, dataset: StreamingDataset, *args, **kwargs):
if not isinstance(dataset, StreamingDataset):
raise TypeError("StreamingDataLoader requires a StreamingDataset")
if type(dataset).__iter__ is not StreamingDataset.__iter__:
raise TypeError(
"StreamingDataLoader does not support StreamingDataset subclasses "
"that override __iter__ because they cannot provide exact "
"per-yield checkpoint state"
)
if kwargs.get("in_order", True) is False:
raise ValueError(
"StreamingDataLoader requires in_order=True for deterministic "
"consumer checkpoints"
)
if kwargs.get("persistent_workers", False):
raise ValueError(
"StreamingDataLoader does not support persistent_workers=True "
"because worker prefetch state cannot be reset from a checkpoint"
)
self._streaming_dataset = dataset
super().__init__(_StreamingDatasetAdapter(dataset), *args, **kwargs)
if self.drop_last:
raise ValueError(
"StreamingDataLoader does not support drop_last=True because "
"discarded worker tails cannot be checkpointed "
"topology-independently"
)
self.collate_fn = _CheckpointCollate(self.collate_fn)
def __iter__(self):
dataset = self._streaming_dataset
previous_lease = dataset._consumer_iterator_lease
owner_token = None
try:
owner_token = dataset._acquire_consumer_iterator()
state = dataset._checkpoint_snapshot()
packed = dataset._pack_sequences is not None
if packed:
blocks = state["blocks_emitted_per_split"]
rank_blocks = [blocks[split] for split in dataset._rank_splits]
if len(set(rank_blocks)) > 1:
raise RuntimeError(
"StreamingDataLoader cannot start from a partial packed "
"logical step; resume from a checkpoint whose splits "
"assigned to this rank have equal emitted-block counts"
)
elif self.num_workers > 1:
samples = state["samples_consumed_per_split"]
rank_samples = [samples[split] for split in dataset._rank_splits]
if len(set(rank_samples)) > 1:
raise RuntimeError(
"StreamingDataLoader cannot start multiple workers from a "
"partial logical step; resume from a checkpoint whose "
"splits assigned to this rank have equal consumed-sample "
"counts"
)
return _ConsumerCommitIterator(
super().__iter__(),
dataset,
owner_token=owner_token,
require_uniform=self.num_workers > 1 or packed,
)
except BaseException:
if owner_token is not None:
dataset._release_consumer_iterator(owner_token)
else:
dataset._release_consumer_iterator_after_failed_acquire(previous_lease)
raise
+3 -3
View File
@@ -1972,7 +1972,7 @@ class Table(ABC):
A mapping with one ``FunctionApplication`` value keeps its scalar
or named-struct result in the named table column. A bare
named-struct application expands its ordered result fields as one
atomic sibling group; aliases come from ``rename(columns=...)``.
atomic binding; aliases come from ``rename(columns=...)``.
Function columns are supported only on LanceDB Cloud and
Enterprise.
computed: Dict[str, str], optional
@@ -6038,7 +6038,7 @@ class AsyncTable:
A mapping with one ``FunctionApplication`` value keeps its scalar
or named-struct result in the named table column. A bare
named-struct application expands its ordered result fields as one
atomic sibling group; aliases come from ``rename(columns=...)``.
atomic binding; aliases come from ``rename(columns=...)``.
Function columns are supported only on LanceDB Cloud and
Enterprise.
computed: Dict[str, str], optional
@@ -6075,7 +6075,7 @@ class AsyncTable:
isinstance(value, FunctionApplication) for value in transforms.values()
):
raise ValueError(
"one add_columns call declares exactly one Function sibling group"
"one add_columns call declares exactly one Function binding"
)
function_output_name, function_application = next(iter(transforms.items()))
@@ -32,6 +32,7 @@ Parameters used throughout:
import dataclasses
import logging
import threading
from unittest.mock import patch
import lancedb
@@ -46,6 +47,7 @@ from utils import (
torch = pytest.importorskip("torch")
streaming = pytest.importorskip("lancedb.streaming")
StreamingDataset = streaming.StreamingDataset
StreamingDataLoader = streaming.StreamingDataLoader
# ---------------------------------------------------------------------------
# Dataset parameters
@@ -92,6 +94,27 @@ class FakeWorkerInfo:
num_workers: int
def _collate_with_first_batch_error(samples):
ids = [sample["id"] for sample in samples]
if ids == [0, 1]:
raise ValueError("first batch fails")
return ids
def _collate_with_first_batch_stop(samples):
ids = [sample["id"] for sample in samples]
if ids == [0, 1]:
raise StopIteration("first batch stopped")
return ids
def _collate_with_first_batch_interrupt(samples):
ids = [sample["id"] for sample in samples]
if ids == [0, 1]:
raise KeyboardInterrupt("first batch interrupted")
return ids
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@@ -1008,6 +1031,565 @@ def test_multi_worker_elastic_det_across_worker_counts(lance_table):
# ── Resumability with num_workers ─────────────────────────────────────────────
def test_streaming_dataloader_commits_only_consumed_worker_batches(tmp_path):
"""Prefetched worker state is committed only as the trainer receives it."""
db = lancedb.connect(tmp_path)
table = db.create_table(
"worker_commit", pa.table({"id": [1, 2, 3, 4, 10, 20, 30, 40]})
)
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
loader = StreamingDataLoader(
dataset,
batch_size=2,
num_workers=2,
multiprocessing_context="spawn",
prefetch_factor=4,
)
iterator = iter(loader)
try:
first = next(iterator)["id"].tolist()
assert first == [1, 2]
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2, 0]
with pytest.raises(RuntimeError, match="complete logical step boundary"):
dataset.state_dict()
second = next(iterator)["id"].tolist()
assert second == [10, 20]
checkpoint = dataset.state_dict()
assert checkpoint["samples_consumed_per_split"] == [2, 2]
uninterrupted = [batch["id"].tolist() for batch in iterator]
finally:
iterator._shutdown_workers()
resumed = StreamingDataset(table, num_splits=2, shuffle=False)
resumed.load_state_dict(checkpoint)
resumed_loader = StreamingDataLoader(
resumed,
batch_size=2,
num_workers=2,
multiprocessing_context="spawn",
prefetch_factor=4,
)
resumed_iterator = iter(resumed_loader)
try:
remaining = [batch["id"].tolist() for batch in resumed_iterator]
finally:
resumed_iterator._shutdown_workers()
assert remaining == uninterrupted == [[3, 4], [30, 40]]
def test_distributed_checkpoint_uses_rank_local_worker_boundary(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("rank_boundary", pa.table({"id": list(range(8))}))
dataset = StreamingDataset(
table,
num_splits=4,
shuffle=False,
rank=0,
world_size=2,
)
loader = StreamingDataLoader(
dataset,
batch_size=1,
num_workers=2,
multiprocessing_context="spawn",
)
iterator = iter(loader)
try:
assert next(iterator)["id"].tolist() == [0]
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [
1,
0,
0,
0,
]
with pytest.raises(RuntimeError, match="complete logical step boundary"):
dataset.state_dict()
assert next(iterator)["id"].tolist() == [2]
checkpoint = dataset.state_dict()
remaining = [batch["id"].tolist() for batch in iterator]
finally:
iterator._shutdown_workers()
assert checkpoint["samples_consumed_per_split"] == [1, 1, 0, 0]
assert remaining == [[1], [3]]
def test_standard_dataloader_rejects_stale_parent_checkpoint(tmp_path):
"""A standard DataLoader must not expose prefetched producer progress."""
db = lancedb.connect(tmp_path)
table = db.create_table("untracked_workers", pa.table({"id": [1, 2, 10, 20]}))
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
# Merely constructing the checkpoint-aware loader must not authorize a
# later plain DataLoader's worker progress.
StreamingDataLoader(dataset, batch_size=2, num_workers=0)
loader = torch.utils.data.DataLoader(
dataset,
batch_size=2,
num_workers=2,
multiprocessing_context="spawn",
)
iterator = iter(loader)
try:
assert next(iterator)["id"].tolist() == [1, 2]
with pytest.raises(RuntimeError, match="Use StreamingDataLoader"):
dataset.state_dict()
list(iterator)
finally:
iterator._shutdown_workers()
def test_streaming_dataloader_rejects_persistent_workers(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("persistent_workers", pa.table({"id": [1, 2]}))
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
with pytest.raises(ValueError, match="persistent_workers=True"):
StreamingDataLoader(
dataset,
batch_size=1,
num_workers=2,
persistent_workers=True,
)
def test_collate_failure_invalidates_consumer_checkpoint(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table(
"collate_failure", pa.table({"id": [0, 1, 2, 3, 100, 101, 102, 103]})
)
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
loader = StreamingDataLoader(
dataset,
batch_size=2,
num_workers=2,
multiprocessing_context="spawn",
collate_fn=_collate_with_first_batch_error,
prefetch_factor=2,
)
iterator = iter(loader)
try:
with pytest.raises(ValueError, match="first batch fails"):
next(iterator)
assert next(iterator) == [100, 101]
assert next(iterator) == [2, 3]
with pytest.raises(RuntimeError, match="failed before it was returned"):
dataset.state_dict()
list(iterator)
finally:
iterator._shutdown_workers()
def test_collate_stop_iteration_invalidates_consumer_checkpoint(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("collate_stop", pa.table({"id": list(range(6))}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(
dataset,
batch_size=2,
num_workers=0,
collate_fn=_collate_with_first_batch_stop,
)
iterator = iter(loader)
with pytest.raises(RuntimeError, match="collate_fn raised StopIteration"):
next(iterator)
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2]
with pytest.raises(RuntimeError, match="failed before it was returned"):
dataset.state_dict()
assert list(iterator) == [[2, 3], [4, 5]]
def test_batch_base_exception_invalidates_consumer_checkpoint(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("collate_interrupt", pa.table({"id": list(range(6))}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(
dataset,
batch_size=2,
num_workers=0,
collate_fn=_collate_with_first_batch_interrupt,
)
iterator = iter(loader)
with pytest.raises(KeyboardInterrupt, match="first batch interrupted"):
next(iterator)
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2]
with pytest.raises(RuntimeError, match="failed before it was returned"):
dataset.state_dict()
assert list(iterator) == [[2, 3], [4, 5]]
def test_parent_commit_base_exception_invalidates_consumer_checkpoint(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("commit_interrupt", pa.table({"id": list(range(4))}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
iterator = iter(loader)
real_commit = dataset._commit_worker_state
def interrupt_after_commit(state, *, require_uniform):
real_commit(state, require_uniform=require_uniform)
raise KeyboardInterrupt("after parent commit")
with patch.object(
dataset, "_commit_worker_state", side_effect=interrupt_after_commit
):
with pytest.raises(KeyboardInterrupt, match="after parent commit"):
next(iterator)
assert dataset._checkpoint_snapshot()["samples_consumed_per_split"] == [2]
with pytest.raises(RuntimeError, match="failed before it was returned"):
dataset.state_dict()
def test_direct_iteration_surfaces_prefetch_failure_before_committing_row(
tmp_path, monkeypatch
):
db = lancedb.connect(tmp_path)
table = db.create_table("prefetch_failure", pa.table({"id": list(range(4))}))
release = threading.Event()
failed = threading.Event()
real_getitems = streaming.Permutation.__getitems__
def controlled_getitems(permutation, indices):
if indices and indices[0] >= 2:
assert release.wait(timeout=5)
failed.set()
raise RuntimeError("later prefetched I/O failed")
return real_getitems(permutation, indices)
class SignalDict(dict):
def __setitem__(self, key, value):
super().__setitem__(key, value)
release.set()
assert failed.wait(timeout=5)
monkeypatch.setattr(streaming.Permutation, "__getitems__", controlled_getitems)
dataset = StreamingDataset(
table,
num_splits=1,
shuffle=False,
read_batch_size=2,
io_queue_depth=2,
)
dataset._resume_positions = SignalDict()
iterator = iter(dataset)
assert next(iterator)["id"] == 0
with pytest.raises(RuntimeError, match="later prefetched I/O failed"):
next(iterator)
checkpoint = dataset.state_dict()
assert checkpoint["samples_consumed_per_split"] == [1]
assert checkpoint["positions_consumed_per_split"] == [1]
@pytest.mark.parametrize("workers", [0, 1, 2])
def test_streaming_dataloader_rejects_drop_last(tmp_path, workers):
db = lancedb.connect(tmp_path)
table = db.create_table("drop_last", pa.table({"id": [0, 1, 2]}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
worker_options = {"multiprocessing_context": "spawn"} if workers else {}
with pytest.raises(ValueError, match="drop_last=True"):
StreamingDataLoader(
dataset,
batch_size=2,
num_workers=workers,
drop_last=True,
**worker_options,
)
def test_streaming_dataloader_owns_one_iterator_until_teardown(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("iterator_owner", pa.table({"id": list(range(4))}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(
dataset,
batch_size=2,
num_workers=1,
multiprocessing_context="spawn",
)
first = iter(loader)
try:
assert next(first)["id"].tolist() == [0, 1]
with pytest.raises(RuntimeError, match="concurrent iteration"):
iter(loader)
finally:
first._shutdown_workers()
second = iter(loader)
try:
assert [batch["id"].tolist() for batch in second] == [[2, 3]]
except BaseException:
second._shutdown_workers()
raise
# Natural exhaustion releases ownership too.
third = iter(loader)
try:
assert list(third) == []
finally:
third._shutdown_workers()
def test_zero_worker_shutdown_closes_inner_iterator_before_release(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("zero_worker_shutdown", pa.table({"id": list(range(6))}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
first = iter(loader)
assert next(first)["id"].tolist() == [0, 1]
first._shutdown_workers()
assert dataset._consumer_iterator_active is False
assert dataset._raw_batches_ref is None
second = iter(loader)
try:
with pytest.raises(StopIteration):
next(first)
assert next(second)["id"].tolist() == [2, 3]
finally:
second._shutdown_workers()
def test_direct_and_loader_admission_share_one_atomic_lease(tmp_path, monkeypatch):
db = lancedb.connect(tmp_path)
table = db.create_table("direct_loader_lease", pa.table({"id": list(range(4))}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
entered = threading.Event()
release = threading.Event()
direct_result = []
direct_error = []
contender = []
real_resolve = dataset._resolve_my_splits
def controlled_resolve():
if threading.current_thread().name == "direct-start":
entered.set()
assert release.wait(timeout=5)
return real_resolve()
def advance_direct(iterator):
try:
direct_result.append(next(iterator)["id"])
except BaseException as exc:
direct_error.append(exc)
monkeypatch.setattr(dataset, "_resolve_my_splits", controlled_resolve)
direct = iter(dataset)
thread = threading.Thread(
target=advance_direct, args=(direct,), name="direct-start"
)
thread.start()
assert entered.wait(timeout=5)
try:
with pytest.raises(RuntimeError, match="concurrent iteration"):
contender.append(iter(loader))
finally:
release.set()
thread.join(timeout=5)
if contender:
contender[0]._shutdown_workers()
direct.close()
assert not thread.is_alive()
assert direct_error == []
assert direct_result == [0]
def test_loader_acquires_before_snapshot_and_cleans_interrupted_acquire(
tmp_path, monkeypatch
):
db = lancedb.connect(tmp_path)
table = db.create_table("lease_snapshot", pa.table({"id": list(range(4))}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(dataset, batch_size=2, num_workers=0)
first = iter(loader)
assert next(first)["id"].tolist() == [0, 1]
entered = threading.Event()
release = threading.Event()
pending = []
pending_errors = []
observed_snapshots = []
real_acquire = dataset._acquire_consumer_iterator
real_snapshot = dataset._checkpoint_snapshot
def controlled_acquire():
if threading.current_thread().name == "stale-start":
entered.set()
assert release.wait(timeout=5)
return real_acquire()
def recording_snapshot():
state = real_snapshot()
if threading.current_thread().name == "stale-start":
observed_snapshots.append(state["samples_consumed_per_split"])
return state
def create_pending_iterator():
try:
pending.append(iter(loader))
except BaseException as exc:
pending_errors.append(exc)
monkeypatch.setattr(dataset, "_acquire_consumer_iterator", controlled_acquire)
monkeypatch.setattr(dataset, "_checkpoint_snapshot", recording_snapshot)
thread = threading.Thread(target=create_pending_iterator, name="stale-start")
thread.start()
assert entered.wait(timeout=5)
assert next(first)["id"].tolist() == [2, 3]
with pytest.raises(StopIteration):
next(first)
release.set()
thread.join(timeout=5)
assert not thread.is_alive()
assert pending_errors == []
assert observed_snapshots == [[4]]
assert len(pending) == 1
assert list(pending[0]) == []
assert dataset.state_dict()["samples_consumed_per_split"] == [4]
def interrupted_acquire():
real_acquire()
raise KeyboardInterrupt("after acquire")
monkeypatch.setattr(dataset, "_acquire_consumer_iterator", interrupted_acquire)
with pytest.raises(KeyboardInterrupt, match="after acquire"):
iter(loader)
assert dataset._consumer_iterator_active is False
def test_consumer_iterator_lease_publication_is_atomic(tmp_path, monkeypatch):
db = lancedb.connect(tmp_path)
table = db.create_table("atomic_lease", pa.table({"id": [0, 1]}))
dataset = StreamingDataset(table, num_splits=1, shuffle=False)
loader = StreamingDataLoader(dataset, batch_size=1, num_workers=0)
real_get_ident = streaming.threading.get_ident
calls = 0
def interrupt_during_publication():
nonlocal calls
calls += 1
if calls == 1:
raise KeyboardInterrupt("during lease mutation")
return real_get_ident()
monkeypatch.setattr(streaming.threading, "get_ident", interrupt_during_publication)
with pytest.raises(KeyboardInterrupt, match="during lease mutation"):
iter(loader)
monkeypatch.setattr(streaming.threading, "get_ident", real_get_ident)
assert dataset._consumer_iterator_active is False
iterator = iter(loader)
try:
assert next(iterator)["id"].tolist() == [0]
finally:
iterator._shutdown_workers()
def test_streaming_dataloader_rejects_dataset_iter_override(tmp_path):
class CustomizedDataset(StreamingDataset):
def __iter__(self):
return iter([1000, 1001])
db = lancedb.connect(tmp_path)
table = db.create_table("custom_iteration", pa.table({"id": [0, 1, 2]}))
dataset = CustomizedDataset(table, num_splits=1, shuffle=False)
assert list(dataset) == [1000, 1001]
with pytest.raises(TypeError, match="override __iter__"):
StreamingDataLoader(
dataset,
batch_size=2,
num_workers=0,
collate_fn=list,
)
def test_interleaved_adapters_do_not_authorize_plain_iteration(tmp_path):
db = lancedb.connect(tmp_path)
table_a = db.create_table("adapter_a", pa.table({"id": [0, 1]}))
table_b = db.create_table("adapter_b", pa.table({"id": [10, 11]}))
dataset_a = StreamingDataset(table_a, num_splits=1, shuffle=False)
dataset_b = StreamingDataset(table_b, num_splits=1, shuffle=False)
initial_state = dataset_a.state_dict()
owner_a = dataset_a._acquire_consumer_iterator()
owner_b = dataset_b._acquire_consumer_iterator()
try:
iterator_a = iter(streaming._StreamingDatasetAdapter(dataset_a))
iterator_b = iter(streaming._StreamingDatasetAdapter(dataset_b))
assert next(iterator_a).data["id"] == 0
assert next(iterator_b).data["id"] == 10
assert [sample.data["id"] for sample in iterator_a] == [1]
assert [sample.data["id"] for sample in iterator_b] == [11]
finally:
dataset_a._release_consumer_iterator(owner_a)
dataset_b._release_consumer_iterator(owner_b)
dataset_a.load_state_dict(initial_state)
with patch(
"lancedb.streaming.get_worker_info",
return_value=FakeWorkerInfo(id=0, num_workers=1),
):
plain_iterator = iter(dataset_a)
assert next(plain_iterator)["id"] == 0
plain_iterator.close()
assert dataset_a._untracked_worker_iteration[0] == 1
with pytest.raises(RuntimeError, match="Use StreamingDataLoader"):
dataset_a.state_dict()
def test_resume_from_partial_split_cycle_preserves_remaining_order(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table("partial_cycle", pa.table({"id": [1, 2, 10, 20]}))
dataset = StreamingDataset(table, num_splits=2, shuffle=False)
iterator = iter(dataset)
assert next(iterator)["id"] == 1
checkpoint = dataset.state_dict()
iterator.close()
assert checkpoint["samples_consumed_per_split"] == [1, 0]
resumed = StreamingDataset(table, num_splits=2, shuffle=False)
resumed.load_state_dict(checkpoint)
assert [row["id"] for row in resumed] == [10, 2, 20]
def test_partial_cycle_resume_preserves_skip_truncation(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table(
"partial_skip", pa.table({"id": [0, 1, 2, 3, 100, 101, 102, 103]})
)
kwargs = dict(
num_splits=2,
shuffle=False,
transform=_failing_transform({1, 2, 3}),
on_transform_error="skip",
)
dataset = StreamingDataset(table, **kwargs)
iterator = iter(dataset)
assert next(iterator)["id"] == 0
checkpoint = dataset.state_dict()
uninterrupted = [row["id"] for row in iterator]
resumed = StreamingDataset(table, **kwargs)
resumed.load_state_dict(checkpoint)
assert [row["id"] for row in resumed] == uninterrupted == [100]
def test_multi_worker_resumability_same_topology(lance_table):
"""Checkpoint with num_workers=2, resume with num_workers=2: exact continuation."""
world_size = 1
@@ -2018,6 +2600,23 @@ def test_merge_state_dicts_validates_consistency(lance_table):
StreamingDataset.merge_state_dicts([])
def test_merge_state_dicts_combines_nonuniform_consumer_progress(lance_table):
dataset = StreamingDataset(
lance_table, num_splits=2, shuffle=False, shuffle_seed=SHUFFLE_SEED
)
rank0 = dataset.state_dict()
rank0["samples_consumed_per_split"] = [2, 0]
rank0["positions_consumed_per_split"] = [2, 0]
rank1 = dataset.state_dict()
rank1["samples_consumed_per_split"] = [0, 2]
rank1["positions_consumed_per_split"] = [0, 2]
merged = StreamingDataset.merge_state_dicts([rank0, rank1])
assert merged["samples_consumed_per_split"] == [2, 2]
assert merged["positions_consumed_per_split"] == [2, 2]
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)."""
@@ -2254,6 +2853,65 @@ def test_pack_sequences_checkpoint_resumes_on_new_topology(tmp_path):
]
def test_packed_checkpoint_requires_complete_split_cycle(tmp_path):
table = _create_token_table(tmp_path, [[1], [2], [10], [20]])
dataset = _packed_dataset(table, pack_sequences=3, blocks_per_epoch=4, num_splits=2)
iterator = iter(dataset)
next(iterator)
with pytest.raises(RuntimeError, match="complete logical step boundary"):
dataset.state_dict()
next(iterator)
assert dataset.state_dict()["blocks_emitted_per_split"] == [1, 1]
iterator.close()
def test_streaming_dataloader_commits_consumed_packed_batches(tmp_path):
table = _create_token_table(
tmp_path,
[[1], [2], [3], [4], [10], [20], [30], [40]],
)
kwargs = dict(pack_sequences=4, blocks_per_epoch=4, num_splits=2)
dataset = _packed_dataset(table, **kwargs)
loader = StreamingDataLoader(
dataset,
batch_size=1,
num_workers=2,
multiprocessing_context="spawn",
prefetch_factor=2,
)
iterator = iter(loader)
try:
next(iterator)
with pytest.raises(RuntimeError, match="complete logical step boundary"):
dataset.state_dict()
next(iterator)
checkpoint = dataset.state_dict()
uninterrupted = [batch["input_ids"].tolist() for batch in iterator]
finally:
iterator._shutdown_workers()
resumed = _packed_dataset(table, **kwargs)
resumed.load_state_dict(checkpoint)
resumed_loader = StreamingDataLoader(
resumed,
batch_size=1,
num_workers=2,
multiprocessing_context="spawn",
prefetch_factor=2,
)
resumed_iterator = iter(resumed_loader)
try:
remaining = [batch["input_ids"].tolist() for batch in resumed_iterator]
finally:
resumed_iterator._shutdown_workers()
assert checkpoint["blocks_emitted_per_split"] == [1, 1]
assert remaining == uninterrupted
def test_pack_sequences_validates_configuration_and_tokens(tmp_path):
table = _create_token_table(tmp_path, [[1, 2]])
@@ -121,7 +121,7 @@ def test_function_version_identity_is_immutable_and_exact():
assert FunctionVersion(**changed) != version
def test_function_version_binds_named_columns_as_one_immutable_group():
def test_function_version_binds_named_columns_as_one_immutable_application():
version = FunctionVersion.from_json(
json.dumps(job_result("remote_function_job.json"))
)
@@ -131,13 +131,10 @@ def test_function_version_binds_named_columns_as_one_immutable_group():
assert application.function.name == version.name
assert application.function.version == version.version
assert application.output is version.signature.output
assert application.group_id.startswith("fg_")
assert [
(value.parameter, value.kind, value.value["path"])
for value in application.inputs
] == [("text", "column", "documents.body")]
with pytest.raises((TypeError, ValueError)):
application.group_id = "fg_changed"
def test_function_version_binding_validates_names_and_direct_columns():
@@ -156,7 +153,7 @@ def test_function_version_binding_validates_names_and_direct_columns():
def test_function_version_keeps_named_struct_outputs_in_one_application():
value = job_result("remote_function_job.json")
value["name"] = "text_features"
value["version"] = "fv_grouped"
value["version"] = "fv_multi_output"
value["signature"] = {
"inputs": [
{"name": "title", "arrow_type": "utf8", "nullable": True},
@@ -221,7 +218,6 @@ def test_function_application_uses_rename_columns_only():
assert application.columns["normalized_text"] == "search_text"
assert renamed.columns["normalized_text"] == "body_normalized"
assert renamed.function == application.function
assert renamed.group_id == application.group_id
assert not hasattr(application, "rename_outputs")
with pytest.raises(TypeError, match="immutable"):
renamed.columns["normalized_text"] = "changed"
@@ -242,7 +238,6 @@ def test_function_application_uses_rename_columns_only():
def test_binding_and_refresh_result_keep_stable_remote_fields():
binding = FunctionBinding.from_json(fixture("remote_function_binding.json"))
assert binding.revision == 3
assert binding.function.version == "fv_01K3TEXT"
assert [output.output_ordinal for output in binding.outputs] == [0, 1]
assert binding.input_schema is not None
@@ -322,7 +317,7 @@ def known_application() -> FunctionApplication:
@pytest.mark.asyncio
async def test_add_columns_routes_struct_as_one_and_grouped_expansion_atomically():
async def test_add_columns_routes_struct_as_one_and_multi_output_binding_atomically():
inner = _FunctionDeclarationInner()
table = AsyncTable(inner)
application = known_application()
@@ -343,12 +338,12 @@ async def test_add_columns_routes_struct_as_one_and_grouped_expansion_atomically
@pytest.mark.asyncio
async def test_add_columns_rejects_mixed_groups_and_unknown_newer_application():
async def test_add_columns_rejects_multiple_bindings_and_unknown_newer_application():
inner = _FunctionDeclarationInner()
table = AsyncTable(inner)
application = known_application()
with pytest.raises(ValueError, match="exactly one Function sibling group"):
with pytest.raises(ValueError, match="exactly one Function binding"):
await table.add_columns({"a": application, "b": application})
future = json.loads(fixture("remote_function_application.json"))
@@ -376,7 +371,6 @@ def test_rename_requires_named_struct_and_keeps_partial_mapping_immutable():
"arrow_type": "list<float32>",
"nullable": False,
},
"group_id": "fg_scalar",
}
)
)
@@ -3,7 +3,12 @@
from __future__ import annotations
import base64
import contextlib
import functools
import importlib.util
import types
from datetime import date
import http.server
import json
from pathlib import Path
@@ -16,6 +21,9 @@ import pytest
import lancedb
from lancedb.functions import UdfDefinition, udf
THRESHOLD = 20
_CACHE = None
FIXTURES = (
Path(__file__).parents[3]
@@ -71,9 +79,396 @@ def test_scalar_udf_matches_shared_registration_golden_and_remains_callable():
_assert_no_secret_values(request)
def _run_packaged(definition, *args):
"""Execute the shipped artifact in a fresh namespace, as a worker would."""
source = base64.b64decode(definition.registration_request.artifact.content.data)
namespace: dict = {}
exec(compile(source, "<udf>", "exec"), namespace)
return namespace[definition.registration_request.artifact.entrypoint](*args)
def test_udf_packages_attribute_access_and_body_imports():
@udf
def word_norm(body: str) -> float:
import numpy as np
try:
words = body.split()
except AttributeError as error:
raise ValueError(str(error)) from error
return float(np.linalg.norm([len(w) for w in words]))
assert _run_packaged(word_norm, "aa bb") == pytest.approx(8**0.5)
def test_udf_packages_module_globals_and_global_caches():
@udf
def label(value: int) -> str:
return "big" if value >= THRESHOLD else "small"
assert _run_packaged(label, 21) == "big"
@udf
def cached(value: int) -> int:
global _CACHE
if _CACHE is None:
_CACHE = 40
return _CACHE + value
assert _run_packaged(cached, 2) == 42
def test_udf_annotations_are_not_runtime_names():
@udf
def identity(value: date) -> date:
return value
assert _run_packaged(identity, date(2026, 8, 25)) == date(2026, 8, 25)
def test_udf_nested_scopes_resolve_lexically():
@udf
def score(value: int) -> int:
offset = 2
def add_offset() -> int:
return value + offset
return add_offset() + sum(v for v in [0])
assert _run_packaged(score, 3) == 5
def test_udf_resolves_module_globals_before_builtins(tmp_path):
module_path = tmp_path / "shadowing_udfs.py"
module_path.write_text(
"max = 7\n"
"len = lambda _: 99\n"
"\n"
"def uses_literal_shadow(value: int) -> int:\n"
" def nested() -> int:\n"
" return max\n"
" return nested() + value\n"
"\n"
"def uses_callable_shadow(value: int) -> int:\n"
" def nested() -> int:\n"
" return len([1])\n"
" return nested() + value\n"
)
spec = importlib.util.spec_from_file_location("shadowing_udfs", module_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
# The module's `max = 7` is what the interpreter would use, so it ships.
assert _run_packaged(udf(module.uses_literal_shadow), 1) == 8
# A callable global cannot ship; it must not be silently swapped for the builtin.
with pytest.raises(TypeError, match="unsupported global value of type function"):
udf(module.uses_callable_shadow)
def test_canonical_arrow_type_is_exactly_the_grammar():
from lancedb.functions import _GRAMMAR_PRIMITIVES, _canonical_arrow_type
golden = json.loads(
(
Path(__file__).parents[3]
/ "rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json"
).read_text()
)
primitives = [
case["arrow_type"] for case in golden["valid"] if "<" not in case["arrow_type"]
]
assert [name for _, name in _GRAMMAR_PRIMITIVES] == primitives
for outside in [
pa.timestamp("us"),
pa.decimal128(10, 2),
pa.large_string(),
pa.large_binary(),
pa.binary(4),
pa.duration("s"),
pa.struct([pa.field("a", pa.int32())]),
pa.list_(pa.float32(), 0),
pa.list_(pa.timestamp("us")),
]:
with pytest.raises(TypeError, match="unsupported Arrow type"):
_canonical_arrow_type(outside)
def test_udf_nested_annotations_are_postponed_in_the_artifact():
@udf
def score(value: int) -> int:
def identity(item: date) -> date:
return item
identity(date(2026, 8, 25))
return value
assert _run_packaged(score, 3) == 3
def test_udf_ships_globals_the_body_deletes():
@udf
def clear(value: int) -> int:
global _CACHE
del _CACHE
return value
assert _run_packaged(clear, 3) == 3
def test_udf_rejects_a_module_global_that_does_not_import_as_itself(tmp_path):
module_path = tmp_path / "fake_module_udfs.py"
module_path.write_text(
"import types\n"
"np = types.ModuleType('numpy')\n"
"np.sqrt = lambda x: 0\n"
"\n"
"def score(value: int) -> int:\n"
" return int(np.sqrt(value))\n"
)
spec = importlib.util.spec_from_file_location("fake_module_udfs", module_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
with pytest.raises(TypeError, match="does not import as 'numpy'"):
udf(module.score)
def test_udf_rejects_a_module_level_namespace_alias(tmp_path):
module_path = tmp_path / "aliasing_udfs.py"
module_path.write_text(
"import builtins as b\n"
"THRESHOLD = 5\n"
"\n"
"def score(value: int) -> int:\n"
" return value + b.vars(b.__import__('aliasing_udfs'))['THRESHOLD']\n"
)
spec = importlib.util.spec_from_file_location("aliasing_udfs", module_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
with pytest.raises(ValueError, match="dynamic namespace access"):
udf(module.score)
@pytest.mark.parametrize(
"access",
[
"globals()['THRESHOLD']",
"eval('THRESHOLD')",
"(lambda g: g()['THRESHOLD'])(globals)",
"__import__('sys').modules[__name__].THRESHOLD",
"sys.modules[__name__].THRESHOLD",
],
)
def test_udf_rejects_dynamic_namespace_access(access):
namespace: dict = {}
exec(
f"def score(value: int) -> int:\n return value + {access}\n",
{"THRESHOLD": 5},
namespace,
)
with pytest.raises(ValueError, match="dynamic namespace access"):
_package_from_text(
"def score(value: int) -> int:\n"
" import sys\n"
f" return value + {access}\n"
)
def _package_from_text(source: str, module_globals: dict | None = None):
"""Load `source` as a real module file so the packager can inspect it."""
import tempfile
directory = tempfile.mkdtemp()
path = Path(directory) / "generated_udf_module.py"
path.write_text(source)
spec = importlib.util.spec_from_file_location(f"generated_udf_{id(source)}", path)
module = importlib.util.module_from_spec(spec)
if module_globals:
module.__dict__.update(module_globals)
spec.loader.exec_module(module)
functions = [
value
for value in vars(module).values()
if callable(value) and getattr(value, "__module__", None) == module.__name__
]
return udf(functions[0])
def test_udf_rejects_a_non_standard_builtins_environment():
def score(value: int) -> int:
return len([1]) + value
score.__globals__ # noqa: B018 -- real function, real globals
import builtins
patched = types.FunctionType(
score.__code__,
{"__builtins__": {**vars(builtins), "len": lambda _: 99}},
"score",
)
patched.__annotations__ = score.__annotations__
assert patched(3) == 102
with pytest.raises(ValueError, match="non-standard builtins environment"):
udf(patched)
class ReportingDict(dict): # reports standard entries, resolves differently
def __missing__(self, key):
return vars(builtins)[key]
disguised = types.FunctionType(
score.__code__, {"__builtins__": ReportingDict(len=lambda _: 99)}, "score"
)
disguised.__annotations__ = score.__annotations__
assert disguised(3) == 102
with pytest.raises(ValueError, match="non-standard builtins environment"):
udf(disguised)
hooked = types.FunctionType(
score.__code__,
{"__builtins__": {**vars(builtins), "__import__": lambda *a, **k: None}},
"score",
)
hooked.__annotations__ = score.__annotations__
with pytest.raises(ValueError, match="non-standard builtins environment"):
udf(hooked)
def test_udf_recursion_versus_a_rebound_module_name(tmp_path):
module_path = tmp_path / "rebound_udfs.py"
module_path.write_text(
"def fact(value: int) -> int:\n"
" return 1 if value <= 1 else value * fact(value - 1)\n"
"\n"
"def score(value: int) -> int:\n"
" return score + value\n"
)
spec = importlib.util.spec_from_file_location("rebound_udfs", module_path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
assert _run_packaged(udf(module.fact), 5) == 120
raw = module.score
module.score = 10
with pytest.raises(ValueError, match="binds that name to another value"):
udf(raw)
# A wrapper that merely exposes __wrapped__ is not the function.
module.score = functools.wraps(raw)(lambda value: 41)
with pytest.raises(ValueError, match="binds that name to another value"):
udf(raw)
# The decorator's own result is; a subclass of it is not.
module.fact = udf(module.fact)
assert _run_packaged(module.fact, 4) == 24
class Twisted(UdfDefinition):
def __call__(self, *args, **kwargs):
return 41
raw_fact = module.fact._function
module.fact = Twisted(
raw_fact,
name=None,
input_schema=None,
output_schema=None,
pip=(),
env={},
secrets=(),
python_version=None,
)
with pytest.raises(ValueError, match="binds that name to another value"):
udf(raw_fact)
def test_canonical_arrow_type_rejects_unrepresentable_list_children():
from lancedb.functions import _canonical_arrow_type
for outside in [
pa.list_(pa.float32()), # pyarrow default: nullable child
pa.list_(pa.field("custom", pa.float32(), nullable=False)),
pa.list_(pa.field("item", pa.float32(), nullable=False, metadata={"k": "v"})),
pa.list_(pa.field("item", pa.float32(), nullable=False), 0),
]:
with pytest.raises(TypeError, match="unsupported Arrow type"):
_canonical_arrow_type(outside)
assert (
_canonical_arrow_type(
pa.list_(pa.field("item", pa.float32(), nullable=False), 3)
)
== "fixed_size_list<float32, 3>"
)
def _calls_missing(value: int) -> int:
return missing(value) # noqa: F821
def _shadows_missing_in_a_comprehension(value: int) -> int:
return missing(value) + sum(missing for missing in ()) # noqa: F821
def _shadows_missing_in_a_lambda(value: int) -> int:
return (lambda missing: missing)(value) + missing # noqa: F821
@pytest.mark.parametrize(
"function",
[_calls_missing, _shadows_missing_in_a_comprehension, _shadows_missing_in_a_lambda],
)
def test_udf_rejects_a_truly_unresolved_global(function):
with pytest.raises(ValueError, match=r"unresolved global names: \['missing'\]"):
udf(function)
def _arrow_type_from_golden(spec: dict) -> pa.DataType:
kind = spec["type"]
if kind in ("list", "large_list", "fixed_size_list"):
item = _arrow_type_from_golden(spec["fields"][0]["type"])
field = pa.field("item", item, nullable=False)
if kind == "list":
return pa.list_(field)
if kind == "large_list":
return pa.large_list(field)
return pa.list_(field, spec["length"])
return {
"null": pa.null(),
"bool": pa.bool_(),
"utf8": pa.string(),
"binary": pa.binary(),
"float16": pa.float16(),
"float32": pa.float32(),
"float64": pa.float64(),
"date32": pa.date32(),
"date64": pa.date64(),
}.get(kind) or getattr(pa, kind)()
def test_arrow_type_grammar_matches_the_shared_golden():
golden = json.loads(
(
Path(__file__).parents[3]
/ "rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json"
).read_text()
)
from lancedb.functions import _canonical_arrow_type
emitted = {
case["arrow_type"]: _canonical_arrow_type(_arrow_type_from_golden(case["json"]))
for case in golden["valid"]
}
assert emitted == {
case["arrow_type"]: case["arrow_type"] for case in golden["valid"]
}
assert not set(emitted) & set(golden["invalid"])
for case in golden["server_only"]:
with pytest.raises(TypeError, match="unsupported Arrow type"):
_canonical_arrow_type(_arrow_type_from_golden(case["json"]))
def test_explicit_arrow_schema_is_deterministic():
input_schema = pa.schema([pa.field("value", pa.float32(), nullable=True)])
output_schema = pa.field("embedding", pa.list_(pa.float32(), 3), nullable=False)
output_schema = pa.field(
"embedding",
pa.list_(pa.field("item", pa.float32(), nullable=False), 3),
nullable=False,
)
@udf(input_schema=input_schema, output_schema=output_schema)
def explicit(value):
@@ -82,7 +477,7 @@ def test_explicit_arrow_schema_is_deterministic():
signature = explicit.registration_request.signature
assert signature.inputs[0].arrow_type == "float32"
assert signature.inputs[0].nullable is True
assert signature.output.arrow_type == "fixed_size_list<float32>[3]"
assert signature.output.arrow_type == "fixed_size_list<float32, 3>"
assert signature.output.nullable is False
+25
View File
@@ -56,6 +56,31 @@ def test_execute_does_not_reenter_background_loop(tmp_path, monkeypatch):
assert permutation_tbl._conn.read_consistency_interval is None
def test_pickled_permutation_reads_pinned_version(tmp_path):
"""An unpickled copy must still read the pinned version, which also covers the
version surviving the ``to_arrow()`` round trip in ``__getstate__``."""
import pickle
db = connect(tmp_path)
tbl = db.create_table("base", pa.table({"idx": range(20)}))
permutation_tbl = permutation_builder(tbl).execute()
perm = Permutation.from_tables(tbl, permutation_tbl)
payload = pickle.dumps(perm)
# Compact so the stored row addresses no longer describe these rows at latest.
tbl.delete("true")
tbl.optimize()
assert tbl.count_rows() == 0
# Unpickle after the mutation: __setstate__ reopens at latest, so this only
# passes if the recorded version is applied on reopen.
restored = pickle.loads(payload)
assert len(restored) == 20
rows = restored.__getitems__(list(range(20)))
assert sorted(row["idx"] for row in rows) == list(range(20))
def test_split_random_counts(mem_db):
"""Test random splitting with absolute counts."""
tbl = mem_db.create_table(
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "lancedb"
version = "0.38.0-beta.6"
version = "0.38.0-beta.9"
edition.workspace = true
description = "LanceDB: A serverless, low-latency vector database for AI applications"
license.workspace = true
+44
View File
@@ -1679,6 +1679,50 @@ mod tests {
assert_eq!(tables, names[..7]);
}
#[tokio::test]
async fn test_list_tables_walks_page_boundaries() {
let tc = new_test_connection().await.unwrap();
if tc.is_remote {
// What resumes a page is the server's to decide, and asserting it here would be
// asserting the server's contract rather than this one.
return;
}
let db = tc.connection;
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int32, false)]));
let mut names = Vec::with_capacity(5);
for _ in 0..5 {
let name = uuid::Uuid::new_v4().to_string();
names.push(name.clone());
db.create_empty_table(name, schema.clone())
.execute()
.await
.unwrap();
}
names.sort();
// Walking in pages has to reach every table exactly once, with nothing lost at a
// page boundary.
let mut seen = Vec::with_capacity(names.len());
let mut page_token = None;
loop {
let page = db
.list_tables(ListTablesRequest {
id: Some(Vec::new()),
limit: Some(2),
page_token,
..Default::default()
})
.await
.unwrap();
seen.extend(page.tables);
page_token = page.page_token.filter(|token| !token.is_empty());
if page_token.is_none() {
break;
}
}
assert_eq!(seen, names);
}
#[tokio::test]
async fn test_open_table() {
let tc = new_test_connection().await.unwrap();
+7 -9
View File
@@ -974,17 +974,15 @@ impl Database for ListingDatabase {
f.drain(0..index);
}
// Determine if there's a next page
let next_page_token = if let Some(limit) = request.limit {
if f.len() > limit as usize {
let token = f[limit as usize].clone();
// Determine if there's a next page. The token is the last name of this page,
// not the first of the next one: the next page resumes strictly after the
// token, so naming the next page's first entry would skip it.
let next_page_token = match request.limit {
Some(limit) if f.len() > limit as usize => {
f.truncate(limit as usize);
Some(token)
} else {
None
f.last().cloned()
}
} else {
None
_ => None,
};
Ok(ListTablesResponse {
@@ -27,6 +27,12 @@ pub const SRC_ROW_ID_COL: &str = "row_id";
pub const SPLIT_NAMES_CONFIG_KEY: &str = "split_names";
/// Base table version the permutation was built against.
pub const BASE_VERSION_CONFIG_KEY: &str = "base_version";
/// Base table branch the permutation was built against. Absent means main.
pub const BASE_BRANCH_CONFIG_KEY: &str = "base_branch";
pub const DEFAULT_MEMORY_LIMIT: usize = 100 * 1024 * 1024;
/// Where to store the permutation table
@@ -214,21 +220,11 @@ impl PermutationBuilder {
Ok(Box::pin(SimpleRecordBatchStream { schema, stream }))
}
fn add_split_names(
fn add_config_metadata(
data: SendableRecordBatchStream,
split_names: &[String],
metadata: HashMap<String, String>,
) -> Result<SendableRecordBatchStream> {
let schema = data
.schema()
.as_ref()
.clone()
.with_metadata(HashMap::from([(
SPLIT_NAMES_CONFIG_KEY.to_string(),
serde_json::to_string(split_names).map_err(|e| Error::Other {
message: format!("Failed to serialize split names: {}", e),
source: Some(e.into()),
})?,
)]));
let schema = data.schema().as_ref().clone().with_metadata(metadata);
let schema = Arc::new(schema);
let schema_clone = schema.clone();
let stream = data.map_ok(move |batch| batch.with_schema(schema.clone()).unwrap());
@@ -239,7 +235,20 @@ impl PermutationBuilder {
}
/// Builds the permutation table and stores it in the given database.
pub async fn build(self) -> Result<Table> {
pub async fn build(mut self) -> Result<Table> {
// Remote tables resolve latest independently for each request. Use a
// separate pinned handle so count, projection, and scan all refer to one
// snapshot without changing the caller's table checkout state. Native
// tables return `None` here and retain their existing behavior.
if let Some(snapshot) = self
.base_table
.base_table()
.snapshot_at_current_version()
.await?
{
self.base_table = Table::from(snapshot);
}
// Unflushed rows have no row id, so a permutation cannot address them.
match self.base_table.base_table().get_lsm_write_spec().await {
Ok(Some(_)) => {
@@ -256,9 +265,14 @@ impl PermutationBuilder {
Err(err) => return Err(err),
}
// The handle above is already pinned to one version. Record which one, so a
// reader -- in a DataLoader worker, against a table that has since moved --
// resolves these row addresses against the same snapshot.
let base_version = self.base_table.version().await?;
let base_branch = self.base_table.current_branch();
// First pass, apply filter and load row ids. `Shuffler` permutes positions, so
// every rank must scan the rows in the same order to build the same permutation.
// TODO: pin the version resolved here; remote does not implement Lazy.
let mut rows = self.base_table.query().select(Select::columns(&[ROW_ID]));
if let Some(filter) = &self.config.filter {
@@ -318,11 +332,24 @@ impl PermutationBuilder {
// Rename _rowid to row_id
let renamed = rename_column(sorted, ROW_ID, SRC_ROW_ID_COL)?;
let streaming_data = if let Some(split_names) = &self.config.split_names {
Self::add_split_names(renamed, split_names)?
} else {
renamed
};
let mut metadata = HashMap::from([(
BASE_VERSION_CONFIG_KEY.to_string(),
base_version.to_string(),
)]);
// Version numbers are per-branch, so the branch is part of the coordinate.
if let Some(branch) = &base_branch {
metadata.insert(BASE_BRANCH_CONFIG_KEY.to_string(), branch.clone());
}
if let Some(split_names) = &self.config.split_names {
metadata.insert(
SPLIT_NAMES_CONFIG_KEY.to_string(),
serde_json::to_string(split_names).map_err(|e| Error::Other {
message: format!("Failed to serialize split names: {}", e),
source: Some(e.into()),
})?,
);
}
let streaming_data = Self::add_config_metadata(renamed, metadata)?;
let (name, database) = match &self.config.destination {
PermutationDestination::Permanent(database, table_name) => {
@@ -409,6 +436,253 @@ mod tests {
assert!(table.base_table().scan_order_is_deterministic());
}
#[cfg(feature = "remote")]
#[tokio::test]
async fn test_remote_permutation_builder_pins_snapshot() {
use std::sync::{
Mutex,
atomic::{AtomicU64, Ordering},
};
use arrow_array::{RecordBatch, UInt64Array};
use arrow_schema::{DataType, Field, Schema};
let row_ids = RecordBatch::try_new(
Arc::new(Schema::new(vec![Field::new(
ROW_ID,
DataType::UInt64,
false,
)])),
vec![Arc::new(UInt64Array::from(vec![100]))],
)
.unwrap();
let mut query_body = Vec::new();
{
let mut writer =
arrow_ipc::writer::FileWriter::try_new(&mut query_body, &row_ids.schema()).unwrap();
writer.write(&row_ids).unwrap();
writer.finish().unwrap();
}
let latest = Arc::new(AtomicU64::new(7));
let expected_snapshot = Arc::new(AtomicU64::new(7));
let planning_versions = Arc::new(Mutex::new(Vec::new()));
let latest_ref = latest.clone();
let expected_snapshot_ref = expected_snapshot.clone();
let planning_versions_ref = planning_versions.clone();
let table = Table::new_with_handler("remote_base", move |request| {
let path = request.url().path();
let body = request
.body()
.and_then(|body| body.as_bytes())
.map(|body| serde_json::from_slice::<serde_json::Value>(body).unwrap());
match path {
"/v1/table/remote_base/describe/" => {
let requested = body.as_ref().and_then(|body| body["version"].as_u64());
let version = requested.unwrap_or_else(|| latest_ref.load(Ordering::SeqCst));
http::Response::builder()
.status(200)
.body(
format!(r#"{{"version":{version},"schema":{{"fields":[]}}}}"#)
.into_bytes(),
)
.unwrap()
}
"/v1/table/remote_base/get_lsm_write_spec/" => http::Response::builder()
.status(200)
.body(br#"{"lsm_write_spec":null}"#.to_vec())
.unwrap(),
"/v1/table/remote_base/count_rows/" => {
let body = body.unwrap();
let version = body["version"].as_u64().unwrap();
assert_eq!(version, expected_snapshot_ref.load(Ordering::SeqCst));
assert_eq!(body["predicate"], "value > 0");
planning_versions_ref.lock().unwrap().push(version);
// Simulate a concurrent append after count_rows. An unpinned
// scan would now resolve version 8 and include different rows.
latest_ref.store(8, Ordering::SeqCst);
http::Response::builder()
.status(200)
.body(b"1".to_vec())
.unwrap()
}
"/v1/table/remote_base/query/" => {
let body = body.unwrap();
let version = body["version"].as_u64().unwrap();
assert_eq!(version, expected_snapshot_ref.load(Ordering::SeqCst));
assert_eq!(body["filter"], "value > 0");
assert_eq!(body["columns"], serde_json::json!([ROW_ID]));
planning_versions_ref.lock().unwrap().push(version);
http::Response::builder()
.status(200)
.header("content-type", "application/vnd.apache.arrow.file")
.body(query_body.clone())
.unwrap()
}
_ => panic!("unexpected request: {path}"),
}
});
let permutation = PermutationBuilder::new(table.clone())
.with_filter("value > 0".to_string())
.build()
.await
.unwrap();
assert_eq!(permutation.count_rows(None).await.unwrap(), 1);
// Building uses a separate handle and must not pin the caller's table.
assert_eq!(table.version().await.unwrap(), 8);
// An explicit checkout is copied as-is and remains checked out afterward.
expected_snapshot.store(6, Ordering::SeqCst);
table.checkout(6).await.unwrap();
let permutation = PermutationBuilder::new(table.clone())
.with_filter("value > 0".to_string())
.build()
.await
.unwrap();
assert_eq!(permutation.count_rows(None).await.unwrap(), 1);
assert_eq!(table.version().await.unwrap(), 6);
assert_eq!(*planning_versions.lock().unwrap(), vec![7, 7, 6, 6]);
}
#[tokio::test]
async fn test_permutation_records_base_version() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(100), BatchCount::from(2));
let data_table = db
.create_table("base_tbl", initial_data)
.execute()
.await
.unwrap();
let build_version = data_table.version().await.unwrap();
let permutation_table = PermutationBuilder::new(data_table.clone())
.build()
.await
.unwrap();
let recorded = permutation_table
.schema()
.await
.unwrap()
.metadata
.get(BASE_VERSION_CONFIG_KEY)
.expect("permutation should record the base version")
.parse::<u64>()
.unwrap();
assert_eq!(recorded, build_version);
// Advancing the base table must not move the recorded version.
let more_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(50), BatchCount::from(1));
data_table.add(more_data).execute().await.unwrap();
assert!(data_table.version().await.unwrap() > recorded);
assert_eq!(
permutation_table
.schema()
.await
.unwrap()
.metadata
.get(BASE_VERSION_CONFIG_KEY)
.unwrap()
.parse::<u64>()
.unwrap(),
recorded,
);
}
/// Version numbers are per-branch, so a permutation built on a branch must record
/// it -- a worker reopens by name and lands on main at the same number.
#[tokio::test]
async fn test_permutation_records_base_branch() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(10), BatchCount::from(1));
let data_table = db
.create_table("base_tbl", initial_data)
.execute()
.await
.unwrap();
let branch = data_table
.create_branch("exp", lance::dataset::refs::Ref::from(("main", 1)))
.await
.unwrap();
let permutation_table = PermutationBuilder::new(branch.clone())
.build()
.await
.unwrap();
let metadata = permutation_table.schema().await.unwrap().metadata.clone();
assert_eq!(
metadata.get(BASE_BRANCH_CONFIG_KEY).map(String::as_str),
Some("exp")
);
// Main records nothing, so an absent key keeps meaning main.
let main_permutation = PermutationBuilder::new(data_table.clone())
.build()
.await
.unwrap();
assert!(
!main_permutation
.schema()
.await
.unwrap()
.metadata
.contains_key(BASE_BRANCH_CONFIG_KEY)
);
}
#[tokio::test]
async fn test_build_does_not_pin_the_callers_table() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let initial_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(100), BatchCount::from(1));
let data_table = db
.create_table("base_tbl", initial_data)
.execute()
.await
.unwrap();
PermutationBuilder::new(data_table.clone())
.build()
.await
.unwrap();
// The builder pins its own handle; the caller's must still track latest.
let more_data = lance_datagen::gen_batch()
.col("col_a", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(50), BatchCount::from(1));
data_table.add(more_data).execute().await.unwrap();
assert_eq!(data_table.count_rows(None).await.unwrap(), 150);
}
#[tokio::test]
async fn test_permutation_builder() {
let temp_dir = tempfile::tempdir().unwrap();
@@ -8,7 +8,9 @@
//! the rows from a source table that correspond to row IDs stored in a separate table.
use crate::arrow::{SendableRecordBatchStream, SimpleRecordBatchStream};
use crate::dataloader::permutation::builder::SRC_ROW_ID_COL;
use crate::dataloader::permutation::builder::{
BASE_BRANCH_CONFIG_KEY, BASE_VERSION_CONFIG_KEY, SRC_ROW_ID_COL,
};
use crate::dataloader::permutation::split::SPLIT_ID_COLUMN;
use crate::error::Error;
use crate::query::{
@@ -23,6 +25,7 @@ use arrow_array::{RecordBatch, UInt64Array};
use arrow_schema::SchemaRef;
use datafusion_expr::{Expr, col, lit};
use futures::{StreamExt, TryStreamExt};
use lance::dataset::refs::MAIN_BRANCH;
use lance::dataset::scanner::DatasetRecordBatchStream;
use lance::io::RecordBatchStream;
use lance_arrow::RecordBatchExt;
@@ -69,6 +72,10 @@ impl PermutationReader {
permutation_table: Option<Arc<dyn BaseTable>>,
split: u64,
) -> Result<Self> {
let base_table = match &permutation_table {
Some(permutation_table) => Self::pin_base_table(base_table, permutation_table).await?,
None => base_table,
};
let mut slf = Self {
base_table,
permutation_table,
@@ -89,6 +96,34 @@ impl PermutationReader {
Ok(slf)
}
/// Pins the base table to the version the permutation was built against.
/// Permutations written before that was recorded carry no key and stay unpinned.
async fn pin_base_table(
base_table: Arc<dyn BaseTable>,
permutation_table: &Arc<dyn BaseTable>,
) -> Result<Arc<dyn BaseTable>> {
let schema = permutation_table.schema().await?;
let Some(raw) = schema.metadata.get(BASE_VERSION_CONFIG_KEY) else {
return Ok(base_table);
};
let version = raw.parse::<u64>().map_err(|e| Error::InvalidInput {
message: format!(
"Permutation table has an unreadable {} of {:?}: {}",
BASE_VERSION_CONFIG_KEY, raw, e
),
})?;
// The recorded branch, not the handle's: a worker reopens by name and lands
// on main, and version numbers are per-branch.
let branch = schema
.metadata
.get(BASE_BRANCH_CONFIG_KEY)
.map(String::as_str)
.unwrap_or(MAIN_BRANCH);
base_table
.checkout_branch_version(branch, Some(version))
.await
}
pub async fn try_from_tables(
base_table: Arc<dyn BaseTable>,
permutation_table: Arc<dyn BaseTable>,
@@ -511,9 +546,13 @@ mod tests {
use lance_datagen::{BatchCount, RowCount};
use rand::seq::SliceRandom;
// Aliased: `test_utils::datagen` exports a trait of the same name.
use crate::arrow::LanceDbDatagenExt as _;
use crate::{
Table,
arrow::SendableRecordBatchStream,
connect,
dataloader::permutation::builder::PermutationBuilder,
query::{ExecutableQuery, QueryBase},
test_utils::datagen::{LanceDbDatagenExt, virtual_table},
};
@@ -545,6 +584,58 @@ mod tests {
.await
}
/// Compaction moves row addresses, so the reader must read the pinned version.
#[tokio::test]
async fn test_reader_pins_base_version() {
let temp_dir = tempfile::tempdir().unwrap();
let db = connect(temp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let data = lance_datagen::gen_batch()
.col("idx", lance_datagen::array::step::<Int32Type>())
.into_ldb_stream(RowCount::from(20), BatchCount::from(1));
let base_table = db.create_table("base_tbl", data).execute().await.unwrap();
let permutation_table = PermutationBuilder::new(base_table.clone())
.build()
.await
.unwrap();
base_table.delete("true").await.unwrap();
base_table
.optimize(crate::table::OptimizeAction::All)
.await
.unwrap();
assert_eq!(base_table.count_rows(None).await.unwrap(), 0);
let reader = PermutationReader::try_from_tables(
base_table.base_table().clone(),
permutation_table.base_table().clone(),
0,
)
.await
.unwrap();
let values = collect_from_stream::<Int32Type>(
reader
.read(
Select::Columns(vec!["idx".to_string()]),
QueryExecutionOptions::default(),
)
.await
.unwrap(),
"idx",
)
.await;
assert_eq!(
values.len(),
20,
"reader should still see the pinned version"
);
}
#[tokio::test]
async fn test_permutation_reader() {
let base_table = lance_datagen::gen_batch()
+3 -18
View File
@@ -446,7 +446,6 @@ pub struct FunctionApplication {
function: FunctionVersionRef,
inputs: Vec<ApplicationInput>,
output: FunctionOutput,
group_id: String,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
columns: BTreeMap<String, String>,
#[serde(default, flatten, skip_serializing)]
@@ -468,10 +467,6 @@ impl FunctionApplication {
&self.output
}
pub fn group_id(&self) -> &str {
&self.group_id
}
pub fn columns(&self) -> &BTreeMap<String, String> {
&self.columns
}
@@ -513,7 +508,7 @@ pub struct InputBinding {
pub nullable: bool,
}
/// Ordered result-field to table-field mapping for a grouped binding.
/// Ordered result-field to table-field mapping for a Function binding.
///
/// Assignment state is not part of the Slice 1 client contract. During the
/// NULL transition there is no public Lance cell-flag identifier to persist.
@@ -527,20 +522,18 @@ pub struct OutputMapping {
pub nullable: bool,
}
/// Immutable grouped binding persisted by the Enterprise table service.
/// Immutable Function binding persisted by the Enterprise table service.
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct FunctionBinding {
binding_id: String,
revision: u64,
function: FunctionVersionRef,
group_id: String,
inputs: Vec<InputBinding>,
outputs: Vec<OutputMapping>,
/// Exact Arrow schema presented to the Function, encoded with the Lance
/// Namespace Arrow JSON representation.
#[serde(default, skip_serializing_if = "Option::is_none")]
input_schema: Option<Value>,
/// Exact physical Arrow schema of the grouped table outputs.
/// Exact physical Arrow schema of the binding's table outputs.
#[serde(default, skip_serializing_if = "Option::is_none")]
output_schema: Option<Value>,
}
@@ -550,18 +543,10 @@ impl FunctionBinding {
&self.binding_id
}
pub fn revision(&self) -> u64 {
self.revision
}
pub fn function(&self) -> &FunctionVersionRef {
&self.function
}
pub fn group_id(&self) -> &str {
&self.group_id
}
pub fn inputs(&self) -> &[InputBinding] {
&self.inputs
}
+55 -3
View File
@@ -37,6 +37,13 @@ pub use refresh::{RefreshMaterializedViewResult, RefreshMode};
/// Schema metadata key holding the view definition, as kind-tagged JSON.
pub const DEFINITION_META_KEY: &str = "mv.definition";
/// Schema metadata key holding the view's incarnation: a token minted at each
/// physical creation of a view table, so a view dropped and recreated under
/// the same name and definition is still told apart from the one a caller
/// captured. A view whose metadata was replaced wholesale, or one declared
/// before tokens existed, carries none until its next refresh mints one.
pub const INCARNATION_META_KEY: &str = "mv.incarnation";
/// Schema metadata key holding the source table version the view was last
/// refreshed to. Absent until the first refresh.
pub const SOURCE_VERSION_META_KEY: &str = "mv.source_version";
@@ -612,8 +619,17 @@ impl PreparedDeclaration {
pub async fn create(self, name: &str) -> Result<MaterializedView> {
let empty: Vec<std::result::Result<arrow_array::RecordBatch, arrow_schema::ArrowError>> =
vec![];
// Minted here, not at preparation: a declaration can be cloned and
// create more than one physical table, and each needs its own token.
let incarnation = uuid::Uuid::new_v4().to_string();
let mut metadata = self.schema.metadata().clone();
metadata.insert(INCARNATION_META_KEY.to_string(), incarnation.clone());
let schema = Arc::new(ArrowSchema::new_with_metadata(
self.schema.fields().clone(),
metadata,
));
let reader: Box<dyn arrow_array::RecordBatchReader + Send> =
Box::new(arrow_array::RecordBatchIterator::new(empty, self.schema));
Box::new(arrow_array::RecordBatchIterator::new(empty, schema));
let mut request = CreateTableRequest::new(name.to_string(), Box::new(reader));
let write_params = request
.write_options
@@ -648,6 +664,7 @@ impl PreparedDeclaration {
Ok(MaterializedView {
table,
definition: self.definition,
incarnation: Some(incarnation),
})
}
}
@@ -878,6 +895,7 @@ impl CreateMaterializedViewBuilder {
pub struct MaterializedView {
table: Table,
definition: MaterializedViewDefinition,
incarnation: Option<String>,
}
impl MaterializedView {
@@ -893,8 +911,13 @@ impl MaterializedView {
});
}
let schema = table.schema().await?;
let incarnation = schema.metadata().get(INCARNATION_META_KEY).cloned();
match materialized_view_kind(schema.metadata())? {
Some(MaterializedViewKind::Select(definition)) => Ok(Self { table, definition }),
Some(MaterializedViewKind::Select(definition)) => Ok(Self {
table,
definition,
incarnation,
}),
Some(MaterializedViewKind::Unrecognized { kind }) => Err(Error::NotSupported {
message: format!(
"materialized view '{}' is defined by '{kind}', which this version of \
@@ -923,6 +946,13 @@ impl MaterializedView {
&self.definition
}
/// The view's incarnation token as of when this handle was opened; see
/// [`RefreshMaterializedViewBuilder::expect_incarnation`]. `None` for a
/// view that has none yet (see [`INCARNATION_META_KEY`]).
pub fn incarnation(&self) -> Option<&str> {
self.incarnation.as_deref()
}
/// Recompute the view from its source.
///
/// By default the refresh is incremental when the source's changes can be
@@ -943,6 +973,7 @@ impl MaterializedView {
view: self.clone(),
full: false,
source_version: None,
expected_incarnation: None,
}
}
}
@@ -952,6 +983,7 @@ pub struct RefreshMaterializedViewBuilder {
view: MaterializedView,
full: bool,
source_version: Option<u64>,
expected_incarnation: Option<String>,
}
impl RefreshMaterializedViewBuilder {
@@ -967,8 +999,28 @@ impl RefreshMaterializedViewBuilder {
self
}
/// Refresh only if the view is still the incarnation that minted `token`
/// (see [`MaterializedView::incarnation`]): a refresh requested against
/// one declaration must not land in a view dropped and recreated since,
/// even under the same name and definition.
///
/// Best effort. The token is read from the latest stored manifest before
/// planning and again immediately before every commit, but it is not part
/// of the commit's own condition, so a recreation that lands between that
/// final read and the commit is not caught.
pub fn expect_incarnation(mut self, token: impl Into<String>) -> Self {
self.expected_incarnation = Some(token.into());
self
}
pub async fn execute(self) -> Result<RefreshMaterializedViewResult> {
refresh::execute_refresh(&self.view.table, self.full, self.source_version).await
refresh::execute_refresh(
&self.view.table,
self.full,
self.source_version,
self.expected_incarnation.as_deref(),
)
.await
}
}
+252 -13
View File
@@ -46,8 +46,8 @@ use lance_table::format::Fragment;
use serde::{Deserialize, Serialize};
use super::{
MaterializedViewDefinition, REFRESHED_AT_MS_META_KEY, SOURCE_ROW_ID_COLUMN,
SOURCE_VERSION_META_KEY,
INCARNATION_META_KEY, MaterializedViewDefinition, REFRESHED_AT_MS_META_KEY,
SOURCE_ROW_ID_COLUMN, SOURCE_VERSION_META_KEY,
};
use crate::database::OpenTableRequest;
use crate::table::{NativeTable, NativeTableExt, Table};
@@ -108,6 +108,7 @@ pub(crate) async fn execute_refresh(
view: &Table,
full: bool,
pinned: Option<u64>,
expected_incarnation: Option<&str>,
) -> Result<RefreshMaterializedViewResult> {
let view_native = view.as_native().ok_or_else(|| Error::NotSupported {
message: "materialized views are supported only on local tables".into(),
@@ -122,6 +123,8 @@ pub(crate) async fn execute_refresh(
view_native.dataset.reload().await?;
let view_ds = view_native.dataset.get().await?.as_ref().clone();
ensure_incarnation(&view_ds, expected_incarnation, view.name()).await?;
// The definition a handle cached at open may since have been replaced;
// what refresh executes and what it stamps must be one generation.
let definition = match super::materialized_view_kind(&view_ds.schema().metadata)? {
@@ -240,6 +243,7 @@ pub(crate) async fn execute_refresh(
increment,
definition,
watermark,
expected_incarnation,
)
.await?;
match reconciled {
@@ -253,6 +257,7 @@ pub(crate) async fn execute_refresh(
source_version,
source_ts,
definition,
expected_incarnation,
)
.await
}
@@ -266,6 +271,7 @@ pub(crate) async fn execute_refresh(
source_version,
source_ts,
definition,
expected_incarnation,
)
.await
}
@@ -595,6 +601,7 @@ async fn incremental(
increment: Increment,
definition: &MaterializedViewDefinition,
watermark: Option<u64>,
expected_incarnation: Option<&str>,
) -> Result<Option<RefreshMaterializedViewResult>> {
let new_fragments = increment.appended;
let watermark_version = watermark.unwrap_or(0);
@@ -671,15 +678,35 @@ async fn incremental(
};
let nothing_to_add = (new_fragments.is_empty() && !updated_rows) || remaining == Some(0);
if nothing_to_add && eviction.is_none() {
result.version =
stamp_watermark(view_native, view_ds.clone(), source_version, source_ts).await?;
result.version = stamp_watermark(
view_native,
view_ds.clone(),
source_version,
source_ts,
expected_incarnation,
)
.await?;
return Ok(Some(result));
}
// Rows left but none arrive: the removals still have to be published.
if nothing_to_add {
let filter = refresh_filter(&empty_keys(view_ds)?)?;
let published = publish(view_ds, eviction, Vec::new(), Some(filter)).await?;
result.version = stamp_watermark(view_native, published, source_version, source_ts).await?;
let published = publish(
view_ds,
eviction,
Vec::new(),
Some(filter),
expected_incarnation,
)
.await?;
result.version = stamp_watermark(
view_native,
published,
source_version,
source_ts,
expected_incarnation,
)
.await?;
return Ok(Some(result));
}
@@ -737,12 +764,20 @@ async fn incremental(
eviction,
Vec::new(),
Some(refresh_filter(&empty_keys(view_ds)?)?),
expected_incarnation,
)
.await?
} else {
view_ds.clone()
};
result.version = stamp_watermark(view_native, published, source_version, source_ts).await?;
result.version = stamp_watermark(
view_native,
published,
source_version,
source_ts,
expected_incarnation,
)
.await?;
return Ok(Some(result));
};
let stream: SendableRecordBatchStream = Box::pin(RecordBatchStreamAdapter::new(
@@ -775,9 +810,23 @@ async fn incremental(
});
};
let filter = refresh_filter(&keys)?;
let appended = publish(view_ds, eviction, new_fragments, Some(filter)).await?;
let appended = publish(
view_ds,
eviction,
new_fragments,
Some(filter),
expected_incarnation,
)
.await?;
result.rows_written = rows_written.load(Ordering::Relaxed);
result.version = stamp_watermark(view_native, appended, source_version, source_ts).await?;
result.version = stamp_watermark(
view_native,
appended,
source_version,
source_ts,
expected_incarnation,
)
.await?;
Ok(Some(result))
}
@@ -788,6 +837,7 @@ async fn rebuild(
source_version: u64,
source_ts: u128,
definition: &MaterializedViewDefinition,
expected_incarnation: Option<&str>,
) -> Result<RefreshMaterializedViewResult> {
let rows_written = Arc::new(AtomicU64::new(0));
let schema = Arc::new(ArrowSchema::from(view_ds.schema()));
@@ -810,8 +860,16 @@ async fn rebuild(
// carries no schema metadata, so it cannot erase a definition update
// that raced in the way an overwrite (which adopts its stream's schema)
// durably would -- and it must land on the planned generation or abort.
let replaced = replace_retaining_indices(view_ds.clone(), stream, keys).await?;
let version = stamp_watermark(view_native, replaced, source_version, source_ts).await?;
let replaced =
replace_retaining_indices(view_ds.clone(), stream, keys, expected_incarnation).await?;
let version = stamp_watermark(
view_native,
replaced,
source_version,
source_ts,
expected_incarnation,
)
.await?;
Ok(RefreshMaterializedViewResult {
mode: RefreshMode::Rebuild,
rows_written: rows_written.load(Ordering::Relaxed),
@@ -828,11 +886,13 @@ async fn replace_retaining_indices(
view_ds: Dataset,
stream: SendableRecordBatchStream,
keys: Arc<StdMutex<KeyExistenceFilterBuilder>>,
expected_incarnation: Option<&str>,
) -> Result<Dataset> {
let ds = Arc::new(view_ds);
let read_version = ds.version().version;
#[cfg(test)]
tests::hold_before_publish(ds.uri()).await;
ensure_incarnation(&ds, expected_incarnation, ds.uri()).await?;
let removed_fragment_ids: Vec<u64> = ds.get_fragments().iter().map(|f| f.id() as u64).collect();
let write_txn = InsertBuilder::new(WriteDestination::Dataset(ds.clone()))
@@ -886,6 +946,32 @@ async fn replace_retaining_indices(
}
/// Record that the view now reflects `source_version`, including the view
/// Refuse to act on a view that is not `expected`'s incarnation, judged from
/// the latest stored manifest. Not a commit condition; see
/// `RefreshMaterializedViewBuilder::expect_incarnation`.
async fn ensure_incarnation(view_ds: &Dataset, expected: Option<&str>, what: &str) -> Result<()> {
let Some(expected) = expected else {
return Ok(());
};
let mut latest = view_ds.clone();
latest.checkout_latest().await?;
match latest.schema().metadata.get(INCARNATION_META_KEY) {
Some(actual) if actual == expected => Ok(()),
Some(_) => Err(Error::Runtime {
message: format!(
"materialized view '{what}' is not the incarnation this refresh was \
requested for: it was dropped and recreated"
),
}),
None => Err(Error::Runtime {
message: format!(
"materialized view '{what}' carries no incarnation token: its schema \
metadata was replaced since the token was captured"
),
}),
}
}
/// version this very commit produces. The version is predicted and then
/// verified; on a mismatch another commit raced in between, and the stamp
/// ABORTS rather than certify that commit as the refresh's own generation.
@@ -895,10 +981,21 @@ async fn stamp_watermark(
mut dataset: Dataset,
source_version: u64,
source_ts: u128,
expected_incarnation: Option<&str>,
) -> Result<u64> {
ensure_incarnation(&dataset, expected_incarnation, dataset.uri()).await?;
let predicted = dataset.version().version + 1;
// A view with no token (declared before tokens existed, or its metadata
// replaced wholesale) starts a new incarnation here.
let incarnation = dataset
.schema()
.metadata
.get(INCARNATION_META_KEY)
.cloned()
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
dataset
.update_schema_metadata([
(INCARNATION_META_KEY.to_string(), Some(incarnation)),
(
SOURCE_VERSION_META_KEY.to_string(),
Some(source_version.to_string()),
@@ -1052,12 +1149,14 @@ async fn publish(
eviction: Option<(Vec<Fragment>, Vec<u64>)>,
new_fragments: Vec<Fragment>,
keys: Option<KeyExistenceFilter>,
expected_incarnation: Option<&str>,
) -> Result<Dataset> {
let planned = view_ds.version().version;
#[cfg(test)]
tests::hold_before_publish(view_ds.uri()).await;
#[cfg(test)]
tests::hold_until_peers_planned();
ensure_incarnation(view_ds, expected_incarnation, view_ds.uri()).await?;
let (updated_fragments, removed_fragment_ids) = eviction.unwrap_or_default();
let committed = CommitBuilder::new(WriteDestination::Dataset(Arc::new(view_ds.clone())))
.execute(Transaction::new(
@@ -1830,7 +1929,9 @@ mod tests {
);
let staged = eviction.finish().await.unwrap();
assert!(staged.is_some(), "four ids over a chunk of two stage twice");
publish(&view_ds, staged, Vec::new(), None).await.unwrap();
publish(&view_ds, staged, Vec::new(), None, None)
.await
.unwrap();
native.dataset.reload().await.unwrap();
assert_eq!(read(view.table(), "x").await, vec![5, 6]);
@@ -2486,6 +2587,144 @@ mod tests {
assert_eq!(read(view.table(), "twice").await, vec![14]);
}
/// A refresh bound to an incarnation refuses a view dropped and recreated
/// since, even under the same name and definition; the recreated view's
/// own token is accepted, and the token survives a refresh's stamp.
#[tokio::test]
async fn test_refresh_refuses_a_recreated_view_incarnation() {
let (conn, _, view) = refreshed_doubled(vec![1]).await;
let token = view.incarnation().unwrap().to_string();
view.refresh()
.expect_incarnation(&token)
.execute()
.await
.unwrap();
let reopened = conn.open_materialized_view("doubled").await.unwrap();
assert_eq!(reopened.incarnation(), Some(token.as_str()));
conn.drop_table("doubled", &[]).await.unwrap();
let recreated = doubled_view(&conn).await;
assert_ne!(recreated.incarnation(), Some(token.as_str()));
let err = recreated
.refresh()
.expect_incarnation(&token)
.execute()
.await
.unwrap_err();
assert!(err.to_string().contains("dropped and recreated"), "{err}");
assert_eq!(read(recreated.table(), "twice").await, Vec::<i32>::new());
recreated
.refresh()
.expect_incarnation(recreated.incarnation().unwrap())
.execute()
.await
.unwrap();
assert_eq!(read(recreated.table(), "twice").await, vec![2]);
}
/// A cloned declaration creates two physical tables; each gets its own
/// token.
#[tokio::test]
async fn test_cloned_declaration_mints_a_fresh_incarnation_per_create() {
let (conn, source) = db_with_source(vec![1]).await;
let prepared = crate::materialized_view::prepare_declaration(
&source,
&[("x".into(), "x".into()), ("twice".into(), "x * 2".into())],
None,
None,
)
.await
.unwrap();
let replacement = prepared.clone();
let first = prepared.create("cloned").await.unwrap();
let first_token = first.incarnation().unwrap().to_string();
conn.drop_table("cloned", &[]).await.unwrap();
let second = replacement.create("cloned").await.unwrap();
assert_ne!(second.incarnation(), Some(first_token.as_str()));
}
/// A recreation that lands after planning but before publication is
/// caught by the pre-commit read: the stale refresh fails and the
/// replacement stays empty under its own token.
#[tokio::test(flavor = "multi_thread")]
async fn test_bound_refresh_cannot_publish_into_a_raced_recreation() {
let _serial = DRIFT_LOCK.lock().await;
let (conn, _) = db_with_source(vec![1]).await;
let view = doubled_view(&conn).await;
let token = view.incarnation().unwrap().to_string();
let uri = view
.table()
.as_native()
.unwrap()
.dataset
.get()
.await
.unwrap()
.uri()
.to_string();
*DRIFT_TARGET.lock().unwrap() = Some(uri);
let refreshing =
tokio::spawn(async move { view.refresh().expect_incarnation(token).execute().await });
tokio::time::timeout(std::time::Duration::from_secs(30), DRIFT_PLANNED.notified())
.await
.expect("refresh never reached publication");
conn.drop_table("doubled", &[]).await.unwrap();
let replacement = doubled_view(&conn).await;
let replacement_token = replacement.incarnation().unwrap().to_string();
DRIFT_RELEASED.notify_one();
let result = refreshing.await.unwrap();
assert!(result.is_err(), "the stale refresh unexpectedly succeeded");
let reopened = conn.open_materialized_view("doubled").await.unwrap();
assert_eq!(reopened.incarnation(), Some(replacement_token.as_str()));
assert_eq!(read(reopened.table(), "twice").await, Vec::<i32>::new());
}
/// Replacing the schema metadata wholesale drops the token. A refresh
/// bound to the old token is refused for that reason, not as a
/// recreation; an unbound refresh mints the view a fresh one.
#[tokio::test]
async fn test_a_view_whose_metadata_was_replaced_starts_a_new_incarnation() {
let (conn, _, view) = refreshed_doubled(vec![1]).await;
let token = view.incarnation().unwrap().to_string();
let mut metadata = HashMap::new();
metadata.insert(
crate::materialized_view::DEFINITION_META_KEY.to_string(),
crate::materialized_view::definition_to_metadata(view.definition()).unwrap(),
);
view.table()
.as_native()
.unwrap()
.replace_schema_metadata(metadata)
.await
.unwrap();
assert_eq!(
conn.open_materialized_view("doubled")
.await
.unwrap()
.incarnation(),
None
);
let err = view
.refresh()
.expect_incarnation(&token)
.execute()
.await
.unwrap_err();
assert!(err.to_string().contains("no incarnation token"), "{err}");
view.refresh().execute().await.unwrap();
let reopened = conn.open_materialized_view("doubled").await.unwrap();
assert!(reopened.incarnation().is_some());
assert_ne!(reopened.incarnation(), Some(token.as_str()));
}
/// In-process refreshes of one view serialize: the loser of the race
/// observes the winner's watermark instead of appending the same rows.
#[tokio::test(flavor = "multi_thread")]
@@ -2528,7 +2767,7 @@ mod tests {
let stale = view_native.dataset.get().await.unwrap().as_ref().clone();
view.table().delete("x = 1").await.unwrap();
let err = stamp_watermark(view_native, stale, 99, 99).await;
let err = stamp_watermark(view_native, stale, 99, 99, None).await;
assert!(err.is_err());
let result = view.refresh().execute().await.unwrap();
+6 -3
View File
@@ -1774,11 +1774,14 @@ mod tests {
.postfilter();
let result = query.execute().await;
let mut stream = result.expect("should have result");
// should only have one batch
let mut num_rows = 0;
while let Some(batch) = stream.next().await {
// post filter should have removed some rows
assert!(batch.expect("should be Ok").num_rows() < 10);
let batch = batch.expect("should be Ok");
let ids: &Int32Array = batch["id"].as_primitive();
assert!(ids.iter().all(|id| id.unwrap() % 2 == 0));
num_rows += batch.num_rows();
}
assert!(num_rows <= 10);
let query = table
.query()
+171 -20
View File
@@ -344,6 +344,62 @@ impl<S: HttpSend> RemoteDatabase<S> {
self.table_cache.remove(&cache_key).await;
Ok((request_id, resp))
}
/// Collect the tables of a namespace in name order, for `table_names`.
///
/// `table_names` promises name order and resumes after a table name, but the namespace
/// route's `page_token` is opaque -- it belongs to the store the listing walks, and a
/// token this client invented would resume from the wrong place. So the whole namespace is
/// walked by handing each response's token straight back, and the name semantics are
/// applied here. Constructing no token is what makes this work against a server on either
/// side of the change: it only ever repeats what the server said.
///
/// This is the cost `table_names` already paid -- the server used to enumerate and sort the
/// namespace on every request -- and it is why `list_tables` replaces it.
async fn table_names_in_namespace(
&self,
request: &TableNamesRequest,
) -> Result<(Vec<String>, ServerVersion)> {
let namespace_id =
build_namespace_identifier(&request.namespace_path, &self.client.id_delimiter);
let path = format!("/v1/namespace/{}/table/list", namespace_id);
let mut names = Vec::new();
// Every page reports the same server, so keep the first page's version.
let mut version: Option<ServerVersion> = None;
let mut page_token: Option<String> = None;
loop {
let mut req = self.client.get(&path);
if let Some(ref token) = page_token {
req = req.query(&[("page_token", token)]);
}
let (request_id, rsp) = self.client.send_with_retry(req, None, true).await?;
let rsp = self.client.check_response(&request_id, rsp).await?;
if version.is_none() {
version = Some(parse_server_version(&request_id, &rsp)?);
}
let response: ListTablesResponse = rsp.json().await.err_to_http(request_id)?;
names.extend(response.tables);
// An empty token is the end of the listing, not a token to send back: a server
// that reads an empty token as "start from the beginning" would hand back the
// first page again.
match response.page_token.filter(|token| !token.is_empty()) {
// A server that repeated a token would never finish; treat that as the end
// rather than looping on it.
Some(token) if Some(&token) != page_token.as_ref() => page_token = Some(token),
_ => break,
}
}
names.sort();
if let Some(ref start_after) = request.start_after {
names.retain(|name| name > start_after);
}
if let Some(limit) = request.limit {
names.truncate(limit as usize);
}
Ok((names, version.unwrap_or_default()))
}
}
#[cfg(all(test, feature = "remote"))]
@@ -621,29 +677,29 @@ impl<S: HttpSend> Database for RemoteDatabase<S> {
}
async fn table_names(&self, request: TableNamesRequest) -> Result<Vec<String>> {
let mut req = if !request.namespace_path.is_empty() {
let namespace_id =
build_namespace_identifier(&request.namespace_path, &self.client.id_delimiter);
self.client
.get(&format!("/v1/namespace/{}/table/list", namespace_id))
let (tables, version) = if request.namespace_path.is_empty() {
// The flat route resumes after a table name and orders by name, which is exactly
// what `start_after` means, so the server does the paging.
let mut req = self.client.get("/v1/table/");
if let Some(limit) = request.limit {
req = req.query(&[("limit", limit)]);
}
if let Some(ref start_after) = request.start_after {
req = req.query(&[("page_token", start_after)]);
}
let (request_id, rsp) = self.client.send_with_retry(req, None, true).await?;
let rsp = self.client.check_response(&request_id, rsp).await?;
let version = parse_server_version(&request_id, &rsp)?;
let tables = rsp
.json::<ListTablesResponse>()
.await
.err_to_http(request_id)?
.tables;
(tables, version)
} else {
self.client.get("/v1/table/")
self.table_names_in_namespace(&request).await?
};
if let Some(limit) = request.limit {
req = req.query(&[("limit", limit)]);
}
if let Some(start_after) = request.start_after {
req = req.query(&[("page_token", start_after)]);
}
let (request_id, rsp) = self.client.send_with_retry(req, None, true).await?;
let rsp = self.client.check_response(&request_id, rsp).await?;
let version = parse_server_version(&request_id, &rsp)?;
let tables = rsp
.json::<ListTablesResponse>()
.await
.err_to_http(request_id)?
.tables;
for table in &tables {
let table_identifier =
build_table_identifier(table, &request.namespace_path, &self.client.id_delimiter);
@@ -1227,6 +1283,101 @@ mod tests {
assert_eq!(names, vec!["table1", "table2"]);
}
#[tokio::test]
async fn test_table_names_in_a_namespace_never_invents_a_page_token() {
// The namespace route's token belongs to the store, so `table_names` cannot build one
// from `start_after`. It walks the namespace on the server's own tokens and applies the
// name semantics itself, which is what keeps it working either side of the change.
let page = Arc::new(AtomicUsize::new(0));
let conn = Connection::new_with_handler(move |request| {
assert_eq!(request.url().path(), "/v1/namespace/ns/table/list");
let query = request.url().query().unwrap_or("");
assert!(
!query.contains("page_token=users"),
"a table name must never be sent as a page token: {query}"
);
match page.fetch_add(1, Ordering::SeqCst) {
0 => {
assert!(
!query.contains("page_token"),
"the walk starts with no token"
);
http::Response::builder()
.status(200)
.body(r#"{"tables": ["users", "orders"], "page_token": "opaque-1"}"#)
.unwrap()
}
_ => {
assert!(query.contains("page_token=opaque-1"));
http::Response::builder()
.status(200)
.body(r#"{"tables": ["widgets"]}"#)
.unwrap()
}
}
});
let names = conn
.table_names()
.namespace(vec!["ns".to_string()])
.start_after("users")
.execute()
.await
.unwrap();
// Name order, resumed after "users": "orders" sorts before it and is dropped.
assert_eq!(names, vec!["widgets"]);
}
#[tokio::test]
async fn test_table_names_in_a_namespace_stops_on_a_repeated_token() {
// A server that handed back the token it was given would never finish the walk.
let conn = Connection::new_with_handler(|_request| {
http::Response::builder()
.status(200)
.body(r#"{"tables": ["a"], "page_token": "same"}"#)
.unwrap()
});
let names = conn
.table_names()
.namespace(vec!["ns".to_string()])
.execute()
.await
.unwrap();
// The guard bounds the walk instead of letting it run forever. The repeat is the
// server breaking the token contract and is not papered over here.
assert_eq!(names, vec!["a", "a"]);
}
#[tokio::test]
async fn test_table_names_in_a_namespace_stops_on_an_empty_token() {
// An empty token ends the listing. Sending it back would ask a server that reads it
// as "start from the beginning" for the first page a second time, and every name on
// that page would be collected twice.
let requests = Arc::new(AtomicUsize::new(0));
let seen = requests.clone();
let conn = Connection::new_with_handler(move |request| {
seen.fetch_add(1, Ordering::SeqCst);
assert!(
!request.url().query().unwrap_or("").contains("page_token"),
"an empty token must never be sent back"
);
http::Response::builder()
.status(200)
.body(r#"{"tables": ["a"], "page_token": ""}"#)
.unwrap()
});
let names = conn
.table_names()
.namespace(vec!["ns".to_string()])
.execute()
.await
.unwrap();
assert_eq!(names, vec!["a"]);
assert_eq!(requests.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn test_table_names_pagination() {
let conn = Connection::new_with_handler(|request| {
+98 -5
View File
@@ -1775,6 +1775,18 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
Ok(())
}
async fn snapshot_at_current_version(&self) -> Result<Option<Arc<dyn BaseTable>>> {
// A checked-out handle already names its snapshot. Otherwise resolve
// latest exactly once before creating the independent pinned handle.
let version = match self.current_version().await {
Some(version) => version,
None => self.describe().await?.version,
};
let snapshot = self.with_branch(self.branch.clone());
*snapshot.version.write().await = Some(version);
Ok(Some(Arc::new(snapshot)))
}
async fn restore(&self) -> Result<()> {
let mut request = self
.client
@@ -4337,6 +4349,43 @@ mod tests {
assert!(!table.base_table().scan_order_is_deterministic());
}
#[tokio::test]
async fn test_checkout_branch_pins_without_touching_the_original() {
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let recorder = seen.clone();
let table = Table::new_with_handler_version(
"my_table",
semver::Version::new(0, 5, 0),
move |request| match request.url().path() {
"/v1/table/my_table/describe/" => http::Response::builder()
.status(200)
.body(br#"{"version": 42, "schema": {"fields": []}}"#.to_vec())
.unwrap(),
"/v1/table/my_table/count_rows/" => {
let body = request_body_json(&request);
recorder.lock().unwrap().push(body["version"].clone());
http::Response::builder()
.status(200)
.body(b"0".to_vec())
.unwrap()
}
path => panic!("unexpected request path: {path}"),
},
);
let pinned = table.checkout_branch("main", Some(42)).await.unwrap();
pinned.count_rows(None).await.unwrap();
table.count_rows(None).await.unwrap();
let seen = seen.lock().unwrap();
assert_eq!(seen[0], 42, "the pinned handle must send its version");
assert!(
seen[1].is_null(),
"the original handle must still track latest, got {:?}",
seen[1]
);
}
#[tokio::test]
async fn test_fetch_blobs_sends_the_checked_out_version() {
let ipc = one_row_blob_ipc_stream("image");
@@ -6738,8 +6787,7 @@ mod tests {
r#"{
"function":{"name":"embed","version":"fv_01K3EXACT"},
"inputs":[{"parameter":"text","kind":"column","value":{"path":"description"}}],
"output":{"kind":"scalar","arrow_type":"list<float32>","nullable":false},
"group_id":"fg_scalar"
"output":{"kind":"scalar","arrow_type":"list<float32>","nullable":false}
}"#,
)
.unwrap();
@@ -6754,7 +6802,53 @@ mod tests {
}
#[tokio::test]
async fn test_add_named_struct_function_expands_one_atomic_sibling_group() {
async fn test_add_fixed_size_list_function_column_declares_the_vector_type() {
let table = Table::new_with_handler("my_table", |request| {
match request.url().path() {
"/v1/table/my_table/describe/" => http::Response::builder()
.status(200)
.body(
r#"{"version":1,"schema":{"fields":[{"name":"description","nullable":true,"type":{"type":"string"}}]}}"#,
)
.unwrap(),
"/v1/table/my_table/add_columns/" => {
let actual: serde_json::Value = serde_json::from_slice(
request.body().unwrap().as_bytes().unwrap(),
)
.unwrap();
let expected: serde_json::Value = serde_json::from_str(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json"
))
.unwrap();
assert_eq!(actual, expected);
http::Response::builder()
.status(200)
.body(r#"{"version":8}"#)
.unwrap()
}
path => panic!("Unexpected path: {path}"),
}
});
let application = crate::function::FunctionApplication::from_json(
r#"{
"function":{"name":"embed","version":"fv_01K3EXACT"},
"inputs":[{"parameter":"text","kind":"column","value":{"path":"description"}}],
"output":{"kind":"scalar","arrow_type":"fixed_size_list<float32, 3>","nullable":false}
}"#,
)
.unwrap();
let result = table
.add_columns()
.function_as("embedding", application)
.execute()
.await
.unwrap();
assert_eq!(result.version, 8);
}
#[tokio::test]
async fn test_add_named_struct_function_expands_one_atomic_binding() {
let table = Table::new_with_handler("my_table", |request| match request.url().path() {
"/v1/table/my_table/describe/" => http::Response::builder()
.status(200)
@@ -6769,7 +6863,7 @@ mod tests {
let actual: serde_json::Value =
serde_json::from_slice(request.body().unwrap().as_bytes().unwrap()).unwrap();
let expected: serde_json::Value = serde_json::from_str(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json"
"../../tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json"
))
.unwrap();
assert_eq!(actual, expected);
@@ -6791,7 +6885,6 @@ mod tests {
{"name":"normalized_text","arrow_type":"utf8","nullable":false},
{"name":"token_count","arrow_type":"int64","nullable":false}
]},
"group_id":"fg_01K3TEXT",
"columns":{"normalized_text":"search_text"}
}"#,
)
+9 -1
View File
@@ -751,7 +751,7 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync {
message: "computed columns are not supported on this table type".into(),
})
}
/// Declare one immutable registered-Function output group.
/// Declare one immutable registered-Function binding.
async fn add_function_columns(
&self,
_application: &crate::function::FunctionApplication,
@@ -792,6 +792,14 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync {
async fn checkout_tag(&self, tag: &str) -> Result<()>;
/// Checkout the latest version of the table.
async fn checkout_latest(&self) -> Result<()>;
/// Return an independent handle pinned to the version currently selected.
///
/// Backends that can advance between requests should override this for
/// multi-request operations that need snapshot consistency. Backends whose
/// existing handles already provide the desired behavior return `None`.
async fn snapshot_at_current_version(&self) -> Result<Option<Arc<dyn BaseTable>>> {
Ok(None)
}
/// Whether repeated identical scans return rows in the same order.
///
/// Callers that assign meaning to a row's position must order the results
+1 -1
View File
@@ -88,7 +88,7 @@ impl AddColumnsBuilder {
}
/// Declare every field of a named-struct Function result as one atomic
/// sibling group. Result-field aliases come from
/// binding. Result-field aliases come from
/// [`FunctionApplication::columns`](crate::function::FunctionApplication::columns).
///
/// ```
+70 -30
View File
@@ -13,7 +13,7 @@
//! self-describing -- both are derived from the expression, so a caller writes
//! neither -- while a kind resolved through a registry cannot be typed without
//! consulting it. Registered Functions use an exact remote version plus a
//! schema-level grouped binding; unknown newer kinds remain readable and fail
//! schema-level Function binding; unknown newer kinds remain readable and fail
//! closed before mutation.
//!
//! [`computed_columns`] and [`computed_column_from_field`] read declarations
@@ -46,16 +46,16 @@ pub const EXPRESSION_META_KEY: &str = "computed_column.expression";
/// Field metadata key holding the column's inputs, as a JSON array of names.
pub const INPUTS_META_KEY: &str = "computed_column.inputs";
/// Field metadata key holding the grouped Function binding identity.
/// Field metadata key holding the Function binding identity.
pub const FUNCTION_BINDING_ID_META_KEY: &str = "computed_column.function.binding_id";
/// Field metadata key holding this sibling's ordered Function output ordinal.
pub const FUNCTION_OUTPUT_ORDINAL_META_KEY: &str = "computed_column.function.output_ordinal";
/// Schema metadata key holding all immutable grouped Function bindings.
/// Schema metadata key holding all immutable Function bindings.
pub const FUNCTION_BINDINGS_META_KEY: &str = "lancedb::function_bindings";
/// Version of the schema-level grouped Function binding envelope.
/// Version of the schema-level Function binding envelope.
pub const FUNCTION_BINDINGS_VERSION: u32 = 1;
/// Value of [`KIND_META_KEY`] for a column defined by a SQL expression.
@@ -81,7 +81,7 @@ pub enum ComputedColumnKind {
/// The expression.
expression: String,
},
/// One physical output in an immutable grouped registered-Function
/// One physical output in an immutable registered-Function
/// binding. The full binding lives in schema metadata.
Function {
/// Shared immutable binding identity.
@@ -159,7 +159,7 @@ struct FunctionBindingEnvelope {
bindings: Vec<Value>,
}
/// Encode immutable grouped bindings for schema-level persistence.
/// Encode immutable Function bindings for schema-level persistence.
pub fn function_bindings_metadata(bindings: &[FunctionBinding]) -> Result<String> {
let bindings = bindings
.iter()
@@ -177,7 +177,7 @@ pub fn function_bindings_metadata(bindings: &[FunctionBinding]) -> Result<String
})
}
/// Decode known grouped Function bindings without rewriting their raw schema
/// Decode known Function bindings without rewriting their raw schema
/// metadata. Unknown envelope versions fail closed.
pub fn function_bindings(schema: &ArrowSchema) -> Result<Vec<FunctionBinding>> {
let Some(envelope) = function_binding_envelope(schema)? else {
@@ -238,21 +238,15 @@ pub(crate) fn ensure_supported_function_metadata(schema: &ArrowSchema) -> Result
message: format!("duplicate Function binding '{}'", binding.binding_id()),
});
}
if binding.revision() == 0 || binding.outputs().is_empty() {
if binding.outputs().is_empty() {
return Err(Error::InvalidInput {
message: format!(
"Function binding '{}' has no immutable revision or outputs",
binding.binding_id()
),
message: format!("Function binding '{}' has no outputs", binding.binding_id()),
});
}
if binding.function().name.is_empty()
|| binding.function().version.is_empty()
|| binding.group_id().is_empty()
{
if binding.function().name.is_empty() || binding.function().version.is_empty() {
return Err(Error::InvalidInput {
message: format!(
"Function binding '{}' has no exact version or group identity",
"Function binding '{}' has no exact version",
binding.binding_id()
),
});
@@ -493,9 +487,7 @@ fn ensure_known_binding_shape(value: &Value) -> Result<()> {
value,
&[
"binding_id",
"revision",
"function",
"group_id",
"inputs",
"outputs",
"input_schema",
@@ -586,6 +578,25 @@ fn canonical_input_arrow_type(field: &JsonArrowField) -> Result<String> {
}
}
/// `fixed_size_list<item, size>` -> (`item`, `size`); the comma must sit outside
/// any nested `<...>`.
fn split_fixed_size_list(raw: &str) -> Option<(&str, i32)> {
let inner = raw.strip_prefix("fixed_size_list<")?.strip_suffix('>')?;
let mut depth = 0_u32;
let mut separator = None;
for (index, byte) in inner.bytes().enumerate() {
match byte {
b'<' => depth += 1,
b'>' => depth = depth.checked_sub(1)?,
b',' if depth == 0 => separator = Some(index),
_ => {}
}
}
let (item, size) = inner.split_at(separator?);
let size: i32 = size[1..].trim().parse().ok()?;
(size > 0).then_some((item.trim(), size))
}
fn parse_output_arrow_type(raw: &str) -> Result<JsonArrowDataType> {
fn parse(raw: &str) -> Result<JsonArrowDataType> {
let raw = raw.trim();
@@ -618,6 +629,16 @@ fn parse_output_arrow_type(raw: &str) -> Result<JsonArrowDataType> {
)]);
return Ok(data_type);
}
if let Some((inner, size)) = split_fixed_size_list(raw) {
let mut data_type = JsonArrowDataType::new("fixed_size_list".to_string());
data_type.fields = Some(vec![JsonArrowField::new(
"item".to_string(),
false,
parse(inner)?,
)]);
data_type.length = Some(i64::from(size));
return Ok(data_type);
}
let normalized = match raw {
"boolean" => "bool",
"string" => "utf8",
@@ -757,12 +778,9 @@ pub(crate) fn plan_function_application(
message: "Function application contains fields from a newer contract".into(),
});
}
if application.function().name.is_empty()
|| application.function().version.is_empty()
|| application.group_id().is_empty()
{
if application.function().name.is_empty() || application.function().version.is_empty() {
return Err(invalid_function(
"Function application requires an exact version and group identity",
"Function application requires an exact version",
));
}
@@ -1322,6 +1340,32 @@ pub(super) async fn add_foreign_kind(table: &crate::Table, name: &str, kind: &st
#[cfg(test)]
mod tests {
#[test]
fn output_arrow_type_grammar_matches_the_shared_golden() {
let golden: serde_json::Value = serde_json::from_str(include_str!(
"../../tests/fixtures/first_class_functions/v1/arrow_types.json"
))
.unwrap();
let valid = golden["valid"].as_array().unwrap().iter();
for case in valid.chain(golden["server_only"].as_array().unwrap()) {
let raw = case["arrow_type"].as_str().unwrap();
let parsed = super::parse_output_arrow_type(raw)
.unwrap_or_else(|error| panic!("{raw}: {error}"));
assert_eq!(
serde_json::to_value(&parsed).unwrap(),
case["json"],
"{raw}"
);
}
for raw in golden["invalid"].as_array().unwrap() {
let raw = raw.as_str().unwrap();
assert!(
super::parse_output_arrow_type(raw).is_err(),
"{raw:?} should be rejected"
);
}
}
use arrow_array::record_batch;
use arrow_schema::DataType;
use futures::TryStreamExt;
@@ -2206,7 +2250,6 @@ mod tests {
{{"name":"normalized_text","arrow_type":"utf8","nullable":false}},
{{"name":"token_count","arrow_type":"int64","nullable":false}}
]}},
"group_id":"fg_exact",
"columns":{columns}
}}"#
))
@@ -2387,8 +2430,7 @@ mod tests {
r#"{
"function":{"name":"f","version":"fv"},
"inputs":[{"parameter":"title","kind":"future_source","value":{"path":"title"}}],
"output":{"kind":"scalar","arrow_type":"int64","nullable":false},
"group_id":"fg"
"output":{"kind":"scalar","arrow_type":"int64","nullable":false}
}"#,
)
.unwrap();
@@ -2401,7 +2443,6 @@ mod tests {
"function":{"name":"f","version":"fv"},
"inputs":[],
"output":{"kind":"scalar","arrow_type":"int64","nullable":false},
"group_id":"fg",
"future_declaration":{"mode":"managed"}
}"#,
)
@@ -2415,8 +2456,7 @@ mod tests {
r#"{
"function":{"name":"f","version":"fv"},
"inputs":[],
"output":{"kind":"scalar","arrow_type":"int64","nullable":false,"assignment":"cell_flag"},
"group_id":"fg"
"output":{"kind":"scalar","arrow_type":"int64","nullable":false,"assignment":"cell_flag"}
}"#,
)
.unwrap();
+3
View File
@@ -70,6 +70,9 @@ async fn can_execute_namespace_query(table: &NativeTable, query: &AnyQuery) -> R
.contains(&NamespaceClientPushdownOperation::QueryTable)
&& table.namespace_client.is_some()
&& table.dataset.current_branch().is_none()
// NsQueryTableRequest has no version field, so a pushed-down query would
// read latest and ignore the pin.
&& table.dataset.time_travel_version().is_none()
&& !requires_local_namespace_execution(query))
{
return Ok(false);
@@ -84,7 +84,6 @@ fn application_and_binding_match_shared_remote_goldens() {
let binding = FunctionBinding::from_json(&fixture("remote_function_binding.json"))
.expect("binding fixture");
assert_eq!(binding.revision(), 3);
assert_eq!(binding.function().version, "fv_01K3TEXT");
assert_eq!(binding.outputs()[0].output_ordinal, 0);
assert_eq!(binding.outputs()[1].output_ordinal, 1);
@@ -0,0 +1,333 @@
{
"valid": [
{
"arrow_type": "bool",
"json": {
"type": "bool"
}
},
{
"arrow_type": "int8",
"json": {
"type": "int8"
}
},
{
"arrow_type": "int16",
"json": {
"type": "int16"
}
},
{
"arrow_type": "int32",
"json": {
"type": "int32"
}
},
{
"arrow_type": "int64",
"json": {
"type": "int64"
}
},
{
"arrow_type": "uint8",
"json": {
"type": "uint8"
}
},
{
"arrow_type": "uint16",
"json": {
"type": "uint16"
}
},
{
"arrow_type": "uint32",
"json": {
"type": "uint32"
}
},
{
"arrow_type": "uint64",
"json": {
"type": "uint64"
}
},
{
"arrow_type": "float16",
"json": {
"type": "float16"
}
},
{
"arrow_type": "float32",
"json": {
"type": "float32"
}
},
{
"arrow_type": "float64",
"json": {
"type": "float64"
}
},
{
"arrow_type": "utf8",
"json": {
"type": "utf8"
}
},
{
"arrow_type": "binary",
"json": {
"type": "binary"
}
},
{
"arrow_type": "date32",
"json": {
"type": "date32"
}
},
{
"arrow_type": "date64",
"json": {
"type": "date64"
}
},
{
"arrow_type": "list<float32>",
"json": {
"type": "list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "float32"
}
}
]
}
},
{
"arrow_type": "large_list<float32>",
"json": {
"type": "large_list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "float32"
}
}
]
}
},
{
"arrow_type": "list<int64>",
"json": {
"type": "list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "int64"
}
}
]
}
},
{
"arrow_type": "large_list<int64>",
"json": {
"type": "large_list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "int64"
}
}
]
}
},
{
"arrow_type": "list<utf8>",
"json": {
"type": "list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "utf8"
}
}
]
}
},
{
"arrow_type": "large_list<utf8>",
"json": {
"type": "large_list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "utf8"
}
}
]
}
},
{
"arrow_type": "fixed_size_list<float32, 384>",
"json": {
"type": "fixed_size_list",
"length": 384,
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "float32"
}
}
]
}
},
{
"arrow_type": "fixed_size_list<float16, 8>",
"json": {
"type": "fixed_size_list",
"length": 8,
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "float16"
}
}
]
}
},
{
"arrow_type": "fixed_size_list<uint8, 1>",
"json": {
"type": "fixed_size_list",
"length": 1,
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "uint8"
}
}
]
}
},
{
"arrow_type": "list<list<float32>>",
"json": {
"type": "list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "float32"
}
}
]
}
}
]
}
},
{
"arrow_type": "fixed_size_list<list<int32>, 2>",
"json": {
"type": "fixed_size_list",
"length": 2,
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "int32"
}
}
]
}
}
]
}
},
{
"arrow_type": "list<fixed_size_list<float32, 3>>",
"json": {
"type": "list",
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "fixed_size_list",
"length": 3,
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "float32"
}
}
]
}
}
]
}
}
],
"server_only": [
{
"arrow_type": "null",
"json": {
"type": "null"
}
}
],
"invalid": [
"",
"list<>",
"list<float32",
"fixed_size_list<float32>[3]",
"fixed_size_list<float32>",
"fixed_size_list<float32, 0>",
"fixed_size_list<float32, x>",
"map<utf8, int32>",
"decimal128(10, 2)",
"timestamp[us]",
"struct<a: int32>"
]
}
@@ -0,0 +1,78 @@
{
"new_columns": [
{
"name": "embedding",
"all_null": true
}
],
"function": {
"application": {
"function": {
"name": "embed",
"version": "fv_01K3EXACT"
},
"inputs": [
{
"parameter": "text",
"kind": "column",
"value": {
"path": "description"
}
}
],
"output": {
"kind": "scalar",
"arrow_type": "fixed_size_list<float32, 3>",
"nullable": false
}
},
"binding_metadata_version": 1,
"input_bindings": [
{
"parameter": "text",
"field_path": "description",
"arrow_type": "utf8",
"nullable": true
}
],
"input_schema": {
"fields": [
{
"name": "text",
"nullable": true,
"type": {
"type": "utf8"
}
}
]
},
"output_schema": {
"fields": [
{
"name": "embedding",
"nullable": true,
"type": {
"type": "fixed_size_list",
"length": 3,
"fields": [
{
"name": "item",
"nullable": false,
"type": {
"type": "float32"
}
}
]
}
}
]
},
"outputs": [
{
"result_field": "$value",
"output_name": "embedding",
"output_ordinal": 0
}
]
}
}
@@ -1 +1 @@
{"columns":{"normalized_text":"search_text","token_count":"search_token_count"},"function":{"name":"text_features","version":"fv_01K3TEXT"},"group_id":"fg_01K3TEXT","inputs":[{"kind":"column","parameter":"title","value":{"path":"title"}},{"kind":"column","parameter":"body","value":{"path":"body"}}],"output":{"fields":[{"arrow_type":"utf8","name":"normalized_text","nullable":false},{"arrow_type":"int64","name":"token_count","nullable":false}],"kind":"named_struct"}}
{"columns":{"normalized_text":"search_text","token_count":"search_token_count"},"function":{"name":"text_features","version":"fv_01K3TEXT"},"inputs":[{"kind":"column","parameter":"title","value":{"path":"title"}},{"kind":"column","parameter":"body","value":{"path":"body"}}],"output":{"fields":[{"arrow_type":"utf8","name":"normalized_text","nullable":false},{"arrow_type":"int64","name":"token_count","nullable":false}],"kind":"named_struct"}}
@@ -11,7 +11,6 @@
{"name": "token_count", "arrow_type": "int64", "nullable": false}
]
},
"group_id": "fg_01K3TEXT",
"columns": {
"normalized_text": "search_text",
"token_count": "search_token_count"
@@ -3,6 +3,5 @@
"inputs": [
{"parameter": "threshold", "kind": "literal", "value": 1e-7}
],
"output": {"kind": "scalar", "arrow_type": "bool", "nullable": false},
"group_id": "fg_01K3FLOAT"
"output": {"kind": "scalar", "arrow_type": "bool", "nullable": false}
}
@@ -1 +1 @@
{"binding_id":"fb_01K3TEXT","function":{"name":"text_features","version":"fv_01K3TEXT"},"group_id":"fg_01K3TEXT","input_schema":{"fields":[{"name":"title","nullable":true,"type":{"type":"utf8"}},{"name":"body","nullable":true,"type":{"type":"utf8"}}]},"inputs":[{"arrow_type":"utf8","field_id":11,"field_path":"title","nullable":true,"parameter":"title"},{"arrow_type":"utf8","field_id":12,"field_path":"body","nullable":true,"parameter":"body"}],"output_schema":{"fields":[{"name":"search_text","nullable":true,"type":{"type":"utf8"}},{"name":"search_token_count","nullable":true,"type":{"type":"int64"}}]},"outputs":[{"arrow_type":"utf8","nullable":false,"output_field_id":21,"output_name":"search_text","output_ordinal":0,"result_field":"normalized_text"},{"arrow_type":"int64","nullable":false,"output_field_id":22,"output_name":"search_token_count","output_ordinal":1,"result_field":"token_count"}],"revision":3}
{"binding_id":"fb_01K3TEXT","function":{"name":"text_features","version":"fv_01K3TEXT"},"input_schema":{"fields":[{"name":"title","nullable":true,"type":{"type":"utf8"}},{"name":"body","nullable":true,"type":{"type":"utf8"}}]},"inputs":[{"arrow_type":"utf8","field_id":11,"field_path":"title","nullable":true,"parameter":"title"},{"arrow_type":"utf8","field_id":12,"field_path":"body","nullable":true,"parameter":"body"}],"output_schema":{"fields":[{"name":"search_text","nullable":true,"type":{"type":"utf8"}},{"name":"search_token_count","nullable":true,"type":{"type":"int64"}}]},"outputs":[{"arrow_type":"utf8","nullable":false,"output_field_id":21,"output_name":"search_text","output_ordinal":0,"result_field":"normalized_text"},{"arrow_type":"int64","nullable":false,"output_field_id":22,"output_name":"search_token_count","output_ordinal":1,"result_field":"token_count"}]}
@@ -1,8 +1,6 @@
{
"binding_id": "fb_01K3TEXT",
"revision": 3,
"function": {"name": "text_features", "version": "fv_01K3TEXT"},
"group_id": "fg_01K3TEXT",
"inputs": [
{"parameter": "title", "field_id": 11, "field_path": "title", "arrow_type": "utf8", "nullable": true},
{"parameter": "body", "field_id": 12, "field_path": "body", "arrow_type": "utf8", "nullable": true}
@@ -23,5 +21,5 @@
{"name": "search_token_count", "nullable": true, "type": {"type": "int64"}}
]
},
"future_binding": {"metadata_revision": 1}
"future_binding": {"mode": "managed"}
}
@@ -17,7 +17,6 @@
{"name": "token_count", "arrow_type": "int64", "nullable": false}
]
},
"group_id": "fg_01K3TEXT",
"columns": {"normalized_text": "search_text"}
},
"binding_metadata_version": 1,
@@ -3,7 +3,7 @@
"job_type": "refresh_function_columns",
"job_state": "DONE",
"creation_ms": 1787270400001,
"spec": {"table": "documents", "binding_revision": 3},
"spec": {"table": "documents", "binding_id": "fb_01K3TEXT"},
"result": {
"rows_assigned": 999998800,
"rows_failed": 0,
@@ -8,8 +8,7 @@
"inputs": [
{"parameter": "text", "kind": "column", "value": {"path": "description"}}
],
"output": {"kind": "scalar", "arrow_type": "list<float32>", "nullable": false},
"group_id": "fg_scalar"
"output": {"kind": "scalar", "arrow_type": "list<float32>", "nullable": false}
},
"binding_metadata_version": 1,
"input_bindings": [