From 105fd73bc6da822e40f2c7206c8612bd009023bd Mon Sep 17 00:00:00 2001 From: "lancedb-gatefixer[bot]" <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 04:22:22 +0800 Subject: [PATCH 01/23] 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 --------- Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Co-authored-by: Xuanwo --- docs/src/python/python.md | 2 + python/python/lancedb/streaming.py | 594 +++++++++++++++- .../python/tests/test_elastic_dataloader.py | 658 ++++++++++++++++++ 3 files changed, 1221 insertions(+), 33 deletions(-) diff --git a/docs/src/python/python.md b/docs/src/python/python.md index 8a24ea199..3cbeee6f0 100644 --- a/docs/src/python/python.md +++ b/docs/src/python/python.md @@ -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 diff --git a/python/python/lancedb/streaming.py b/python/python/lancedb/streaming.py index c54940ab7..2c0e2d5c1 100644 --- a/python/python/lancedb/streaming.py +++ b/python/python/lancedb/streaming.py @@ -19,6 +19,7 @@ above. """ import ctypes +import heapq import logging import os import random @@ -29,12 +30,12 @@ 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, @@ -55,6 +56,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 +534,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 +562,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 +691,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 +737,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 +746,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 +1059,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 +1105,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 +1157,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 +1311,7 @@ class StreamingDataset(IterableDataset): "_local_consumed_ref", ): state[key] = None + state["_consumer_iterator_lock"] = None return state def __setstate__(self, state): @@ -1074,6 +1322,7 @@ 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: @@ -1083,10 +1332,18 @@ class StreamingDataset(IterableDataset): 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 +1352,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 +1402,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 +1556,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 +1583,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 +1606,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 +1721,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 +1738,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 diff --git a/python/python/tests/test_elastic_dataloader.py b/python/python/tests/test_elastic_dataloader.py index ff65d100c..da27e5bfc 100644 --- a/python/python/tests/test_elastic_dataloader.py +++ b/python/python/tests/test_elastic_dataloader.py @@ -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]]) From 93f47b8aab4f888ab6cd4da75403e2d41be80da4 Mon Sep 17 00:00:00 2001 From: Will Jones Date: Mon, 24 Aug 2026 13:35:31 -0700 Subject: [PATCH 02/23] 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) --- rust/lancedb/src/remote/db.rs | 191 ++++++++++++++++++++++++++++++---- 1 file changed, 171 insertions(+), 20 deletions(-) diff --git a/rust/lancedb/src/remote/db.rs b/rust/lancedb/src/remote/db.rs index 8265d22c0..3f4216bc8 100644 --- a/rust/lancedb/src/remote/db.rs +++ b/rust/lancedb/src/remote/db.rs @@ -344,6 +344,62 @@ impl RemoteDatabase { 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, 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 = None; + let mut page_token: Option = 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 Database for RemoteDatabase { } async fn table_names(&self, request: TableNamesRequest) -> Result> { - 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::() + .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::() - .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| { From c72f5b2960d08a1d6caee703191c360ecfe6f3ff Mon Sep 17 00:00:00 2001 From: Wyatt Alt Date: Mon, 24 Aug 2026 14:06:40 -0700 Subject: [PATCH 03/23] 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. --- rust/lancedb/src/materialized_view.rs | 58 +++- rust/lancedb/src/materialized_view/refresh.rs | 265 +++++++++++++++++- 2 files changed, 307 insertions(+), 16 deletions(-) diff --git a/rust/lancedb/src/materialized_view.rs b/rust/lancedb/src/materialized_view.rs index 1d7cbb5e7..b28d52931 100644 --- a/rust/lancedb/src/materialized_view.rs +++ b/rust/lancedb/src/materialized_view.rs @@ -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 { let empty: Vec> = 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 = - 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, } 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, + expected_incarnation: Option, } 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) -> Self { + self.expected_incarnation = Some(token.into()); + self + } + pub async fn execute(self) -> Result { - 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 } } diff --git a/rust/lancedb/src/materialized_view/refresh.rs b/rust/lancedb/src/materialized_view/refresh.rs index 24d828e01..735751c27 100644 --- a/rust/lancedb/src/materialized_view/refresh.rs +++ b/rust/lancedb/src/materialized_view/refresh.rs @@ -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, + expected_incarnation: Option<&str>, ) -> Result { 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, + expected_incarnation: Option<&str>, ) -> Result> { 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 { 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>, + expected_incarnation: Option<&str>, ) -> Result { 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 = 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 { + 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, Vec)>, new_fragments: Vec, keys: Option, + expected_incarnation: Option<&str>, ) -> Result { 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::::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::::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(); From 71f85a8d9fab89c238648c9963e7c89e1452e74d Mon Sep 17 00:00:00 2001 From: Lance Release Date: Mon, 24 Aug 2026 21:08:39 +0000 Subject: [PATCH 04/23] =?UTF-8?q?Bump=20version:=200.38.0-beta.6=20?= =?UTF-8?q?=E2=86=92=200.38.0-beta.7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .bumpversion.toml | 2 +- Cargo.lock | 6 +++--- docs/src/java/java.md | 2 +- java/lancedb-core/pom.xml | 2 +- java/pom.xml | 2 +- nodejs/Cargo.toml | 2 +- nodejs/npm/darwin-arm64/package.json | 2 +- nodejs/npm/linux-arm64-gnu/package.json | 2 +- nodejs/npm/linux-arm64-musl/package.json | 2 +- nodejs/npm/linux-x64-gnu/package.json | 2 +- nodejs/npm/linux-x64-musl/package.json | 2 +- nodejs/npm/win32-arm64-msvc/package.json | 2 +- nodejs/npm/win32-x64-msvc/package.json | 2 +- nodejs/package-lock.json | 4 ++-- nodejs/package.json | 2 +- python/Cargo.toml | 2 +- rust/lancedb/Cargo.toml | 2 +- 17 files changed, 20 insertions(+), 20 deletions(-) diff --git a/.bumpversion.toml b/.bumpversion.toml index f5f981b2c..bcf714002 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.38.0-beta.6" +current_version = "0.38.0-beta.7" parse = """(?x) (?P0|[1-9]\\d*)\\. (?P0|[1-9]\\d*)\\. diff --git a/Cargo.lock b/Cargo.lock index 6c6dd4aea..a60a5ddf1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5402,7 +5402,7 @@ dependencies = [ [[package]] name = "lancedb" -version = "0.38.0-beta.6" +version = "0.38.0-beta.7" dependencies = [ "ahash", "anyhow", @@ -5490,7 +5490,7 @@ dependencies = [ [[package]] name = "lancedb-nodejs" -version = "0.38.0-beta.6" +version = "0.38.0-beta.7" 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.7" dependencies = [ "arrow", "async-trait", diff --git a/docs/src/java/java.md b/docs/src/java/java.md index d6c285fc3..59ee8a7d6 100644 --- a/docs/src/java/java.md +++ b/docs/src/java/java.md @@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`: com.lancedb lancedb-core - 0.38.0-beta.6 + 0.38.0-beta.7 ``` diff --git a/java/lancedb-core/pom.xml b/java/lancedb-core/pom.xml index 23d58124d..b24ad61ab 100644 --- a/java/lancedb-core/pom.xml +++ b/java/lancedb-core/pom.xml @@ -8,7 +8,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.6 + 0.38.0-beta.7 ../pom.xml diff --git a/java/pom.xml b/java/pom.xml index a07414ec0..0ee59bcb7 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -6,7 +6,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.6 + 0.38.0-beta.7 pom ${project.artifactId} LanceDB Java SDK Parent POM diff --git a/nodejs/Cargo.toml b/nodejs/Cargo.toml index ead873855..8732d1965 100644 --- a/nodejs/Cargo.toml +++ b/nodejs/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "lancedb-nodejs" edition.workspace = true -version = "0.38.0-beta.6" +version = "0.38.0-beta.7" publish = false license.workspace = true description.workspace = true diff --git a/nodejs/npm/darwin-arm64/package.json b/nodejs/npm/darwin-arm64/package.json index d89eb5ad1..56480c0cb 100644 --- a/nodejs/npm/darwin-arm64/package.json +++ b/nodejs/npm/darwin-arm64/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-darwin-arm64", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "os": ["darwin"], "cpu": ["arm64"], "main": "lancedb.darwin-arm64.node", diff --git a/nodejs/npm/linux-arm64-gnu/package.json b/nodejs/npm/linux-arm64-gnu/package.json index 1721c970e..a1102ee29 100644 --- a/nodejs/npm/linux-arm64-gnu/package.json +++ b/nodejs/npm/linux-arm64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-gnu", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-gnu.node", diff --git a/nodejs/npm/linux-arm64-musl/package.json b/nodejs/npm/linux-arm64-musl/package.json index 177804dce..695862baf 100644 --- a/nodejs/npm/linux-arm64-musl/package.json +++ b/nodejs/npm/linux-arm64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-musl", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-musl.node", diff --git a/nodejs/npm/linux-x64-gnu/package.json b/nodejs/npm/linux-x64-gnu/package.json index 9af443583..9099ee4de 100644 --- a/nodejs/npm/linux-x64-gnu/package.json +++ b/nodejs/npm/linux-x64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-gnu", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-gnu.node", diff --git a/nodejs/npm/linux-x64-musl/package.json b/nodejs/npm/linux-x64-musl/package.json index c98aa3812..f04c2a884 100644 --- a/nodejs/npm/linux-x64-musl/package.json +++ b/nodejs/npm/linux-x64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-musl", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-musl.node", diff --git a/nodejs/npm/win32-arm64-msvc/package.json b/nodejs/npm/win32-arm64-msvc/package.json index 05b1086f5..f9b35af32 100644 --- a/nodejs/npm/win32-arm64-msvc/package.json +++ b/nodejs/npm/win32-arm64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-arm64-msvc", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "os": [ "win32" ], diff --git a/nodejs/npm/win32-x64-msvc/package.json b/nodejs/npm/win32-x64-msvc/package.json index 50432bdea..cc5340105 100644 --- a/nodejs/npm/win32-x64-msvc/package.json +++ b/nodejs/npm/win32-x64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-x64-msvc", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "os": ["win32"], "cpu": ["x64"], "main": "lancedb.win32-x64-msvc.node", diff --git a/nodejs/package-lock.json b/nodejs/package-lock.json index 219025534..6604da6cd 100644 --- a/nodejs/package-lock.json +++ b/nodejs/package-lock.json @@ -1,12 +1,12 @@ { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "cpu": [ "x64", "arm64" diff --git a/nodejs/package.json b/nodejs/package.json index 63c51a799..66777d60c 100644 --- a/nodejs/package.json +++ b/nodejs/package.json @@ -11,7 +11,7 @@ "ann" ], "private": false, - "version": "0.38.0-beta.6", + "version": "0.38.0-beta.7", "main": "dist/index.js", "exports": { ".": "./dist/index.js", diff --git a/python/Cargo.toml b/python/Cargo.toml index 2e4f8996f..d2ccf3064 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb-python" -version = "0.38.0-beta.6" +version = "0.38.0-beta.7" publish = false edition.workspace = true description = "Python bindings for LanceDB" diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index d32d88ecd..071c876d9 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb" -version = "0.38.0-beta.6" +version = "0.38.0-beta.7" edition.workspace = true description = "LanceDB: A serverless, low-latency vector database for AI applications" license.workspace = true From 5013c176dd9a664813888c8bd8cf99161c047ec7 Mon Sep 17 00:00:00 2001 From: "lancedb-gatefixer[bot]" <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 07:14:33 +0800 Subject: [PATCH 05/23] fix: pin remote snapshots during permutation construction (#4022) Fixes #4015 ## 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 --- .../src/dataloader/permutation/builder.rs | 128 +++++++++++++++++- rust/lancedb/src/remote/table.rs | 12 ++ rust/lancedb/src/table.rs | 8 ++ 3 files changed, 146 insertions(+), 2 deletions(-) diff --git a/rust/lancedb/src/dataloader/permutation/builder.rs b/rust/lancedb/src/dataloader/permutation/builder.rs index a3700d0ef..641c191d5 100644 --- a/rust/lancedb/src/dataloader/permutation/builder.rs +++ b/rust/lancedb/src/dataloader/permutation/builder.rs @@ -239,7 +239,20 @@ impl PermutationBuilder { } /// Builds the permutation table and stores it in the given database. - pub async fn build(self) -> Result { + pub async fn build(mut self) -> Result
{ + // 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(_)) => { @@ -258,7 +271,6 @@ impl PermutationBuilder { // 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 { @@ -409,6 +421,118 @@ 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::(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_builder() { let temp_dir = tempfile::tempdir().unwrap(); diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index c6b078282..633d0ffce 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -1775,6 +1775,18 @@ impl BaseTable for RemoteTable { Ok(()) } + async fn snapshot_at_current_version(&self) -> Result>> { + // 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 diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index 41551bd37..356c87bbe 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -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>> { + 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 From fce45ba9fc5254d3e7fff8fc90c89b99706f120c Mon Sep 17 00:00:00 2001 From: Will Jones Date: Mon, 24 Aug 2026 16:25:14 -0700 Subject: [PATCH 06/23] feat(nodejs): add listTables, deprecate tableNames (#4041) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `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) --- docs/src/js/classes/Connection.md | 74 ++++++++++++++++- docs/src/js/globals.md | 2 + docs/src/js/interfaces/ListTablesOptions.md | 34 ++++++++ docs/src/js/interfaces/ListTablesResponse.md | 23 ++++++ docs/src/js/interfaces/TableNamesOptions.md | 11 ++- nodejs/__test__/connection.test.ts | 69 +++++++++++++++- nodejs/lancedb/connection.ts | 85 ++++++++++++++++++++ nodejs/lancedb/index.ts | 2 + nodejs/src/connection.rs | 34 ++++++++ 9 files changed, 329 insertions(+), 5 deletions(-) create mode 100644 docs/src/js/interfaces/ListTablesOptions.md create mode 100644 docs/src/js/interfaces/ListTablesResponse.md diff --git a/docs/src/js/classes/Connection.md b/docs/src/js/classes/Connection.md index 92cfd2568..1c5abd89f 100644 --- a/docs/src/js/classes/Connection.md +++ b/docs/src/js/classes/Connection.md @@ -584,6 +584,70 @@ Child namespace names and *** +### listTables() + +#### listTables(options) + +```ts +abstract listTables(options?): Promise +``` + +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`<[`ListTablesOptions`](../interfaces/ListTablesOptions.md)> + Pagination options + (`pageToken`, `limit`). + +##### Returns + +`Promise`<[`ListTablesResponse`](../interfaces/ListTablesResponse.md)> + +A page of table names and an + optional token for the tables after it. + +#### listTables(namespacePath, options) + +```ts +abstract listTables(namespacePath?, options?): Promise +``` + +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`<[`ListTablesOptions`](../interfaces/ListTablesOptions.md)> + Pagination options + (`pageToken`, `limit`). + +##### Returns + +`Promise`<[`ListTablesResponse`](../interfaces/ListTablesResponse.md)> + +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`<`string`[]> +##### 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`<`string`[]> + +##### Deprecated + +Use [Connection.listTables](Connection.md#listtables) instead. diff --git a/docs/src/js/globals.md b/docs/src/js/globals.md index 6a8f644eb..e0635ab65 100644 --- a/docs/src/js/globals.md +++ b/docs/src/js/globals.md @@ -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) diff --git a/docs/src/js/interfaces/ListTablesOptions.md b/docs/src/js/interfaces/ListTablesOptions.md new file mode 100644 index 000000000..52ace47e5 --- /dev/null +++ b/docs/src/js/interfaces/ListTablesOptions.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. diff --git a/docs/src/js/interfaces/ListTablesResponse.md b/docs/src/js/interfaces/ListTablesResponse.md new file mode 100644 index 000000000..76cac2b23 --- /dev/null +++ b/docs/src/js/interfaces/ListTablesResponse.md @@ -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[]; +``` diff --git a/docs/src/js/interfaces/TableNamesOptions.md b/docs/src/js/interfaces/TableNamesOptions.md index 45fa3d1b0..9254e9fa7 100644 --- a/docs/src/js/interfaces/TableNamesOptions.md +++ b/docs/src/js/interfaces/TableNamesOptions.md @@ -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; diff --git a/nodejs/__test__/connection.test.ts b/nodejs/__test__/connection.test.ts index af471b478..816f8c3de 100644 --- a/nodejs/__test__/connection.test.ts +++ b/nodejs/__test__/connection.test.ts @@ -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 })); diff --git a/nodejs/lancedb/connection.ts b/nodejs/lancedb/connection.ts index e5528f8c6..263a338ab 100644 --- a/nodejs/lancedb/connection.ts +++ b/nodejs/lancedb/connection.ts @@ -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} options - options to control the * paging / start point (backwards compatibility) * + * @deprecated Use {@link Connection.listTables} instead. */ abstract tableNames(options?: Partial): Promise; /** @@ -241,12 +266,53 @@ export abstract class Connection { * @param {Partial} options - options to control the * paging / start point * + * @deprecated Use {@link Connection.listTables} instead. */ abstract tableNames( namespacePath?: string[], options?: Partial, ): Promise; + /** + * 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} options - Pagination options + * (`pageToken`, `limit`). + * @returns {Promise} A page of table names and an + * optional token for the tables after it. + */ + abstract listTables( + options?: Partial, + ): Promise; + /** + * 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} options - Pagination options + * (`pageToken`, `limit`). + * @returns {Promise} A page of table names and an + * optional token for the tables after it. + */ + abstract listTables( + namespacePath?: string[], + options?: Partial, + ): Promise; + /** * 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, + options?: Partial, + ): Promise { + // 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[], diff --git a/nodejs/lancedb/index.ts b/nodejs/lancedb/index.ts index da13b434f..ebc8cda8d 100644 --- a/nodejs/lancedb/index.ts +++ b/nodejs/lancedb/index.ts @@ -81,11 +81,13 @@ export { Connection, CreateTableOptions, TableNamesOptions, + ListTablesOptions, OpenTableOptions, ListNamespacesOptions, CreateNamespaceOptions, DropNamespaceOptions, ListNamespacesResponse, + ListTablesResponse, CreateNamespaceResponse, DropNamespaceResponse, DescribeNamespaceResponse, diff --git a/nodejs/src/connection.rs b/nodejs/src/connection.rs index 44eb68f32..5cf676256 100644 --- a/nodejs/src/connection.rs +++ b/nodejs/src/connection.rs @@ -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, } +#[napi(object)] +pub struct ListTablesResponse { + pub tables: Vec, + pub page_token: Option, +} + #[napi(object)] pub struct CreateNamespaceResponse { pub properties: Option>, @@ -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>, + page_token: Option, + limit: Option, + ) -> napi::Result { + 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: From c1a8c3f089fbab95883bd58d9a568d4c30d41074 Mon Sep 17 00:00:00 2001 From: "lancedb-gatefixer[bot]" <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Mon, 24 Aug 2026 16:41:05 -0700 Subject: [PATCH 07/23] 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 --------- Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> --- nodejs/__test__/arrow.test.ts | 131 ++++++++ nodejs/lancedb/arrow.ts | 340 ++------------------ nodejs/lancedb/arrow_type.ts | 40 +++ nodejs/lancedb/schema.ts | 566 ++++++++++++++++++++++++++++++++++ 4 files changed, 770 insertions(+), 307 deletions(-) create mode 100644 nodejs/lancedb/arrow_type.ts create mode 100644 nodejs/lancedb/schema.ts diff --git a/nodejs/__test__/arrow.test.ts b/nodejs/__test__/arrow.test.ts index 160f4d5ef..c5bbbf169 100644 --- a/nodejs/__test__/arrow.test.ts +++ b/nodejs/__test__/arrow.test.ts @@ -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] 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) => diff --git a/nodejs/lancedb/arrow.ts b/nodejs/lancedb/arrow.ts index 8b388d593..b52ab50ef 100644 --- a/nodejs/lancedb/arrow.ts +++ b/nodejs/lancedb/arrow.ts @@ -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>, - schema: Schema | undefined, - opts: MakeArrowTableOptions, -): Schema { - // We will collect all fields we see in the data. - const pathTree = new PathTree(); - - 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): 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, - ): 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, - 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 { return ( typeof value === "object" && @@ -577,146 +459,19 @@ function isObject(value: unknown): value is Record { ); } -function getFieldForPath(schema: Schema, path: string[]): Field | undefined { - let current: Field | Schema = schema; +function valueAtPath(datum: Record, 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 { - map: Map>; - - 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 = this; - for (const part of path) { - if (!(ref instanceof PathTree) || !ref.map.has(part)) { - return false; - } - ref = ref.map.get(part) as PathTree; - } - return true; - } - get(path: string[]): V | undefined { - let ref: PathTree = this; - for (const part of path) { - if (!(ref instanceof PathTree) || !ref.map.has(part)) { - return undefined; - } - ref = ref.map.get(part) as PathTree; - } - return ref as V; - } - set(path: string[], value: V): void { - let ref: PathTree = this; - for (const part of path.slice(0, path.length - 1)) { - if (!ref.map.has(part)) { - ref.map.set(part, new PathTree()); - } - ref = ref.map.get(part) as PathTree; - } - 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[], }); 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 { 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; } } diff --git a/nodejs/lancedb/arrow_type.ts b/nodejs/lancedb/arrow_type.ts new file mode 100644 index 000000000..35da346ef --- /dev/null +++ b/nodejs/lancedb/arrow_type.ts @@ -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; +} diff --git a/nodejs/lancedb/schema.ts b/nodejs/lancedb/schema.ts new file mode 100644 index 000000000..e4749ef37 --- /dev/null +++ b/nodejs/lancedb/schema.ts @@ -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; +}; + +/** + * 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>, + 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>): 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; +type FieldConflict = { path: string[]; value: FieldNode }; + +/** Nested field state, kept separate from Arrow's eventual Struct types. */ +class FieldTree { + private readonly children = new Map(); + + 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, + 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 { + 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"); +} From 6ed3074d4cf98b08bec11269cd2c1fba5bf8f7d3 Mon Sep 17 00:00:00 2001 From: Jack Ye Date: Mon, 24 Aug 2026 17:16:42 -0700 Subject: [PATCH 08/23] feat: pin the base table version for data loader reads (#3982) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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` 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>>` 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. --- python/python/lancedb/permutation.py | 26 ++- python/python/lancedb/streaming.py | 4 + python/python/tests/test_permutation.py | 25 +++ .../src/dataloader/permutation/builder.rs | 186 ++++++++++++++++-- .../src/dataloader/permutation/reader.rs | 93 ++++++++- rust/lancedb/src/remote/table.rs | 37 ++++ rust/lancedb/src/table/query.rs | 3 + 7 files changed, 350 insertions(+), 24 deletions(-) diff --git a/python/python/lancedb/permutation.py b/python/python/lancedb/permutation.py index 5d7a685ef..8287ef2fe 100644 --- a/python/python/lancedb/permutation.py +++ b/python/python/lancedb/permutation.py @@ -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 diff --git a/python/python/lancedb/streaming.py b/python/python/lancedb/streaming.py index 2c0e2d5c1..76b3702a1 100644 --- a/python/python/lancedb/streaming.py +++ b/python/python/lancedb/streaming.py @@ -41,6 +41,7 @@ from .permutation import ( Permutation, Transforms, permutation_builder, + _drop_base_version, _table_from_pickle_state, _table_to_pickle_state, ) @@ -1327,6 +1328,9 @@ class StreamingDataset(IterableDataset): 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: diff --git a/python/python/tests/test_permutation.py b/python/python/tests/test_permutation.py index 135742c84..142fc84f1 100644 --- a/python/python/tests/test_permutation.py +++ b/python/python/tests/test_permutation.py @@ -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( diff --git a/rust/lancedb/src/dataloader/permutation/builder.rs b/rust/lancedb/src/dataloader/permutation/builder.rs index 641c191d5..ba0d0e180 100644 --- a/rust/lancedb/src/dataloader/permutation/builder.rs +++ b/rust/lancedb/src/dataloader/permutation/builder.rs @@ -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, ) -> Result { - 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()); @@ -269,6 +265,12 @@ 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. let mut rows = self.base_table.query().select(Select::columns(&[ROW_ID])); @@ -330,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) => { @@ -533,6 +548,141 @@ mod tests { 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::()) + .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::() + .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::()) + .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::() + .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::()) + .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::()) + .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::()) + .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(); diff --git a/rust/lancedb/src/dataloader/permutation/reader.rs b/rust/lancedb/src/dataloader/permutation/reader.rs index 6da92e986..9757dc552 100644 --- a/rust/lancedb/src/dataloader/permutation/reader.rs +++ b/rust/lancedb/src/dataloader/permutation/reader.rs @@ -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>, split: u64, ) -> Result { + 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, + permutation_table: &Arc, + ) -> Result> { + 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::().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, permutation_table: Arc, @@ -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::()) + .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::( + 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() diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 633d0ffce..10980d55a 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -4349,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"); diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 9feb9d5ab..9b81786cd 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -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); From 0e65123bd8910548c1ad60b4038f0674c33ed940 Mon Sep 17 00:00:00 2001 From: Wyatt Alt Date: Mon, 24 Aug 2026 23:15:07 -0700 Subject: [PATCH 09/23] 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[n]`, `timestamp[us]`, `struct<...>`, zero-sized lists). It now emits exactly the grammar, with the server's `fixed_size_list` 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. --- python/python/lancedb/functions.py | 248 ++++++++--- .../tests/test_first_class_function_slice2.py | 399 +++++++++++++++++- rust/lancedb/src/remote/table.rs | 47 +++ rust/lancedb/src/table/computed_columns.rs | 55 +++ .../first_class_functions/v1/arrow_types.json | 333 +++++++++++++++ ...remote_fixed_size_declaration_request.json | 79 ++++ 6 files changed, 1100 insertions(+), 61 deletions(-) create mode 100644 rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json create mode 100644 rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json diff --git a/python/python/lancedb/functions.py b/python/python/lancedb/functions.py index c4ff9a9b3..33a67b413 100644 --- a/python/python/lancedb/functions.py +++ b/python/python/lancedb/functions.py @@ -12,10 +12,13 @@ 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 @@ -489,59 +492,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 +591,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 +738,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, "", "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 +860,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 +1046,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 diff --git a/python/python/tests/test_first_class_function_slice2.py b/python/python/tests/test_first_class_function_slice2.py index a20347771..c67e17520 100644 --- a/python/python/tests/test_first_class_function_slice2.py +++ b/python/python/tests/test_first_class_function_slice2.py @@ -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, "", "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" + ) + + +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[3]" + assert signature.output.arrow_type == "fixed_size_list" assert signature.output.nullable is False diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 10980d55a..35ac56970 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -6802,6 +6802,53 @@ mod tests { assert_eq!(result.version, 8); } + #[tokio::test] + 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","nullable":false}, + "group_id":"fg_fixed" + }"#, + ) + .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_sibling_group() { let table = Table::new_with_handler("my_table", |request| match request.url().path() { diff --git a/rust/lancedb/src/table/computed_columns.rs b/rust/lancedb/src/table/computed_columns.rs index b1c0b8370..133f37a55 100644 --- a/rust/lancedb/src/table/computed_columns.rs +++ b/rust/lancedb/src/table/computed_columns.rs @@ -586,6 +586,25 @@ fn canonical_input_arrow_type(field: &JsonArrowField) -> Result { } } +/// `fixed_size_list` -> (`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 { fn parse(raw: &str) -> Result { let raw = raw.trim(); @@ -618,6 +637,16 @@ fn parse_output_arrow_type(raw: &str) -> Result { )]); 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", @@ -1322,6 +1351,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; diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json new file mode 100644 index 000000000..ff26e4c08 --- /dev/null +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json @@ -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", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + }, + { + "arrow_type": "large_list", + "json": { + "type": "large_list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + }, + { + "arrow_type": "list", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "int64" + } + } + ] + } + }, + { + "arrow_type": "large_list", + "json": { + "type": "large_list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "int64" + } + } + ] + } + }, + { + "arrow_type": "list", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "utf8" + } + } + ] + } + }, + { + "arrow_type": "large_list", + "json": { + "type": "large_list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "utf8" + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list", + "json": { + "type": "fixed_size_list", + "length": 384, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list", + "json": { + "type": "fixed_size_list", + "length": 8, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float16" + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list", + "json": { + "type": "fixed_size_list", + "length": 1, + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "uint8" + } + } + ] + } + }, + { + "arrow_type": "list>", + "json": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "list", + "fields": [ + { + "name": "item", + "nullable": false, + "type": { + "type": "float32" + } + } + ] + } + } + ] + } + }, + { + "arrow_type": "fixed_size_list, 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>", + "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[3]", + "fixed_size_list", + "fixed_size_list", + "fixed_size_list", + "map", + "decimal128(10, 2)", + "timestamp[us]", + "struct" + ] +} \ No newline at end of file diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json new file mode 100644 index 000000000..7d92f1454 --- /dev/null +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json @@ -0,0 +1,79 @@ +{ + "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", + "nullable": false + }, + "group_id": "fg_fixed" + }, + "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 + } + ] + } +} \ No newline at end of file From 2fea7cd48d86ef149ab7eee93f7751efeb410233 Mon Sep 17 00:00:00 2001 From: Lance Release Date: Tue, 25 Aug 2026 06:26:01 +0000 Subject: [PATCH 10/23] =?UTF-8?q?Bump=20version:=200.38.0-beta.7=20?= =?UTF-8?q?=E2=86=92=200.38.0-beta.8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .bumpversion.toml | 2 +- Cargo.lock | 6 +++--- docs/src/java/java.md | 2 +- java/lancedb-core/pom.xml | 2 +- java/pom.xml | 2 +- nodejs/Cargo.toml | 2 +- nodejs/npm/darwin-arm64/package.json | 2 +- nodejs/npm/linux-arm64-gnu/package.json | 2 +- nodejs/npm/linux-arm64-musl/package.json | 2 +- nodejs/npm/linux-x64-gnu/package.json | 2 +- nodejs/npm/linux-x64-musl/package.json | 2 +- nodejs/npm/win32-arm64-msvc/package.json | 2 +- nodejs/npm/win32-x64-msvc/package.json | 2 +- nodejs/package-lock.json | 4 ++-- nodejs/package.json | 2 +- python/Cargo.toml | 2 +- rust/lancedb/Cargo.toml | 2 +- 17 files changed, 20 insertions(+), 20 deletions(-) diff --git a/.bumpversion.toml b/.bumpversion.toml index bcf714002..86a5064f7 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.38.0-beta.7" +current_version = "0.38.0-beta.8" parse = """(?x) (?P0|[1-9]\\d*)\\. (?P0|[1-9]\\d*)\\. diff --git a/Cargo.lock b/Cargo.lock index a60a5ddf1..8d0058569 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5402,7 +5402,7 @@ dependencies = [ [[package]] name = "lancedb" -version = "0.38.0-beta.7" +version = "0.38.0-beta.8" dependencies = [ "ahash", "anyhow", @@ -5490,7 +5490,7 @@ dependencies = [ [[package]] name = "lancedb-nodejs" -version = "0.38.0-beta.7" +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.7" +version = "0.38.0-beta.8" dependencies = [ "arrow", "async-trait", diff --git a/docs/src/java/java.md b/docs/src/java/java.md index 59ee8a7d6..28de11c30 100644 --- a/docs/src/java/java.md +++ b/docs/src/java/java.md @@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`: com.lancedb lancedb-core - 0.38.0-beta.7 + 0.38.0-beta.8 ``` diff --git a/java/lancedb-core/pom.xml b/java/lancedb-core/pom.xml index b24ad61ab..6f4b0744b 100644 --- a/java/lancedb-core/pom.xml +++ b/java/lancedb-core/pom.xml @@ -8,7 +8,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.7 + 0.38.0-beta.8 ../pom.xml diff --git a/java/pom.xml b/java/pom.xml index 0ee59bcb7..ac6196a45 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -6,7 +6,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.7 + 0.38.0-beta.8 pom ${project.artifactId} LanceDB Java SDK Parent POM diff --git a/nodejs/Cargo.toml b/nodejs/Cargo.toml index 8732d1965..e85da66ec 100644 --- a/nodejs/Cargo.toml +++ b/nodejs/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "lancedb-nodejs" edition.workspace = true -version = "0.38.0-beta.7" +version = "0.38.0-beta.8" publish = false license.workspace = true description.workspace = true diff --git a/nodejs/npm/darwin-arm64/package.json b/nodejs/npm/darwin-arm64/package.json index 56480c0cb..f920c9e91 100644 --- a/nodejs/npm/darwin-arm64/package.json +++ b/nodejs/npm/darwin-arm64/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-darwin-arm64", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "os": ["darwin"], "cpu": ["arm64"], "main": "lancedb.darwin-arm64.node", diff --git a/nodejs/npm/linux-arm64-gnu/package.json b/nodejs/npm/linux-arm64-gnu/package.json index a1102ee29..ad6926039 100644 --- a/nodejs/npm/linux-arm64-gnu/package.json +++ b/nodejs/npm/linux-arm64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-gnu", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-gnu.node", diff --git a/nodejs/npm/linux-arm64-musl/package.json b/nodejs/npm/linux-arm64-musl/package.json index 695862baf..ccd39f18e 100644 --- a/nodejs/npm/linux-arm64-musl/package.json +++ b/nodejs/npm/linux-arm64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-musl", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-musl.node", diff --git a/nodejs/npm/linux-x64-gnu/package.json b/nodejs/npm/linux-x64-gnu/package.json index 9099ee4de..50446b597 100644 --- a/nodejs/npm/linux-x64-gnu/package.json +++ b/nodejs/npm/linux-x64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-gnu", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-gnu.node", diff --git a/nodejs/npm/linux-x64-musl/package.json b/nodejs/npm/linux-x64-musl/package.json index f04c2a884..d4da82dc7 100644 --- a/nodejs/npm/linux-x64-musl/package.json +++ b/nodejs/npm/linux-x64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-musl", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-musl.node", diff --git a/nodejs/npm/win32-arm64-msvc/package.json b/nodejs/npm/win32-arm64-msvc/package.json index f9b35af32..dddc3a470 100644 --- a/nodejs/npm/win32-arm64-msvc/package.json +++ b/nodejs/npm/win32-arm64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-arm64-msvc", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "os": [ "win32" ], diff --git a/nodejs/npm/win32-x64-msvc/package.json b/nodejs/npm/win32-x64-msvc/package.json index cc5340105..23ac122f7 100644 --- a/nodejs/npm/win32-x64-msvc/package.json +++ b/nodejs/npm/win32-x64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-x64-msvc", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "os": ["win32"], "cpu": ["x64"], "main": "lancedb.win32-x64-msvc.node", diff --git a/nodejs/package-lock.json b/nodejs/package-lock.json index 6604da6cd..21deb8dd8 100644 --- a/nodejs/package-lock.json +++ b/nodejs/package-lock.json @@ -1,12 +1,12 @@ { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "cpu": [ "x64", "arm64" diff --git a/nodejs/package.json b/nodejs/package.json index 66777d60c..96f4bafd5 100644 --- a/nodejs/package.json +++ b/nodejs/package.json @@ -11,7 +11,7 @@ "ann" ], "private": false, - "version": "0.38.0-beta.7", + "version": "0.38.0-beta.8", "main": "dist/index.js", "exports": { ".": "./dist/index.js", diff --git a/python/Cargo.toml b/python/Cargo.toml index d2ccf3064..bffac2ef1 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb-python" -version = "0.38.0-beta.7" +version = "0.38.0-beta.8" publish = false edition.workspace = true description = "Python bindings for LanceDB" diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index 071c876d9..7290694f1 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb" -version = "0.38.0-beta.7" +version = "0.38.0-beta.8" edition.workspace = true description = "LanceDB: A serverless, low-latency vector database for AI applications" license.workspace = true From c988e4848dade1ebca7ff73a669190761c3e34ac Mon Sep 17 00:00:00 2001 From: Xuanwo Date: Tue, 25 Aug 2026 18:30:09 +0800 Subject: [PATCH 11/23] 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. --- python/python/lancedb/functions.py | 11 ++--- python/python/lancedb/table.py | 6 +-- .../tests/test_first_class_function_slice1.py | 16 +++---- rust/lancedb/src/function.rs | 21 ++------- rust/lancedb/src/query.rs | 9 ++-- rust/lancedb/src/remote/table.rs | 11 ++--- rust/lancedb/src/table.rs | 2 +- rust/lancedb/src/table/add_columns.rs | 2 +- rust/lancedb/src/table/computed_columns.rs | 45 +++++++------------ .../tests/first_class_function_slice1.rs | 1 - ...remote_fixed_size_declaration_request.json | 5 +-- ...remote_function_application.canonical.json | 2 +- .../v1/remote_function_application.json | 1 - .../v1/remote_function_application_float.json | 3 +- .../v1/remote_function_binding.canonical.json | 2 +- .../v1/remote_function_binding.json | 4 +- ...ote_multi_output_declaration_request.json} | 1 - .../v1/remote_refresh_job.json | 2 +- .../v1/remote_scalar_declaration_request.json | 3 +- 19 files changed, 49 insertions(+), 98 deletions(-) rename rust/lancedb/tests/fixtures/first_class_functions/v1/{remote_grouped_declaration_request.json => remote_multi_output_declaration_request.json} (97%) diff --git a/python/python/lancedb/functions.py b/python/python/lancedb/functions.py index 33a67b413..dd8cecfa9 100644 --- a/python/python/lancedb/functions.py +++ b/python/python/lancedb/functions.py @@ -25,7 +25,6 @@ import re import sys import textwrap import types -import uuid from collections.abc import Mapping from datetime import date, datetime from typing import ( @@ -276,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], @@ -326,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}", ) @@ -370,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 @@ -380,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]: @@ -452,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 diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index 75b7b1db8..fe25d5353 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -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())) diff --git a/python/python/tests/test_first_class_function_slice1.py b/python/python/tests/test_first_class_function_slice1.py index 1bb8feaf4..ca0b30ede 100644 --- a/python/python/tests/test_first_class_function_slice1.py +++ b/python/python/tests/test_first_class_function_slice1.py @@ -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", "nullable": False, }, - "group_id": "fg_scalar", } ) ) diff --git a/rust/lancedb/src/function.rs b/rust/lancedb/src/function.rs index f02158871..433a7e04c 100644 --- a/rust/lancedb/src/function.rs +++ b/rust/lancedb/src/function.rs @@ -446,7 +446,6 @@ pub struct FunctionApplication { function: FunctionVersionRef, inputs: Vec, output: FunctionOutput, - group_id: String, #[serde(default, skip_serializing_if = "BTreeMap::is_empty")] columns: BTreeMap, #[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 { &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, outputs: Vec, /// 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, - /// 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, } @@ -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 } diff --git a/rust/lancedb/src/query.rs b/rust/lancedb/src/query.rs index b2c5fefbe..654777adb 100644 --- a/rust/lancedb/src/query.rs +++ b/rust/lancedb/src/query.rs @@ -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() diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 35ac56970..742ab547d 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -6787,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","nullable":false}, - "group_id":"fg_scalar" + "output":{"kind":"scalar","arrow_type":"list","nullable":false} }"#, ) .unwrap(); @@ -6834,8 +6833,7 @@ mod tests { r#"{ "function":{"name":"embed","version":"fv_01K3EXACT"}, "inputs":[{"parameter":"text","kind":"column","value":{"path":"description"}}], - "output":{"kind":"scalar","arrow_type":"fixed_size_list","nullable":false}, - "group_id":"fg_fixed" + "output":{"kind":"scalar","arrow_type":"fixed_size_list","nullable":false} }"#, ) .unwrap(); @@ -6850,7 +6848,7 @@ mod tests { } #[tokio::test] - async fn test_add_named_struct_function_expands_one_atomic_sibling_group() { + 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) @@ -6865,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); @@ -6887,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"} }"#, ) diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index 356c87bbe..ecd95f161 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -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, diff --git a/rust/lancedb/src/table/add_columns.rs b/rust/lancedb/src/table/add_columns.rs index 3b91c30e4..1ac0c6b4f 100644 --- a/rust/lancedb/src/table/add_columns.rs +++ b/rust/lancedb/src/table/add_columns.rs @@ -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). /// /// ``` diff --git a/rust/lancedb/src/table/computed_columns.rs b/rust/lancedb/src/table/computed_columns.rs index 133f37a55..2b95cb34c 100644 --- a/rust/lancedb/src/table/computed_columns.rs +++ b/rust/lancedb/src/table/computed_columns.rs @@ -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, } -/// Encode immutable grouped bindings for schema-level persistence. +/// Encode immutable Function bindings for schema-level persistence. pub fn function_bindings_metadata(bindings: &[FunctionBinding]) -> Result { let bindings = bindings .iter() @@ -177,7 +177,7 @@ pub fn function_bindings_metadata(bindings: &[FunctionBinding]) -> Result Result> { 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", @@ -786,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", )); } @@ -2261,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} }}"# )) @@ -2442,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(); @@ -2456,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"} }"#, ) @@ -2470,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(); diff --git a/rust/lancedb/tests/first_class_function_slice1.rs b/rust/lancedb/tests/first_class_function_slice1.rs index dc565d309..aec264650 100644 --- a/rust/lancedb/tests/first_class_function_slice1.rs +++ b/rust/lancedb/tests/first_class_function_slice1.rs @@ -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); diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json index 7d92f1454..fd3944412 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_fixed_size_declaration_request.json @@ -24,8 +24,7 @@ "kind": "scalar", "arrow_type": "fixed_size_list", "nullable": false - }, - "group_id": "fg_fixed" + } }, "binding_metadata_version": 1, "input_bindings": [ @@ -76,4 +75,4 @@ } ] } -} \ No newline at end of file +} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json index 05da9fe39..b91dd1061 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.canonical.json @@ -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"}} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json index 44aeff460..177821aaa 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application.json @@ -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" diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json index 47724eee0..c23c8f3ca 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_application_float.json @@ -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} } diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json index 7bf93b8a8..143190f78 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.canonical.json @@ -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"}]} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json index 1a2053e42..dec3fa3f8 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_binding.json @@ -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"} } diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json similarity index 97% rename from rust/lancedb/tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json rename to rust/lancedb/tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json index d0b42cc99..c4ef5009f 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_grouped_declaration_request.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_multi_output_declaration_request.json @@ -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, diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json index bb2490bc8..c2d49827d 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_refresh_job.json @@ -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, diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json index 0aaa0cf72..357834de0 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_scalar_declaration_request.json @@ -8,8 +8,7 @@ "inputs": [ {"parameter": "text", "kind": "column", "value": {"path": "description"}} ], - "output": {"kind": "scalar", "arrow_type": "list", "nullable": false}, - "group_id": "fg_scalar" + "output": {"kind": "scalar", "arrow_type": "list", "nullable": false} }, "binding_metadata_version": 1, "input_bindings": [ From 81c3f108ce35f148123b831c399c335c0c79b46a Mon Sep 17 00:00:00 2001 From: Lance Release Date: Tue, 25 Aug 2026 10:31:36 +0000 Subject: [PATCH 12/23] =?UTF-8?q?Bump=20version:=200.38.0-beta.8=20?= =?UTF-8?q?=E2=86=92=200.38.0-beta.9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .bumpversion.toml | 2 +- Cargo.lock | 6 +++--- docs/src/java/java.md | 2 +- java/lancedb-core/pom.xml | 2 +- java/pom.xml | 2 +- nodejs/Cargo.toml | 2 +- nodejs/npm/darwin-arm64/package.json | 2 +- nodejs/npm/linux-arm64-gnu/package.json | 2 +- nodejs/npm/linux-arm64-musl/package.json | 2 +- nodejs/npm/linux-x64-gnu/package.json | 2 +- nodejs/npm/linux-x64-musl/package.json | 2 +- nodejs/npm/win32-arm64-msvc/package.json | 2 +- nodejs/npm/win32-x64-msvc/package.json | 2 +- nodejs/package-lock.json | 4 ++-- nodejs/package.json | 2 +- python/Cargo.toml | 2 +- rust/lancedb/Cargo.toml | 2 +- 17 files changed, 20 insertions(+), 20 deletions(-) diff --git a/.bumpversion.toml b/.bumpversion.toml index 86a5064f7..fb1323d7e 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.38.0-beta.8" +current_version = "0.38.0-beta.9" parse = """(?x) (?P0|[1-9]\\d*)\\. (?P0|[1-9]\\d*)\\. diff --git a/Cargo.lock b/Cargo.lock index 8d0058569..80db934d2 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5402,7 +5402,7 @@ dependencies = [ [[package]] name = "lancedb" -version = "0.38.0-beta.8" +version = "0.38.0-beta.9" dependencies = [ "ahash", "anyhow", @@ -5490,7 +5490,7 @@ dependencies = [ [[package]] name = "lancedb-nodejs" -version = "0.38.0-beta.8" +version = "0.38.0-beta.9" dependencies = [ "arrow-array", "arrow-buffer", @@ -5515,7 +5515,7 @@ dependencies = [ [[package]] name = "lancedb-python" -version = "0.38.0-beta.8" +version = "0.38.0-beta.9" dependencies = [ "arrow", "async-trait", diff --git a/docs/src/java/java.md b/docs/src/java/java.md index 28de11c30..7aa12d0a0 100644 --- a/docs/src/java/java.md +++ b/docs/src/java/java.md @@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`: com.lancedb lancedb-core - 0.38.0-beta.8 + 0.38.0-beta.9 ``` diff --git a/java/lancedb-core/pom.xml b/java/lancedb-core/pom.xml index 6f4b0744b..58d75d214 100644 --- a/java/lancedb-core/pom.xml +++ b/java/lancedb-core/pom.xml @@ -8,7 +8,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.8 + 0.38.0-beta.9 ../pom.xml diff --git a/java/pom.xml b/java/pom.xml index ac6196a45..69c8f5ce1 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -6,7 +6,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.8 + 0.38.0-beta.9 pom ${project.artifactId} LanceDB Java SDK Parent POM diff --git a/nodejs/Cargo.toml b/nodejs/Cargo.toml index e85da66ec..99decbd9d 100644 --- a/nodejs/Cargo.toml +++ b/nodejs/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "lancedb-nodejs" edition.workspace = true -version = "0.38.0-beta.8" +version = "0.38.0-beta.9" publish = false license.workspace = true description.workspace = true diff --git a/nodejs/npm/darwin-arm64/package.json b/nodejs/npm/darwin-arm64/package.json index f920c9e91..911d6495f 100644 --- a/nodejs/npm/darwin-arm64/package.json +++ b/nodejs/npm/darwin-arm64/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-darwin-arm64", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "os": ["darwin"], "cpu": ["arm64"], "main": "lancedb.darwin-arm64.node", diff --git a/nodejs/npm/linux-arm64-gnu/package.json b/nodejs/npm/linux-arm64-gnu/package.json index ad6926039..078b32dd1 100644 --- a/nodejs/npm/linux-arm64-gnu/package.json +++ b/nodejs/npm/linux-arm64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-gnu", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-gnu.node", diff --git a/nodejs/npm/linux-arm64-musl/package.json b/nodejs/npm/linux-arm64-musl/package.json index ccd39f18e..0fc54cdee 100644 --- a/nodejs/npm/linux-arm64-musl/package.json +++ b/nodejs/npm/linux-arm64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-musl", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-musl.node", diff --git a/nodejs/npm/linux-x64-gnu/package.json b/nodejs/npm/linux-x64-gnu/package.json index 50446b597..ecaf22680 100644 --- a/nodejs/npm/linux-x64-gnu/package.json +++ b/nodejs/npm/linux-x64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-gnu", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-gnu.node", diff --git a/nodejs/npm/linux-x64-musl/package.json b/nodejs/npm/linux-x64-musl/package.json index d4da82dc7..081c4c8ba 100644 --- a/nodejs/npm/linux-x64-musl/package.json +++ b/nodejs/npm/linux-x64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-musl", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-musl.node", diff --git a/nodejs/npm/win32-arm64-msvc/package.json b/nodejs/npm/win32-arm64-msvc/package.json index dddc3a470..609f701b7 100644 --- a/nodejs/npm/win32-arm64-msvc/package.json +++ b/nodejs/npm/win32-arm64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-arm64-msvc", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "os": [ "win32" ], diff --git a/nodejs/npm/win32-x64-msvc/package.json b/nodejs/npm/win32-x64-msvc/package.json index 23ac122f7..59eb8e06e 100644 --- a/nodejs/npm/win32-x64-msvc/package.json +++ b/nodejs/npm/win32-x64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-x64-msvc", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "os": ["win32"], "cpu": ["x64"], "main": "lancedb.win32-x64-msvc.node", diff --git a/nodejs/package-lock.json b/nodejs/package-lock.json index 21deb8dd8..d8b38ce0f 100644 --- a/nodejs/package-lock.json +++ b/nodejs/package-lock.json @@ -1,12 +1,12 @@ { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "cpu": [ "x64", "arm64" diff --git a/nodejs/package.json b/nodejs/package.json index 96f4bafd5..2c0ff7d69 100644 --- a/nodejs/package.json +++ b/nodejs/package.json @@ -11,7 +11,7 @@ "ann" ], "private": false, - "version": "0.38.0-beta.8", + "version": "0.38.0-beta.9", "main": "dist/index.js", "exports": { ".": "./dist/index.js", diff --git a/python/Cargo.toml b/python/Cargo.toml index bffac2ef1..8c73b602d 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb-python" -version = "0.38.0-beta.8" +version = "0.38.0-beta.9" publish = false edition.workspace = true description = "Python bindings for LanceDB" diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index 7290694f1..6f194556a 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb" -version = "0.38.0-beta.8" +version = "0.38.0-beta.9" edition.workspace = true description = "LanceDB: A serverless, low-latency vector database for AI applications" license.workspace = true From d0bcc6c6fe71b057f53be07a75be4bdefd463fb8 Mon Sep 17 00:00:00 2001 From: Xuanwo Date: Tue, 25 Aug 2026 18:35:27 +0800 Subject: [PATCH 13/23] fix: remove unsupported Function secrets contract (#4047) ## Problem The unreleased First-Class Function authoring API exposed `secrets=[...]` and serialized `required_secrets`, promising runtime resolution and injection that Sophon does not implement. ## Behavior Remove the secrets dimension from the public Python decorator, Python and Rust registration/version models, and shared wire fixtures. Existing successfully registered Functions retain stable identity: the field could only be empty and empty values were already omitted from canonical serialization. A stable LanceDB release has not shipped this Function API, so this contracts the surface before it becomes a published compatibility commitment. ## Validation gap The Rust shared-golden Function suites and Python formatting/lint checks pass locally. Python pytest was not run because the local environment lacks its runtime dependencies and the frozen native editable build did not complete in practical time. --- python/python/lancedb/functions.py | 37 +++--------------- .../tests/test_first_class_function_slice1.py | 25 ------------ .../tests/test_first_class_function_slice2.py | 29 -------------- rust/lancedb/src/function.rs | 20 +--------- .../tests/first_class_function_slice1.rs | 38 ------------------- .../tests/first_class_function_slice2.rs | 26 ------------- .../v1/remote_function_job.json | 1 - ...nction_registration_request.canonical.json | 2 +- .../remote_function_registration_request.json | 5 +-- .../v1/remote_function_version.canonical.json | 2 +- 10 files changed, 10 insertions(+), 175 deletions(-) diff --git a/python/python/lancedb/functions.py b/python/python/lancedb/functions.py index dd8cecfa9..237ed73c2 100644 --- a/python/python/lancedb/functions.py +++ b/python/python/lancedb/functions.py @@ -4,7 +4,7 @@ """Canonical Function values exchanged with LanceDB Enterprise services. These immutable models contain client/wire state only. Catalog persistence, -environment bake, secret resolution, and execution are owned by Sophon. +environment bake, and execution are owned by Sophon. ``RefreshColumnResult`` is also the backend-neutral result of a local expression-backed refresh job. """ @@ -228,7 +228,7 @@ class PythonEnvironmentSpec(_RemoteValue): class PythonRuntimeSpec(_RemoteValue): - """Remote runtime definition with non-secret environment values. + """Remote runtime definition with environment values. V1 supports ``kind="python"``. Newer runtime kinds remain readable, while their unknown payload fields are intentionally not retained by the client. @@ -267,7 +267,6 @@ class FunctionVersion(_RemoteValue): runtime: PythonRuntimeSpec runtime_digest: str environment_digest: str - required_secrets: tuple[str, ...] = () created_at: str def __call__(self, **inputs: Any) -> FunctionApplication: @@ -329,17 +328,12 @@ class FunctionVersion(_RemoteValue): class FunctionRegistrationRequest(_RemoteValue): - """Stable remote registration envelope produced by :func:`udf`. - - Only secret names are represented. Secret values are resolved inside the - remote service and have no client request field. - """ + """Stable remote registration envelope produced by :func:`udf`.""" name: str artifact: FunctionArtifactRequest signature: FunctionSignature runtime: PythonRuntimeSpec - required_secrets: tuple[str, ...] = () class FunctionVersionRef(_OpenRemoteValue): @@ -484,7 +478,6 @@ class RefreshColumnResult(_RemoteValue): _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 = ( @@ -915,7 +908,6 @@ class UdfDefinition: output_schema: Optional[pa.DataType | pa.Field | pa.Schema], pip: tuple[str, ...], env: Mapping[str, str], - secrets: tuple[str, ...], python_version: Optional[str], ): function_name = name or function.__name__ @@ -930,18 +922,6 @@ class UdfDefinition: for key, value in environment.items() ): raise TypeError("Function env keys and values must be strings") - required_secrets = tuple(sorted(set(secrets))) - invalid_secrets = [ - secret for secret in required_secrets if not _SECRET_NAME.fullmatch(secret) - ] - if invalid_secrets: - raise ValueError(f"invalid Function secret names: {invalid_secrets!r}") - overlap = set(environment) & set(required_secrets) - if overlap: - raise ValueError( - f"Function env and secret names must be disjoint: {sorted(overlap)!r}" - ) - signature = _infer_signature(function, input_schema, output_schema) source = _package_source(function) digest = f"sha256:{hashlib.sha256(source).hexdigest()}" @@ -970,7 +950,6 @@ class UdfDefinition: ), signature=signature, runtime=runtime, - required_secrets=required_secrets, ) functools.update_wrapper(self, function) @@ -996,7 +975,6 @@ def udf( output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None, pip: tuple[str, ...] | list[str] = (), env: Optional[Mapping[str, str]] = None, - secrets: tuple[str, ...] | list[str] = (), python_version: Optional[str] = None, ) -> Callable[[Callable[..., Any]], UdfDefinition]: ... @@ -1009,7 +987,6 @@ def udf( output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None, pip: tuple[str, ...] | list[str] = (), env: Optional[Mapping[str, str]] = None, - secrets: tuple[str, ...] | list[str] = (), python_version: Optional[str] = None, ): """Prepare a scalar Python callable for remote Function registration. @@ -1034,10 +1011,7 @@ def udf( pip : sequence of str, optional Pip requirements for the remote environment. env : mapping of str to str, optional - Non-secret environment variables. Use ``secrets`` for credentials. - secrets : sequence of str, optional - Names of secrets resolved by the remote service. Secret values are not - accepted by this API or included in the registration request. + Environment variables included in the Function definition. python_version : str, optional Remote Python major/minor version. Defaults to the client version. @@ -1059,7 +1033,7 @@ def udf( Examples -------- >>> from lancedb import udf - >>> @udf(pip=["numpy==2.2.0"], secrets=["MODEL_TOKEN"]) + >>> @udf(pip=["numpy==2.2.0"]) ... def score(value: float) -> float: ... return value * 2 >>> score(1.5) @@ -1074,7 +1048,6 @@ def udf( output_schema=output_schema, pip=tuple(pip), env={} if env is None else env, - secrets=tuple(secrets), python_version=python_version, ) diff --git a/python/python/tests/test_first_class_function_slice1.py b/python/python/tests/test_first_class_function_slice1.py index ca0b30ede..89172ba3f 100644 --- a/python/python/tests/test_first_class_function_slice1.py +++ b/python/python/tests/test_first_class_function_slice1.py @@ -37,21 +37,6 @@ def job_result(name: str) -> dict: return json.loads(fixture(name))["result"] -def assert_no_secret_values(value): - if isinstance(value, dict): - for key, child in value.items(): - assert key not in { - "secret_value", - "secret_values", - "resolved_secret", - "resolved_secrets", - } - assert_no_secret_values(child) - elif isinstance(value, list): - for child in value: - assert_no_secret_values(child) - - def test_public_function_values_are_in_api_reference(): docs = Path(__file__).parents[3] / "docs" / "src" / "python" / "python.md" rendered = docs.read_text() @@ -109,7 +94,6 @@ def test_function_version_identity_is_immutable_and_exact(): version = FunctionVersion.from_json(json.dumps(value)) assert version.name == "embed" assert version.version == "fv_01K3EXACT" - assert version.required_secrets == ("HF_TOKEN",) with pytest.raises((TypeError, ValueError)): version.version = "fv_changed" @@ -292,15 +276,6 @@ def test_refresh_result_rejects_non_u64_values(field): RefreshColumnResult.from_json(json.dumps(value)) -def test_canonical_client_values_contain_secret_names_only(): - version = FunctionVersion.from_json( - json.dumps(job_result("remote_function_job.json")) - ) - canonical = json.loads(version.to_canonical_json()) - assert canonical["required_secrets"] == ["HF_TOKEN"] - assert_no_secret_values(canonical) - - class _FunctionDeclarationInner: def __init__(self): self.calls = [] diff --git a/python/python/tests/test_first_class_function_slice2.py b/python/python/tests/test_first_class_function_slice2.py index c67e17520..55257d322 100644 --- a/python/python/tests/test_first_class_function_slice2.py +++ b/python/python/tests/test_first_class_function_slice2.py @@ -39,28 +39,12 @@ FIXTURES = ( @udf( pip=["numpy>=2"], env={"MODE": "test"}, - secrets=["API_TOKEN"], python_version="3.12", ) def normalize_score(value: float) -> float: return value / 100.0 -def _assert_no_secret_values(value): - if isinstance(value, dict): - for key, child in value.items(): - assert key not in { - "secret_value", - "secret_values", - "resolved_secret", - "resolved_secrets", - } - _assert_no_secret_values(child) - elif isinstance(value, list): - for child in value: - _assert_no_secret_values(child) - - def test_scalar_udf_matches_shared_registration_golden_and_remains_callable(): assert isinstance(normalize_score, UdfDefinition) assert normalize_score(25.0) == 0.25 @@ -75,8 +59,6 @@ def test_scalar_udf_matches_shared_registration_golden_and_remains_callable(): "kind": "scalar_to_arrow_batch", "version": 1, } - assert request["required_secrets"] == ["API_TOKEN"] - _assert_no_secret_values(request) def _run_packaged(definition, *args): @@ -370,7 +352,6 @@ def test_udf_recursion_versus_a_rebound_module_name(tmp_path): output_schema=None, pip=(), env={}, - secrets=(), python_version=None, ) with pytest.raises(ValueError, match="binds that name to another value"): @@ -525,14 +506,6 @@ def test_annotation_and_explicit_schema_validation_fail_closed(): return value -def test_environment_rejects_secret_value_overlap(): - with pytest.raises(ValueError, match="must be disjoint"): - - @udf(env={"TOKEN": "plaintext"}, secrets=["TOKEN"]) - def overlapping(value: int) -> int: - return value - - def test_local_function_catalog_operations_are_not_supported(tmp_path): db = lancedb.connect(tmp_path) message = "Function catalog operations are not supported by this database" @@ -569,7 +542,6 @@ def _mock_remote_function_catalog(): "runtime": body["runtime"], "runtime_digest": "sha256:runtime", "environment_digest": "sha256:environment", - "required_secrets": body.get("required_secrets", []), "created_at": "2026-08-21T00:00:00Z", } response = {"job_id": "job-register"} @@ -628,7 +600,6 @@ def test_remote_registration_job_and_exact_version_reopen_round_trip(): assert create_request == json.loads( normalize_score.registration_request.to_canonical_json() ) - _assert_no_secret_values(create_request) def test_blocking_remote_registration_returns_function_version(): diff --git a/rust/lancedb/src/function.rs b/rust/lancedb/src/function.rs index 433a7e04c..52c70a4b1 100644 --- a/rust/lancedb/src/function.rs +++ b/rust/lancedb/src/function.rs @@ -5,7 +5,7 @@ //! backend-neutral terminal result of a computed-column refresh. //! //! This module contains client/wire values only. Catalog persistence, -//! environment bake, secret resolution, and execution are owned by Sophon. +//! environment bake, and execution are owned by Sophon. use std::collections::BTreeMap; @@ -195,9 +195,6 @@ pub struct PythonEnvironmentSpec { } /// Reproducible Python runtime definition understood by Sophon. -/// -/// `env` contains non-secret values. Secret values have no client model; -/// [`FunctionVersion::required_secrets`] contains names only. #[derive(Debug, Clone, PartialEq, Eq)] #[non_exhaustive] pub enum PythonRuntimeSpec { @@ -239,7 +236,7 @@ impl PythonRuntimeSpec { } } - /// Non-secret environment variables, or `None` for an unknown kind. + /// Environment variables, or `None` for an unknown kind. pub fn env(&self) -> Option<&BTreeMap> { match self { Self::Python { env, .. } => Some(env), @@ -324,8 +321,6 @@ pub struct FunctionVersion { runtime: PythonRuntimeSpec, runtime_digest: String, environment_digest: String, - #[serde(default, skip_serializing_if = "Vec::is_empty")] - required_secrets: Vec, created_at: String, } @@ -358,11 +353,6 @@ impl FunctionVersion { &self.environment_digest } - /// Required secret names. Resolved values exist only inside Sophon. - pub fn required_secrets(&self) -> &[String] { - &self.required_secrets - } - pub fn created_at(&self) -> &str { &self.created_at } @@ -404,18 +394,12 @@ pub struct FunctionArtifactRequest { } /// Stable request envelope for remote immutable Function registration. -/// -/// Secret values deliberately have no field in this model. The only secret -/// material the client may send is the ordered set of names Sophon resolves -/// inside the remote runtime. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct FunctionRegistrationRequest { pub name: String, pub artifact: FunctionArtifactRequest, pub signature: FunctionSignature, pub runtime: PythonRuntimeSpec, - #[serde(default, skip_serializing_if = "Vec::is_empty")] - pub required_secrets: Vec, } impl_json!(FunctionRegistrationRequest); diff --git a/rust/lancedb/tests/first_class_function_slice1.rs b/rust/lancedb/tests/first_class_function_slice1.rs index aec264650..ce020bd53 100644 --- a/rust/lancedb/tests/first_class_function_slice1.rs +++ b/rust/lancedb/tests/first_class_function_slice1.rs @@ -20,25 +20,6 @@ fn job_result(name: &str) -> Value { serde_json::from_str::(&fixture(name)).expect("remote Job fixture")["result"].clone() } -fn assert_no_secret_values(value: &Value) { - match value { - Value::Object(values) => { - for (key, value) in values { - assert!( - !matches!( - key.as_str(), - "secret_value" | "secret_values" | "resolved_secret" | "resolved_secrets" - ), - "client canonical value must not model resolved secret material" - ); - assert_no_secret_values(value); - } - } - Value::Array(values) => values.iter().for_each(assert_no_secret_values), - _ => {} - } -} - #[test] fn function_version_job_result_matches_shared_canonical_golden() { let result = job_result("remote_function_job.json"); @@ -47,7 +28,6 @@ fn function_version_job_result_matches_shared_canonical_golden() { assert_eq!(version.name(), "embed"); assert_eq!(version.version(), "fv_01K3EXACT"); assert_eq!(version.runtime_digest(), "sha256:runtime"); - assert_eq!(version.required_secrets(), &["HF_TOKEN"]); assert_eq!( version.to_canonical_json().expect("canonical JSON"), fixture("remote_function_version.canonical.json").trim() @@ -162,21 +142,3 @@ fn floating_point_application_literals_are_rejected_consistently() { .contains("floating-point Function literals") ); } - -#[test] -fn canonical_client_values_contain_secret_names_only() { - let result = job_result("remote_function_job.json"); - let version = FunctionVersion::from_json(&result.to_string()).expect("FunctionVersion result"); - let canonical: Value = serde_json::from_str( - &version - .to_canonical_json() - .expect("canonical FunctionVersion"), - ) - .expect("canonical JSON"); - - assert_eq!( - canonical["required_secrets"], - serde_json::json!(["HF_TOKEN"]) - ); - assert_no_secret_values(&canonical); -} diff --git a/rust/lancedb/tests/first_class_function_slice2.rs b/rust/lancedb/tests/first_class_function_slice2.rs index 3bae57122..93252dde4 100644 --- a/rust/lancedb/tests/first_class_function_slice2.rs +++ b/rust/lancedb/tests/first_class_function_slice2.rs @@ -6,7 +6,6 @@ use std::path::PathBuf; use lancedb::Error; use lancedb::function::FunctionRegistrationRequest; -use serde_json::Value; fn fixture(name: &str) -> String { let path = PathBuf::from(env!("CARGO_MANIFEST_DIR")) @@ -15,25 +14,6 @@ fn fixture(name: &str) -> String { fs::read_to_string(path).expect("fixture must be readable") } -fn assert_no_secret_values(value: &Value) { - match value { - Value::Object(values) => { - for (key, value) in values { - assert!( - !matches!( - key.as_str(), - "secret_value" | "secret_values" | "resolved_secret" | "resolved_secrets" - ), - "registration requests must not model resolved secret material" - ); - assert_no_secret_values(value); - } - } - Value::Array(values) => values.iter().for_each(assert_no_secret_values), - _ => {} - } -} - #[test] fn registration_request_matches_shared_canonical_golden() { let request = FunctionRegistrationRequest::from_json(&fixture( @@ -42,16 +22,10 @@ fn registration_request_matches_shared_canonical_golden() { .expect("registration request"); assert_eq!(request.name, "normalize_score"); assert_eq!(request.artifact.adapter.kind, "scalar_to_arrow_batch"); - assert_eq!(request.required_secrets, ["API_TOKEN"]); assert_eq!( request.to_canonical_json().expect("canonical request"), fixture("remote_function_registration_request.canonical.json").trim() ); - - let value: Value = - serde_json::from_str(&request.to_canonical_json().expect("canonical request")) - .expect("request JSON"); - assert_no_secret_values(&value); } #[tokio::test] diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json index 6ba4eb226..39a279692 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_job.json @@ -24,7 +24,6 @@ }, "runtime_digest": "sha256:runtime", "environment_digest": "sha256:environment", - "required_secrets": ["HF_TOKEN"], "created_at": "2026-08-21T00:00:00Z" }, "future_job": {"trace_id": "trace-1"} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json index 24fa2cf30..a2f2c4c21 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.canonical.json @@ -1 +1 @@ -{"artifact":{"adapter":{"kind":"scalar_to_arrow_batch","version":1},"content":{"data":"ZnJvbSBfX2Z1dHVyZV9fIGltcG9ydCBhbm5vdGF0aW9ucwoKZGVmIG5vcm1hbGl6ZV9zY29yZSh2YWx1ZTogZmxvYXQpIC0+IGZsb2F0OgogICAgcmV0dXJuIHZhbHVlIC8gMTAwLjAK","encoding":"base64"},"digest":"sha256:760784bdcef57b802f389b97804cc0b618aae39e86044733451bcd0b13089a7f","entrypoint":"normalize_score","kind":"python_callable"},"name":"normalize_score","required_secrets":["API_TOKEN"],"runtime":{"env":{"MODE":"test"},"environment":{"kind":"pip","packages":["numpy>=2"]},"kind":"python","python_version":"3.12"},"signature":{"inputs":[{"arrow_type":"float64","name":"value","nullable":false}],"output":{"arrow_type":"float64","kind":"scalar","nullable":false}}} +{"artifact":{"adapter":{"kind":"scalar_to_arrow_batch","version":1},"content":{"data":"ZnJvbSBfX2Z1dHVyZV9fIGltcG9ydCBhbm5vdGF0aW9ucwoKZGVmIG5vcm1hbGl6ZV9zY29yZSh2YWx1ZTogZmxvYXQpIC0+IGZsb2F0OgogICAgcmV0dXJuIHZhbHVlIC8gMTAwLjAK","encoding":"base64"},"digest":"sha256:760784bdcef57b802f389b97804cc0b618aae39e86044733451bcd0b13089a7f","entrypoint":"normalize_score","kind":"python_callable"},"name":"normalize_score","runtime":{"env":{"MODE":"test"},"environment":{"kind":"pip","packages":["numpy>=2"]},"kind":"python","python_version":"3.12"},"signature":{"inputs":[{"arrow_type":"float64","name":"value","nullable":false}],"output":{"arrow_type":"float64","kind":"scalar","nullable":false}}} diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json index bbfec3169..092d76dc2 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_registration_request.json @@ -39,8 +39,5 @@ "env": { "MODE": "test" } - }, - "required_secrets": [ - "API_TOKEN" - ] + } } diff --git a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json index 7ab632a98..2670ad0b2 100644 --- a/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json +++ b/rust/lancedb/tests/fixtures/first_class_functions/v1/remote_function_version.canonical.json @@ -1 +1 @@ -{"artifact":{"digest":"sha256:code","entrypoint":"embed","kind":"python_callable"},"created_at":"2026-08-21T00:00:00Z","environment_digest":"sha256:environment","name":"embed","required_secrets":["HF_TOKEN"],"runtime":{"env":{"TOKENIZERS_PARALLELISM":"false"},"environment":{"kind":"pip","packages":["sentence-transformers>=3"]},"kind":"python","python_version":"3.12"},"runtime_digest":"sha256:runtime","signature":{"inputs":[{"arrow_type":"utf8","name":"text","nullable":true}],"output":{"arrow_type":"list","kind":"scalar","nullable":false}},"version":"fv_01K3EXACT"} +{"artifact":{"digest":"sha256:code","entrypoint":"embed","kind":"python_callable"},"created_at":"2026-08-21T00:00:00Z","environment_digest":"sha256:environment","name":"embed","runtime":{"env":{"TOKENIZERS_PARALLELISM":"false"},"environment":{"kind":"pip","packages":["sentence-transformers>=3"]},"kind":"python","python_version":"3.12"},"runtime_digest":"sha256:runtime","signature":{"inputs":[{"arrow_type":"utf8","name":"text","nullable":true}],"output":{"arrow_type":"list","kind":"scalar","nullable":false}},"version":"fv_01K3EXACT"} From ec4ad54ba253e11c4cfd894cde6d47e788447886 Mon Sep 17 00:00:00 2001 From: Lance Release Date: Tue, 25 Aug 2026 10:36:33 +0000 Subject: [PATCH 14/23] =?UTF-8?q?Bump=20version:=200.38.0-beta.9=20?= =?UTF-8?q?=E2=86=92=200.38.0-beta.10?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .bumpversion.toml | 2 +- Cargo.lock | 6 +++--- docs/src/java/java.md | 2 +- java/lancedb-core/pom.xml | 2 +- java/pom.xml | 2 +- nodejs/Cargo.toml | 2 +- nodejs/npm/darwin-arm64/package.json | 2 +- nodejs/npm/linux-arm64-gnu/package.json | 2 +- nodejs/npm/linux-arm64-musl/package.json | 2 +- nodejs/npm/linux-x64-gnu/package.json | 2 +- nodejs/npm/linux-x64-musl/package.json | 2 +- nodejs/npm/win32-arm64-msvc/package.json | 2 +- nodejs/npm/win32-x64-msvc/package.json | 2 +- nodejs/package-lock.json | 4 ++-- nodejs/package.json | 2 +- python/Cargo.toml | 2 +- rust/lancedb/Cargo.toml | 2 +- 17 files changed, 20 insertions(+), 20 deletions(-) diff --git a/.bumpversion.toml b/.bumpversion.toml index fb1323d7e..1c4aea809 100644 --- a/.bumpversion.toml +++ b/.bumpversion.toml @@ -1,5 +1,5 @@ [tool.bumpversion] -current_version = "0.38.0-beta.9" +current_version = "0.38.0-beta.10" parse = """(?x) (?P0|[1-9]\\d*)\\. (?P0|[1-9]\\d*)\\. diff --git a/Cargo.lock b/Cargo.lock index 80db934d2..33e43f9fa 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5402,7 +5402,7 @@ dependencies = [ [[package]] name = "lancedb" -version = "0.38.0-beta.9" +version = "0.38.0-beta.10" dependencies = [ "ahash", "anyhow", @@ -5490,7 +5490,7 @@ dependencies = [ [[package]] name = "lancedb-nodejs" -version = "0.38.0-beta.9" +version = "0.38.0-beta.10" dependencies = [ "arrow-array", "arrow-buffer", @@ -5515,7 +5515,7 @@ dependencies = [ [[package]] name = "lancedb-python" -version = "0.38.0-beta.9" +version = "0.38.0-beta.10" dependencies = [ "arrow", "async-trait", diff --git a/docs/src/java/java.md b/docs/src/java/java.md index 7aa12d0a0..06dc267f3 100644 --- a/docs/src/java/java.md +++ b/docs/src/java/java.md @@ -14,7 +14,7 @@ Add the following dependency to your `pom.xml`: com.lancedb lancedb-core - 0.38.0-beta.9 + 0.38.0-beta.10 ``` diff --git a/java/lancedb-core/pom.xml b/java/lancedb-core/pom.xml index 58d75d214..25e3b10e3 100644 --- a/java/lancedb-core/pom.xml +++ b/java/lancedb-core/pom.xml @@ -8,7 +8,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.9 + 0.38.0-beta.10 ../pom.xml diff --git a/java/pom.xml b/java/pom.xml index 69c8f5ce1..4d83153ab 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -6,7 +6,7 @@ com.lancedb lancedb-parent - 0.38.0-beta.9 + 0.38.0-beta.10 pom ${project.artifactId} LanceDB Java SDK Parent POM diff --git a/nodejs/Cargo.toml b/nodejs/Cargo.toml index 99decbd9d..fd08e7a5e 100644 --- a/nodejs/Cargo.toml +++ b/nodejs/Cargo.toml @@ -1,7 +1,7 @@ [package] name = "lancedb-nodejs" edition.workspace = true -version = "0.38.0-beta.9" +version = "0.38.0-beta.10" publish = false license.workspace = true description.workspace = true diff --git a/nodejs/npm/darwin-arm64/package.json b/nodejs/npm/darwin-arm64/package.json index 911d6495f..ff6347c4d 100644 --- a/nodejs/npm/darwin-arm64/package.json +++ b/nodejs/npm/darwin-arm64/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-darwin-arm64", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "os": ["darwin"], "cpu": ["arm64"], "main": "lancedb.darwin-arm64.node", diff --git a/nodejs/npm/linux-arm64-gnu/package.json b/nodejs/npm/linux-arm64-gnu/package.json index 078b32dd1..ed99fee05 100644 --- a/nodejs/npm/linux-arm64-gnu/package.json +++ b/nodejs/npm/linux-arm64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-gnu", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-gnu.node", diff --git a/nodejs/npm/linux-arm64-musl/package.json b/nodejs/npm/linux-arm64-musl/package.json index 0fc54cdee..5b215dcc0 100644 --- a/nodejs/npm/linux-arm64-musl/package.json +++ b/nodejs/npm/linux-arm64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-arm64-musl", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["arm64"], "main": "lancedb.linux-arm64-musl.node", diff --git a/nodejs/npm/linux-x64-gnu/package.json b/nodejs/npm/linux-x64-gnu/package.json index ecaf22680..e0f5a9f26 100644 --- a/nodejs/npm/linux-x64-gnu/package.json +++ b/nodejs/npm/linux-x64-gnu/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-gnu", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-gnu.node", diff --git a/nodejs/npm/linux-x64-musl/package.json b/nodejs/npm/linux-x64-musl/package.json index 081c4c8ba..d42541707 100644 --- a/nodejs/npm/linux-x64-musl/package.json +++ b/nodejs/npm/linux-x64-musl/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-linux-x64-musl", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "os": ["linux"], "cpu": ["x64"], "main": "lancedb.linux-x64-musl.node", diff --git a/nodejs/npm/win32-arm64-msvc/package.json b/nodejs/npm/win32-arm64-msvc/package.json index 609f701b7..496a40720 100644 --- a/nodejs/npm/win32-arm64-msvc/package.json +++ b/nodejs/npm/win32-arm64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-arm64-msvc", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "os": [ "win32" ], diff --git a/nodejs/npm/win32-x64-msvc/package.json b/nodejs/npm/win32-x64-msvc/package.json index 59eb8e06e..734013343 100644 --- a/nodejs/npm/win32-x64-msvc/package.json +++ b/nodejs/npm/win32-x64-msvc/package.json @@ -1,6 +1,6 @@ { "name": "@lancedb/lancedb-win32-x64-msvc", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "os": ["win32"], "cpu": ["x64"], "main": "lancedb.win32-x64-msvc.node", diff --git a/nodejs/package-lock.json b/nodejs/package-lock.json index d8b38ce0f..732a7c01c 100644 --- a/nodejs/package-lock.json +++ b/nodejs/package-lock.json @@ -1,12 +1,12 @@ { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "@lancedb/lancedb", - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "cpu": [ "x64", "arm64" diff --git a/nodejs/package.json b/nodejs/package.json index 2c0ff7d69..262d4c4a7 100644 --- a/nodejs/package.json +++ b/nodejs/package.json @@ -11,7 +11,7 @@ "ann" ], "private": false, - "version": "0.38.0-beta.9", + "version": "0.38.0-beta.10", "main": "dist/index.js", "exports": { ".": "./dist/index.js", diff --git a/python/Cargo.toml b/python/Cargo.toml index 8c73b602d..3a8a05522 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb-python" -version = "0.38.0-beta.9" +version = "0.38.0-beta.10" publish = false edition.workspace = true description = "Python bindings for LanceDB" diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index 6f194556a..8276e5bb3 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "lancedb" -version = "0.38.0-beta.9" +version = "0.38.0-beta.10" edition.workspace = true description = "LanceDB: A serverless, low-latency vector database for AI applications" license.workspace = true From 1d880f11ff12edbf6b5c26505b6914bd1817e9e0 Mon Sep 17 00:00:00 2001 From: Xuanwo Date: Tue, 25 Aug 2026 19:51:40 +0800 Subject: [PATCH 15/23] fix(python): use dev profile for editable builds (#4049) ## Problem `uv run ... maturin develop` synchronizes the project as an editable package before running the command. Maturin's editable build otherwise uses the release profile, which enables the repository's fat LTO configuration during local bootstrap. ## Behavior Editable Python builds now explicitly use Cargo's dev profile. The minimum maturin version is raised to 1.10, where `editable-profile` support was introduced. --- python/pyproject.toml | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/python/pyproject.toml b/python/pyproject.toml index fad3b1001..22a41a8a9 100644 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -101,9 +101,12 @@ azure = ["adlfs>=2024.2.0"] [tool.maturin] python-source = "python" module-name = "lancedb._lancedb" +# uv installs the project as an editable package before `uv run`, so keep that +# bootstrap build consistent with `maturin develop`. +editable-profile = "dev" [build-system] -requires = ["maturin>=1.9.4"] +requires = ["maturin>=1.10"] build-backend = "maturin" [tool.ruff.lint] From a614400755b00c472eba6cf743802e4f0c734132 Mon Sep 17 00:00:00 2001 From: Drew Date: Tue, 25 Aug 2026 07:59:53 -0700 Subject: [PATCH 16/23] feat: accept blob URI writes (#3954) #3528 added blob declarations and binary coercion. String values were still rejected. They now coerce to the blob `uri` child. ```python table.add([{"id": 1, "image": "s3://bucket/media/cat.jpg"}]) payload = table.fetch_blobs("image", table.search().to_arrow()) ``` A URI under a registered base writes with no extra options. An unregistered URI fails. `allow_external_blob_outside_bases` is a local escape hatch that stores an absolute URI. It does not register a base. Remote `add` rejects that flag before making a request. String input still coerces and is sent as a `uri` struct. `add_bases` is a follow-up. `merge_insert` does not coerce string blob input. ### Testing - `cargo test -p lancedb --test blob_integration` - `cargo test -p lancedb blob_coerce` - `cargo test -p lancedb --features remote --lib add_rejects_external_blob_flag add_string_blob_becomes_uri_struct` - `cd python && uv run --extra tests pytest python/tests/test_blob.py -k uri -q` Co-authored-by: Xuanwo --- python/python/lancedb/_lancedb.pyi | 1 + python/python/lancedb/remote/table.py | 4 + python/python/lancedb/table.py | 15 ++ python/python/tests/test_blob.py | 68 ++++++ python/src/table.rs | 8 +- rust/lancedb/src/remote/table.rs | 96 +++++++- rust/lancedb/src/table.rs | 5 +- rust/lancedb/src/table/add_data.rs | 14 ++ .../src/table/datafusion/blob_coerce.rs | 153 +++++++++--- rust/lancedb/tests/blob_integration.rs | 230 +++++++++++++++++- 10 files changed, 552 insertions(+), 42 deletions(-) diff --git a/python/python/lancedb/_lancedb.pyi b/python/python/lancedb/_lancedb.pyi index 1e314ede8..593bceffa 100644 --- a/python/python/lancedb/_lancedb.pyi +++ b/python/python/lancedb/_lancedb.pyi @@ -283,6 +283,7 @@ class Table: mode: Literal["append", "overwrite"], progress: Optional[Any] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: ... async def update( self, updates: Dict[str, str], where: Optional[str] diff --git a/python/python/lancedb/remote/table.py b/python/python/lancedb/remote/table.py index 41394ed71..d0bf9f67a 100644 --- a/python/python/lancedb/remote/table.py +++ b/python/python/lancedb/remote/table.py @@ -610,6 +610,7 @@ class RemoteTable(Table): fill_value: float = 0.0, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add more data to the [Table][lancedb.table.Table]. @@ -642,6 +643,8 @@ class RemoteTable(Table): data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Not supported on LanceDB Cloud. Setting this raises. Returns ------- @@ -658,6 +661,7 @@ class RemoteTable(Table): fill_value=fill_value, progress=progress, write_parallelism=write_parallelism, + allow_external_blob_outside_bases=allow_external_blob_outside_bases, ) ) finally: diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index fe25d5353..b3cab006e 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -1269,6 +1269,7 @@ class Table(ABC): fill_value: float = 0.0, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add more data to the [Table][lancedb.table.Table]. @@ -1320,6 +1321,10 @@ class Table(ABC): data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Store blob URIs that sit outside registered blob bases. The row + keeps a reference, so the object has to stay readable. Local + tables only. Returns ------- @@ -3409,6 +3414,7 @@ class LanceTable(Table): fill_value: float = 0.0, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add data to the table. If vector columns are missing and the table @@ -3436,6 +3442,9 @@ class LanceTable(Table): data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Allow blob URIs outside registered bases. See :meth:`Table.add`. + Local tables only. Returns ------- @@ -3452,6 +3461,7 @@ class LanceTable(Table): fill_value=fill_value, progress=progress, write_parallelism=write_parallelism, + allow_external_blob_outside_bases=allow_external_blob_outside_bases, ) ) finally: @@ -5365,6 +5375,7 @@ class AsyncTable: fill_value: Optional[float] = None, progress: Optional[Union[bool, Callable, Any]] = None, write_parallelism: Optional[int] = None, + allow_external_blob_outside_bases: bool = False, ) -> AddResult: """Add more data to the [AsyncTable][lancedb.table.AsyncTable]. @@ -5395,6 +5406,9 @@ class AsyncTable: data in flight. Defaults to an estimate based on the data size, capped at the number of CPU cores. Lower this if bulk ingestion is using too much memory. + allow_external_blob_outside_bases: bool, default False + Allow blob URIs outside registered bases. See :meth:`Table.add`. + Local tables only. """ schema = await self.schema() @@ -5431,6 +5445,7 @@ class AsyncTable: mode or "append", progress=progress, write_parallelism=write_parallelism, + allow_external_blob_outside_bases=allow_external_blob_outside_bases, ) except RuntimeError as e: if "Cast error" in str(e): diff --git a/python/python/tests/test_blob.py b/python/python/tests/test_blob.py index b205c20a6..5d7682f24 100644 --- a/python/python/tests/test_blob.py +++ b/python/python/tests/test_blob.py @@ -617,3 +617,71 @@ def test_fetch_blobs_nested_path_survives_sort_after_query(): def _identifiable_payload(size: int) -> bytes: block = 256 return b"".join(bytes([i % 256]) * block for i in range(size // block)) + + +def _external_uri_blob_array(uris): + blob_type = lancedb.blob("image").type + storage_type = blob_type.storage_type + child_names = [field.name for field in storage_type] + assert "uri" in child_names, "blob layout no longer has a uri child" + children = [ + pa.array(uris if field.name == "uri" else [None] * len(uris), type=field.type) + for field in storage_type + ] + storage = pa.StructArray.from_arrays(children, fields=list(storage_type)) + return pa.ExtensionArray.from_storage(blob_type, storage) + + +def _external_uri_table_and_rows(name, uris): + db = lancedb.connect("memory:///") + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table(name, schema=schema) + rows = pa.Table.from_arrays( + [ + pa.array(range(len(uris)), type=pa.int64()), + _external_uri_blob_array(uris), + ], + schema=schema, + ) + return table, rows + + +def test_add_external_uri_struct_round_trips_with_flag(tmp_path): + payload = b"external-uri-bytes" + blob_path = tmp_path / "payload.bin" + blob_path.write_bytes(payload) + + table, rows = _external_uri_table_and_rows("external_struct", [blob_path.as_uri()]) + table.add(rows, allow_external_blob_outside_bases=True) + + hits = table.search().to_arrow() + blobs = table.fetch_blobs("image", hits) + assert blobs[0].as_py() == payload + + +def test_add_external_uri_without_flag_raises(tmp_path): + blob_path = tmp_path / "payload.bin" + blob_path.write_bytes(b"unreachable") + + table, rows = _external_uri_table_and_rows("external_no_flag", [blob_path.as_uri()]) + with pytest.raises(ValueError, match="allow_external_blob_outside_bases"): + table.add(rows) + assert table.count_rows() == 0 + + +def test_add_external_uri_string_round_trips_with_flag(tmp_path): + payload = b"external-uri-bytes" + blob_path = tmp_path / "payload.bin" + blob_path.write_bytes(payload) + + db = lancedb.connect("memory:///") + schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")]) + table = db.create_table("external_string", schema=schema) + table.add( + [{"id": 1, "image": blob_path.as_uri()}], + allow_external_blob_outside_bases=True, + ) + + hits = table.search().to_arrow() + blobs = table.fetch_blobs("image", hits) + assert blobs[0].as_py() == payload diff --git a/python/src/table.rs b/python/src/table.rs index b225b191f..784d29136 100644 --- a/python/src/table.rs +++ b/python/src/table.rs @@ -780,15 +780,19 @@ impl Table { }) } - #[pyo3(signature = (data, mode, progress=None, write_parallelism=None))] + #[pyo3(signature = (data, mode, progress=None, write_parallelism=None, allow_external_blob_outside_bases=false))] pub fn add<'a>( self_: PyRef<'a, Self>, data: PyScannable, mode: String, progress: Option>, write_parallelism: Option, + allow_external_blob_outside_bases: bool, ) -> PyResult> { - let mut op = self_.inner_ref()?.add(data); + let mut op = self_ + .inner_ref()? + .add(data) + .allow_external_blob_outside_bases(allow_external_blob_outside_bases); if mode == "append" { op = op.mode(AddDataMode::Append); } else if mode == "overwrite" { diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 742ab547d..44c92310f 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -2161,6 +2161,15 @@ impl BaseTable for RemoteTable { async fn add(&self, mut add: AddDataBuilder) -> Result { self.check_mutable().await?; + if add.allow_external_blob_outside_bases { + return Err(Error::NotSupported { + message: "allow_external_blob_outside_bases is only supported on local tables" + .to_string(), + }); + } + // String blob values still coerce to the uri child in into_plan. + // Remote and local share that input shape. + let table_schema = self.schema().await?; crate::table::computed_columns::ensure_supported_function_metadata(table_schema.as_ref())?; let table_def = TableDefinition::try_from_rich_schema(table_schema.clone())?; @@ -3269,7 +3278,10 @@ mod tests { use arrow::{array::AsArray, compute::concat_batches, datatypes::Int32Type}; use arrow_array::Array; use arrow_array::builder::LargeBinaryBuilder; - use arrow_array::{BinaryArray, Int32Array, RecordBatch, RecordBatchIterator, record_batch}; + use arrow_array::{ + BinaryArray, Int32Array, Int64Array, RecordBatch, RecordBatchIterator, StringArray, + StructArray, record_batch, + }; use arrow_schema::{DataType, Field, Schema}; use chrono::{DateTime, Utc}; use futures::{StreamExt, TryFutureExt, future::BoxFuture}; @@ -3621,6 +3633,88 @@ mod tests { assert_eq!(&body, &expected_body); } + #[tokio::test] + async fn add_rejects_external_blob_flag_before_any_request() { + let table = Table::new_with_handler::("my_table", |request| { + panic!("Unexpected request: {}", request.url().path()) + }); + let data = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])), + vec![Arc::new(Int32Array::from(vec![1]))], + ) + .unwrap(); + + let err = table + .add(data) + .allow_external_blob_outside_bases(true) + .execute() + .await + .unwrap_err(); + + assert!(matches!(err, Error::NotSupported { .. }), "got {err:?}"); + assert!(err.to_string().contains("local tables")); + } + + #[tokio::test] + async fn add_string_blob_becomes_uri_struct_without_the_local_flag() { + let table_schema = Schema::new(vec![ + Field::new("id", DataType::Int64, false), + crate::blob("image", true), + ]); + let describe_body = describe_response(&table_schema); + let input = RecordBatch::try_new( + Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int64, false), + Field::new("image", DataType::Utf8, true), + ])), + vec![ + Arc::new(Int64Array::from(vec![1])), + Arc::new(StringArray::from(vec![Some("s3://bucket/key")])), + ], + ) + .unwrap(); + + let (sender, receiver) = std::sync::mpsc::channel(); + let table = + Table::new_with_handler("my_table", move |mut request| match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(describe_body.clone()) + .unwrap(), + "/v1/table/my_table/insert/" => { + let mut body_out = reqwest::Body::from(Vec::new()); + std::mem::swap(request.body_mut().as_mut().unwrap(), &mut body_out); + sender.send(body_out).unwrap(); + http::Response::builder() + .status(200) + .body(r#"{"version": 2}"#.to_string()) + .unwrap() + } + path => panic!("Unexpected path: {path}"), + }); + + table.add(input).execute().await.unwrap(); + + let body = collect_body(receiver.recv().unwrap()).await; + let mut reader = + arrow_ipc::reader::StreamReader::try_new(std::io::Cursor::new(body), None).unwrap(); + let batch = reader.next().unwrap().unwrap(); + let image = batch + .column_by_name("image") + .unwrap() + .as_any() + .downcast_ref::() + .expect("remote add should send the coerced blob struct"); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "s3://bucket/key"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + } + #[rstest] #[case(true)] #[case(false)] diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index ecd95f161..efc36d260 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -3208,7 +3208,7 @@ impl BaseTable for NativeTable { let output = add.into_plan(&table_schema, &table_def)?; - let lance_params = output + let mut lance_params = output .write_options .lance_write_params .unwrap_or(WriteParams { @@ -3218,6 +3218,9 @@ impl BaseTable for NativeTable { }, ..Default::default() }); + if output.allow_external_blob_outside_bases { + lance_params.allow_external_blob_outside_bases = true; + } // Repartition for write parallelism if beneficial. let plan = if num_partitions > 1 { diff --git a/rust/lancedb/src/table/add_data.rs b/rust/lancedb/src/table/add_data.rs index 11ba43dd6..15ccf9f66 100644 --- a/rust/lancedb/src/table/add_data.rs +++ b/rust/lancedb/src/table/add_data.rs @@ -60,6 +60,7 @@ pub struct AddDataBuilder { pub(crate) embedding_registry: Option>, pub(crate) progress_callback: Option, pub(crate) write_parallelism: Option, + pub(crate) allow_external_blob_outside_bases: bool, } impl std::fmt::Debug for AddDataBuilder { @@ -87,6 +88,7 @@ impl AddDataBuilder { embedding_registry, progress_callback: None, write_parallelism: None, + allow_external_blob_outside_bases: false, } } @@ -141,6 +143,16 @@ impl AddDataBuilder { self } + /// Store blob URIs that sit outside registered blob bases. + /// + /// The row keeps a reference, so the object has to stay readable. + /// [`crate::table::Table::fetch_blobs`] reads from that location. + /// Defaults to `false`. Local tables only. + pub fn allow_external_blob_outside_bases(mut self, allow: bool) -> Self { + self.allow_external_blob_outside_bases = allow; + self + } + pub async fn execute(self) -> Result { if self.write_parallelism.map(|p| p == 0).unwrap_or(false) { return Err(Error::InvalidInput { @@ -199,6 +211,7 @@ impl AddDataBuilder { write_options: self.write_options, mode: self.mode, tracker, + allow_external_blob_outside_bases: self.allow_external_blob_outside_bases, }) } } @@ -212,6 +225,7 @@ pub struct PreprocessingOutput { pub write_options: WriteOptions, pub mode: AddDataMode, pub tracker: Option>, + pub allow_external_blob_outside_bases: bool, } /// Check that the input schema is valid for insert. diff --git a/rust/lancedb/src/table/datafusion/blob_coerce.rs b/rust/lancedb/src/table/datafusion/blob_coerce.rs index b29b2423b..cb984f7f4 100644 --- a/rust/lancedb/src/table/datafusion/blob_coerce.rs +++ b/rust/lancedb/src/table/datafusion/blob_coerce.rs @@ -7,7 +7,7 @@ use std::sync::Arc; -use arrow_schema::{DataType, Field, FieldRef}; +use arrow_schema::{DataType, Field, FieldRef, Fields}; use datafusion::functions::core::{get_field, named_struct}; use datafusion_common::ScalarValue; use datafusion_common::config::ConfigOptions; @@ -35,8 +35,9 @@ pub(super) fn coerce_blob_expr( }); }; - let input_struct_children = match input_field.data_type() { - DataType::Binary | DataType::LargeBinary | DataType::BinaryView => None, + let input_shape = match input_field.data_type() { + DataType::Binary | DataType::LargeBinary | DataType::BinaryView => BlobInputShape::Bytes, + DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View => BlobInputShape::String, DataType::Struct(children) => { if !children .iter() @@ -49,13 +50,15 @@ pub(super) fn coerce_blob_expr( ), }); } - Some(children) + BlobInputShape::Struct(children) } other => { return Err(Error::InvalidInput { message: format!( "cannot coerce column '{}' with type {} into a blob v2 struct. \ - expected Binary, LargeBinary, BinaryView, or a Struct with a 'data' or 'uri' child", + expected binary bytes (Binary, LargeBinary, BinaryView), \ + strings (Utf8, LargeUtf8, Utf8View), \ + or a Struct with a 'data' or 'uri' child", table_field.name(), other, ), @@ -69,9 +72,8 @@ pub(super) fn coerce_blob_expr( declared.name().as_str(), )))); - let value: Arc = match input_struct_children { - // Raw binary lands in `data` and everything else is a typed null. - None => { + let value: Arc = match &input_shape { + BlobInputShape::Bytes => { if declared.name() == "data" { Arc::new(CastExpr::new( input_expr.clone(), @@ -82,30 +84,43 @@ pub(super) fn coerce_blob_expr( typed_null(declared.data_type())? } } - Some(children) => match children.iter().find(|c| c.name() == declared.name()) { - Some(child) => { - let field_expr: Arc = Arc::new(ScalarFunctionExpr::new( - &format!("get_field({})", declared.name()), - get_field(), - vec![ - input_expr.clone(), - Arc::new(Literal::new(ScalarValue::from(declared.name().as_str()))), - ], - Arc::new(child.as_ref().clone()), - config.clone(), - )); - if child.data_type() == declared.data_type() { - field_expr - } else { - Arc::new(CastExpr::new( - field_expr, - declared.data_type().clone(), - None, - )) - } + BlobInputShape::String => { + if declared.name() == "uri" { + Arc::new(CastExpr::new( + input_expr.clone(), + declared.data_type().clone(), + None, + )) + } else { + typed_null(declared.data_type())? } - None => typed_null(declared.data_type())?, - }, + } + BlobInputShape::Struct(children) => { + match children.iter().find(|c| c.name() == declared.name()) { + Some(child) => { + let field_expr: Arc = Arc::new(ScalarFunctionExpr::new( + &format!("get_field({})", declared.name()), + get_field(), + vec![ + input_expr.clone(), + Arc::new(Literal::new(ScalarValue::from(declared.name().as_str()))), + ], + Arc::new(child.as_ref().clone()), + config.clone(), + )); + if child.data_type() == declared.data_type() { + field_expr + } else { + Arc::new(CastExpr::new( + field_expr, + declared.data_type().clone(), + None, + )) + } + } + None => typed_null(declared.data_type())?, + } + } }; ns_args.push(value); } @@ -120,6 +135,12 @@ pub(super) fn coerce_blob_expr( Ok((expr, table_field.clone())) } +enum BlobInputShape<'a> { + Bytes, + String, + Struct(&'a Fields), +} + fn typed_null(data_type: &DataType) -> Result> { let scalar = ScalarValue::try_from(data_type).map_err(|e| Error::InvalidInput { message: format!("cannot build null literal for blob child type {data_type}: {e}"), @@ -134,7 +155,7 @@ mod tests { use crate::blob::blob; use arrow_array::{ Array, ArrayRef, BinaryArray, BinaryViewArray, Int32Array, Int64Array, LargeBinaryArray, - RecordBatch, StringArray, StructArray, UInt8Array, UInt64Array, + RecordBatch, StringArray, StringViewArray, StructArray, UInt8Array, UInt64Array, }; use arrow_schema::Schema; use datafusion::prelude::SessionContext; @@ -436,14 +457,78 @@ mod tests { #[tokio::test] async fn unsupported_input_type_is_rejected_with_column_name() { let batch = batch_with_image( - Field::new("image", DataType::Utf8, true), - Arc::new(StringArray::from(vec!["not bytes"])), + Field::new("image", DataType::Int64, true), + Arc::new(Int64Array::from(vec![42])), ); let err = coerce_err(batch, &blob_table_schema()).await; assert!(matches!(err, Error::InvalidInput { .. }), "got {err:?}"); assert!(err.to_string().contains("image")); } + #[tokio::test] + async fn utf8_string_coerces_to_uri_child() { + let batch = batch_with_image( + Field::new("image", DataType::Utf8, true), + Arc::new(StringArray::from(vec![Some("s3://bucket/key"), None])), + ); + let coerced = coerce(batch, &blob_table_schema()).await; + let image = image_struct(&coerced); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "s3://bucket/key"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + assert!(uri.is_null(1)); + } + + #[tokio::test] + async fn large_utf8_string_coerces_into_four_child_blob_layout() { + use arrow_array::LargeStringArray; + + let table_schema = Schema::new(vec![ + Field::new("id", DataType::Int64, false), + wide_blob_field("image"), + ]); + let batch = batch_with_image( + Field::new("image", DataType::LargeUtf8, true), + Arc::new(LargeStringArray::from(vec!["file:///tmp/blob.bin"])), + ); + let coerced = coerce(batch, &table_schema).await; + let image = image_struct(&coerced); + assert_eq!(image.num_columns(), 4); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "file:///tmp/blob.bin"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + assert!(image.column_by_name("position").unwrap().is_null(0)); + assert!(image.column_by_name("size").unwrap().is_null(0)); + } + + #[tokio::test] + async fn utf8_view_string_coerces_to_uri_child() { + let batch = batch_with_image( + Field::new("image", DataType::Utf8View, true), + Arc::new(StringViewArray::from(vec![Some("s3://bucket/view-key")])), + ); + let coerced = coerce(batch, &blob_table_schema()).await; + let image = image_struct(&coerced); + let uri: &StringArray = image + .column_by_name("uri") + .unwrap() + .as_any() + .downcast_ref() + .unwrap(); + assert_eq!(uri.value(0), "s3://bucket/view-key"); + assert!(image.column_by_name("data").unwrap().is_null(0)); + } + #[tokio::test] async fn blob_metadata_survives_cast_of_sibling_column() { let batch = RecordBatch::try_new( diff --git a/rust/lancedb/tests/blob_integration.rs b/rust/lancedb/tests/blob_integration.rs index b92f961f4..7b709b645 100644 --- a/rust/lancedb/tests/blob_integration.rs +++ b/rust/lancedb/tests/blob_integration.rs @@ -5,12 +5,14 @@ use std::sync::Arc; use arrow_array::{ Array, ArrayRef, BinaryArray, Int64Array, LargeBinaryArray, RecordBatch, StringArray, - StructArray, UInt64Array, + StructArray, UInt64Array, new_null_array, }; use arrow_schema::{DataType, Field, Fields, Schema}; use futures::TryStreamExt; use lance::Dataset; +use lance::dataset::WriteParams; use lance_file::version::{ConcreteFileVersion, LanceFileVersion}; +use lance_table::format::BasePath; use lancedb::{ Connection, Error, Result, Table, blob::{BlobRangeRequest, blob}, @@ -19,7 +21,7 @@ use lancedb::{ ListingDatabaseOptions, NewTableConfig, OPT_NEW_TABLE_ENABLE_STABLE_ROW_IDS, }, query::{ExecutableQuery, QueryBase}, - table::{AddDataMode, CompactionOptions, OptimizeAction, OptimizeStats}, + table::{AddDataMode, CompactionOptions, OptimizeAction, OptimizeStats, WriteOptions}, }; use tempfile::tempdir; @@ -261,11 +263,11 @@ async fn add_rejects_uncoercible_blob_input() -> Result<()> { let batch = RecordBatch::try_new( Arc::new(Schema::new(vec![ Field::new("id", DataType::Int64, false), - Field::new("image", DataType::Utf8, true), + Field::new("image", DataType::Int64, true), ])), vec![ Arc::new(Int64Array::from(vec![1])), - Arc::new(StringArray::from(vec!["not bytes"])), + Arc::new(Int64Array::from(vec![42])), ], ) .unwrap(); @@ -1332,3 +1334,223 @@ async fn optimize_preserves_blob_v2_null_and_empty_distinction() -> Result<()> { ); Ok(()) } + +fn uri_struct_batch(id: i64, uri: &str) -> RecordBatch { + let image_field = blob("image", true); + let DataType::Struct(child_fields) = image_field.data_type().clone() else { + unreachable!("blob field is a struct"); + }; + let children: Vec = child_fields + .iter() + .map(|field| match field.name().as_str() { + "uri" => Arc::new(StringArray::from(vec![Some(uri)])) as ArrayRef, + _ => new_null_array(field.data_type(), 1), + }) + .collect(); + let image = StructArray::new(child_fields, children, None); + RecordBatch::try_new( + Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int64, false), + image_field, + ])), + vec![Arc::new(Int64Array::from(vec![id])), Arc::new(image)], + ) + .unwrap() +} + +fn uri_string_batch(id: i64, uri: &str) -> RecordBatch { + RecordBatch::try_new( + Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int64, false), + Field::new("image", DataType::Utf8, true), + ])), + vec![ + Arc::new(Int64Array::from(vec![id])), + Arc::new(StringArray::from(vec![Some(uri)])), + ], + ) + .unwrap() +} + +fn write_payload_file_uri(dir: &std::path::Path, name: &str, payload: &[u8]) -> String { + let path = dir.join(name); + std::fs::write(&path, payload).unwrap(); + url::Url::from_file_path(&path).unwrap().to_string() +} + +#[tokio::test] +async fn external_uri_struct_round_trips_with_flag() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let payload: &[u8] = b"external-struct-payload"; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", payload); + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + table + .add(uri_struct_batch(1, &uri)) + .allow_external_blob_outside_bases(true) + .execute() + .await?; + + let ids = collect_row_ids(&table).await?; + let bytes = table.fetch_blobs("image", &ids).await?; + assert_eq!(bytes.value(0), payload); + Ok(()) +} + +#[tokio::test] +async fn external_uri_add_requires_opt_in() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", b"unreachable"); + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + let err = table + .add(uri_struct_batch(1, &uri)) + .execute() + .await + .unwrap_err(); + + assert!( + err.to_string() + .contains("allow_external_blob_outside_bases"), + "got: {err}" + ); + assert_eq!(table.count_rows(None).await?, 0); + Ok(()) +} + +#[tokio::test] +async fn string_uri_input_round_trips_as_external_reference() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let payload: &[u8] = b"external-string-payload"; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", payload); + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + table + .add(uri_string_batch(1, &uri)) + .allow_external_blob_outside_bases(true) + .execute() + .await?; + + let ids = collect_row_ids(&table).await?; + let bytes = table.fetch_blobs("image", &ids).await?; + assert_eq!(bytes.value(0), payload); + + let files = table.fetch_blob_files("image", &ids).await?; + let file = files[0].as_ref().expect("missing blob file"); + assert_eq!(file.uri(), Some(uri.as_str())); + Ok(()) +} + +#[tokio::test] +async fn string_uri_inside_registered_base_does_not_need_the_flag() -> Result<()> { + let tmp = tempdir().unwrap(); + let db_path = tmp.path().join("db"); + let external_base = tmp.path().join("external_base"); + let object_dir = external_base.join("objects"); + std::fs::create_dir_all(&object_dir).unwrap(); + let payload: &[u8] = b"mapped-in-base"; + let object_path = object_dir.join("mapped.bin"); + std::fs::write(&object_path, payload).unwrap(); + let object_uri = url::Url::from_file_path(&object_path).unwrap().to_string(); + let base_uri = url::Url::from_file_path(&external_base) + .unwrap() + .to_string(); + + let db = connect(db_path.to_str().unwrap()).execute().await?; + let table = db + .create_empty_table("t", blob_table_schema()) + .write_options(WriteOptions { + lance_write_params: Some(WriteParams { + initial_bases: Some(vec![BasePath { + id: 1, + name: Some("external".to_string()), + path: base_uri, + is_dataset_root: false, + }]), + ..Default::default() + }), + }) + .execute() + .await?; + + table + .add(uri_string_batch(1, &object_uri)) + .execute() + .await?; + + let ids = collect_row_ids(&table).await?; + let bytes = table.fetch_blobs("image", &ids).await?; + assert_eq!(bytes.value(0), payload); + Ok(()) +} + +#[tokio::test] +async fn external_uri_rows_mix_with_inline_rows() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let external_payload: &[u8] = b"external-bytes"; + let uri = write_payload_file_uri(tmp.path(), "payload.bin", external_payload); + let table = + create_inline_blob_table(&db, "t", &[1], &[Some(b"inline-bytes".as_slice())]).await?; + + table + .add(uri_string_batch(2, &uri)) + .allow_external_blob_outside_bases(true) + .execute() + .await?; + + let pairs = collect_id_rowid(&table).await?; + let row_ids: Vec = pairs.iter().map(|(_, r)| *r).collect(); + let bytes = table.fetch_blobs("image", &row_ids).await?; + for (i, (id, _)) in pairs.iter().enumerate() { + match id { + 1 => assert_eq!(bytes.value(i), b"inline-bytes"), + 2 => assert_eq!(bytes.value(i), external_payload), + _ => unreachable!(), + } + } + Ok(()) +} + +#[tokio::test] +async fn malformed_string_uri_is_rejected_at_write() -> Result<()> { + let tmp = tempdir().unwrap(); + let db = connect(tmp.path().join("db").to_str().unwrap()) + .execute() + .await?; + let table = db + .create_empty_table("t", blob_table_schema()) + .execute() + .await?; + + let err = table + .add(uri_string_batch(1, "not a uri")) + .allow_external_blob_outside_bases(true) + .execute() + .await + .unwrap_err(); + + assert!(err.to_string().contains("not a uri"), "got: {err}"); + assert_eq!(table.count_rows(None).await?, 0); + Ok(()) +} From a57fb68891a081aac6e6466f772ff50b21445172 Mon Sep 17 00:00:00 2001 From: "lancedb-gatefixer[bot]" <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Wed, 26 Aug 2026 00:31:14 +0800 Subject: [PATCH 17/23] docs(python): fix Azure storage options examples (#3899) ## Summary - document that Azure Blob Storage credentials can be passed directly through `storage_options` - provide valid quoted `account_name` and `account_key` examples for both sync and async Python connections - execute the option dictionaries during doctests so the original unquoted-key mistake is caught ## Root cause The historical Python storage guide used `account_name` and `account_key` as bare identifiers in dictionary literals. Following that example either raised `NameError` or, when those names were predefined, produced incorrect option keys. The runtime already accepts direct Azure credentials, but the current Python API reference did not contain a corrected Azure example. ## Validation - `python/.venv/bin/ruff format --check python/python/lancedb/__init__.py` - `python/.venv/bin/ruff check .` - `cd python && uv run --no-sync pytest --doctest-modules python/lancedb/__init__.py -q` - `cd python && uv run --no-sync pytest python/tests/test_import.py -q` Fixes #2236 Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Co-authored-by: Xuanwo --- python/python/lancedb/__init__.py | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/python/python/lancedb/__init__.py b/python/python/lancedb/__init__.py index df8950686..8cb85a3ed 100644 --- a/python/python/lancedb/__init__.py +++ b/python/python/lancedb/__init__.py @@ -179,6 +179,18 @@ def connect( ... }, ... ) + For Azure Blob Storage, credentials can be passed directly without setting + environment variables: + + >>> azure_storage_options = { + ... "account_name": "some-account", + ... "account_key": "some-key", + ... } + >>> db = lancedb.connect( # doctest: +SKIP + ... "az://my-container/my-database", + ... storage_options=azure_storage_options, + ... ) + For tests and temporary data, use an in-memory database: >>> db = lancedb.connect("memory://") @@ -465,6 +477,10 @@ async def connect_async( -------- >>> import lancedb + >>> azure_storage_options = { + ... "account_name": "some-account", + ... "account_key": "some-key", + ... } >>> async def doctest_example(): ... # For a local directory, provide a path to the database ... db = await lancedb.connect_async("~/.lancedb") @@ -472,6 +488,11 @@ async def connect_async( ... db = await lancedb.connect_async("s3://my-bucket/lancedb", ... storage_options={ ... "aws_access_key_id": "***"}) + ... # Azure credentials can also be passed directly + ... db = await lancedb.connect_async( + ... "az://my-container/my-database", + ... storage_options=azure_storage_options, + ... ) ... # For tests and temporary data, use an in-memory database ... db = await lancedb.connect_async("memory://") ... # Connect to LanceDB cloud From 35b5d015ac8a38313d8322c26f4b2f86ef6f4e23 Mon Sep 17 00:00:00 2001 From: "lancedb-gatefixer[bot]" <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Wed, 26 Aug 2026 04:17:32 +0800 Subject: [PATCH 18/23] fix(node): preserve embedding registration in server bundles (#3806) ## Summary - lazily initialize built-in OpenAI and Hugging Face providers when consumers call the public embedding registry API - choose automatic vector versus FTS search from embedding metadata on a fresh pinned table revision for every execution - expose automatic string searches as an `AutoQuery` with only operations common to both native query families - keep the registry shared and built-in registration safe across duplicated module graphs ## Root cause Nitro treats dependency modules as side-effect-free and removes the bare OpenAI provider import from its generated route. Registration therefore never runs, so `getRegistry().get("openai")` remains undefined even when the registry itself is shared globally. Bundlers may also duplicate the provider and registry module graphs. The public embedding entry point now initializes built-in providers only when `getRegistry()` is explicitly called, keeping initialization on a live path that Nitro retains. Each terminal automatic-search execution pins the exact table revision visible at dispatch, reads embedding metadata and computes an embedding from that snapshot, replays the builder operations, and constructs and executes the selected native query against the same snapshot. Pinned native snapshots execute locally when namespace pushdown cannot carry their revision, while remote snapshots are seeded directly from one version-and-schema response. The public `AutoQuery` builder exposes only the operations shared by FTS and vector search, so runtime class narrowing cannot expose invalid vector-only methods. Repeated built-in registration replaces stale constructors from duplicated module graphs while public `register()` retains its duplicate-alias error. ## Validation - `cargo fmt --all` - `cargo check --quiet --features remote --tests --examples` - `cargo clippy --quiet --features remote --tests --examples` - `pnpm build` - `pnpm lint` - `pnpm run docs` - `pnpm test --runInBand` (783 passed, 5 skipped) - serial examples suite with a local OpenAI mock (11 passed), including `sentence-transformers.test.ts` - packaged Nitro 2.13.4 server route using the reported imports returned `{"registered":true}` - fresh-process FTS fixture initialized both public built-ins and confirmed automatic string search still returned the indexed row - schema-consistency regressions cover read-consistency refresh, checkout, checkoutLatest, restore, runtime class narrowing, concurrent overwrite during embedding computation, and reused automatic-search builders - focused regressions confirm pinned native snapshots bypass unversioned namespace pushdown and remote snapshots use one describe request Fixes #2429 --------- Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Co-authored-by: Xuanwo --- docs/src/js/classes/AutoQuery.md | 518 ++++++++++++++++++ docs/src/js/classes/Table.md | 4 +- docs/src/js/globals.md | 1 + .../embedding/functions/getRegistry.md | 14 +- nodejs/__test__/embedding_registry.test.ts | 95 ++++ nodejs/__test__/fixtures/auto_fts_search.cjs | 33 ++ nodejs/__test__/table.test.ts | 191 +++++++ nodejs/lancedb/embedding/index.ts | 44 +- nodejs/lancedb/embedding/openai.ts | 5 +- nodejs/lancedb/embedding/registry.ts | 60 +- nodejs/lancedb/embedding/transformers.ts | 5 +- nodejs/lancedb/index.ts | 1 + nodejs/lancedb/query.ts | 139 +++-- nodejs/lancedb/table.ts | 44 +- nodejs/src/table.rs | 6 + rust/lancedb/src/remote/table.rs | 38 ++ rust/lancedb/src/table.rs | 32 ++ rust/lancedb/src/table/query.rs | 49 +- 18 files changed, 1211 insertions(+), 68 deletions(-) create mode 100644 docs/src/js/classes/AutoQuery.md create mode 100644 nodejs/__test__/embedding_registry.test.ts create mode 100644 nodejs/__test__/fixtures/auto_fts_search.cjs diff --git a/docs/src/js/classes/AutoQuery.md b/docs/src/js/classes/AutoQuery.md new file mode 100644 index 000000000..1d7ea6952 --- /dev/null +++ b/docs/src/js/classes/AutoQuery.md @@ -0,0 +1,518 @@ +[**@lancedb/lancedb**](../README.md) • **Docs** + +*** + +[@lancedb/lancedb](../globals.md) / AutoQuery + +# Class: AutoQuery + +A builder for automatic string searches. + +Automatic search determines whether to use full-text or vector search from +the table revision selected for each execution. This builder exposes the +common operations supported by both query families. + +## Extends + +- `StandardQueryBase`<`NativeQuery` \| `NativeVectorQuery`> + +## Properties + +### inner + +```ts +protected inner: Query | VectorQuery | Promise; +``` + +#### Inherited from + +`StandardQueryBase.inner` + +## Methods + +### analyzePlan() + +```ts +analyzePlan(distributedMetrics?): Promise +``` + +Executes the query and returns the physical query plan annotated with runtime metrics. + +This is useful for debugging and performance analysis, as it shows how the query was executed +and includes metrics such as elapsed time, rows processed, and I/O statistics. + +#### Parameters + +* **distributedMetrics?**: [`AnalyzePlanDistributedMetrics`](../type-aliases/AnalyzePlanDistributedMetrics.md) + How distributed worker metrics are displayed for remote query plans. + Defaults to `"aggregate"`. + +#### Returns + +`Promise`<`string`> + +A query execution plan with runtime metrics for each step. + +#### Example + +```ts +import * as lancedb from "@lancedb/lancedb" + +const db = await lancedb.connect("./.lancedb"); +const table = await db.createTable("my_table", [ + { vector: [1.1, 0.9], id: "1" }, +]); + +const plan = await table.query().nearestTo([0.5, 0.2]).analyzePlan(); + +Example output (with runtime metrics inlined): +AnalyzeExec verbose=true, metrics=[] + ProjectionExec: expr=[id@3 as id, vector@0 as vector, _distance@2 as _distance], metrics=[output_rows=1, elapsed_compute=3.292µs] + Take: columns="vector, _rowid, _distance, (id)", metrics=[output_rows=1, elapsed_compute=66.001µs, batches_processed=1, bytes_read=8, iops=1, requests=1] + CoalesceBatchesExec: target_batch_size=1024, metrics=[output_rows=1, elapsed_compute=3.333µs] + GlobalLimitExec: skip=0, fetch=10, metrics=[output_rows=1, elapsed_compute=167ns] + FilterExec: _distance@2 IS NOT NULL, metrics=[output_rows=1, elapsed_compute=8.542µs] + SortExec: TopK(fetch=10), expr=[_distance@2 ASC NULLS LAST], metrics=[output_rows=1, elapsed_compute=63.25µs, row_replacements=1] + KNNVectorDistance: metric=l2, metrics=[output_rows=1, elapsed_compute=114.333µs, output_batches=1] + LanceScan: uri=/path/to/data, projection=[vector], row_id=true, row_addr=false, ordered=false, metrics=[output_rows=1, elapsed_compute=103.626µs, bytes_read=549, iops=2, requests=2] +``` + +#### Inherited from + +`StandardQueryBase.analyzePlan` + +*** + +### execute() + +```ts +protected execute(options?): AsyncGenerator, void, unknown> +``` + +Execute the query and return the results as an + +#### Parameters + +* **options?**: `Partial`<[`QueryExecutionOptions`](../interfaces/QueryExecutionOptions.md)> + +#### Returns + +`AsyncGenerator`<`RecordBatch`<`any`>, `void`, `unknown`> + +#### See + + - AsyncIterator +of + - RecordBatch. + +By default, LanceDb will use many threads to calculate results and, when +the result set is large, multiple batches will be processed at one time. +This readahead is limited however and backpressure will be applied if this +stream is consumed slowly (this constrains the maximum memory used by a +single query) + +#### Inherited from + +`StandardQueryBase.execute` + +*** + +### explainPlan() + +```ts +explainPlan(verbose): Promise +``` + +Generates an explanation of the query execution plan. + +#### Parameters + +* **verbose**: `boolean` = `false` + If true, provides a more detailed explanation. Defaults to false. + +#### Returns + +`Promise`<`string`> + +A Promise that resolves to a string containing the query execution plan explanation. + +#### Example + +```ts +import * as lancedb from "@lancedb/lancedb" +const db = await lancedb.connect("./.lancedb"); +const table = await db.createTable("my_table", [ + { vector: [1.1, 0.9], id: "1" }, +]); +const plan = await table.query().nearestTo([0.5, 0.2]).explainPlan(); +``` + +#### Inherited from + +`StandardQueryBase.explainPlan` + +*** + +### fastSearch() + +```ts +fastSearch(): this +``` + +Skip searching un-indexed data. This can make search faster, but will miss +any data that is not yet indexed. + +Use [Table#optimize](Table.md#optimize) to index all un-indexed data. + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.fastSearch` + +*** + +### ~~filter()~~ + +```ts +filter(predicate): this +``` + +A filter statement to be applied to this query. + +#### Parameters + +* **predicate**: `string` + +#### Returns + +`this` + +#### See + +where + +#### Deprecated + +Use `where` instead + +#### Inherited from + +`StandardQueryBase.filter` + +*** + +### fullTextSearch() + +```ts +fullTextSearch(query, options?): this +``` + +#### Parameters + +* **query**: `string` \| [`FullTextQuery`](../interfaces/FullTextQuery.md) + +* **options?**: `Partial`<[`FullTextSearchOptions`](../interfaces/FullTextSearchOptions.md)> + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.fullTextSearch` + +*** + +### limit() + +```ts +limit(limit): this +``` + +Set the maximum number of results to return. + +By default, a plain search has no limit. If this method is not +called then every valid row from the table will be returned. + +#### Parameters + +* **limit**: `number` + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.limit` + +*** + +### offset() + +```ts +offset(offset): this +``` + +Set the number of rows to skip before returning results. + +This is useful for pagination. + +#### Parameters + +* **offset**: `number` + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.offset` + +*** + +### orderBy() + +```ts +orderBy(ordering): this +``` + +Sort the results by the specified column(s). + +#### Parameters + +* **ordering**: [`ColumnOrdering`](../interfaces/ColumnOrdering.md) \| [`ColumnOrdering`](../interfaces/ColumnOrdering.md)[] + +#### Returns + +`this` + +This query builder. + +#### Inherited from + +`StandardQueryBase.orderBy` + +*** + +### outputSchema() + +```ts +outputSchema(): Promise> +``` + +Returns the schema of the output that will be returned by this query. + +This can be used to inspect the types and names of the columns that will be +returned by the query before executing it. + +#### Returns + +`Promise`<`Schema`<`any`>> + +An Arrow Schema describing the output columns. + +#### Inherited from + +`StandardQueryBase.outputSchema` + +*** + +### select() + +```ts +select(columns): this +``` + +Return only the specified columns. + +By default a query will return all columns from the table. However, this can have +a very significant impact on latency. LanceDb stores data in a columnar fashion. This +means we can finely tune our I/O to select exactly the columns we need. + +As a best practice you should always limit queries to the columns that you need. If you +pass in an array of column names then only those columns will be returned. + +You can also use this method to create new "dynamic" columns based on your existing columns. +For example, you may not care about "a" or "b" but instead simply want "a + b". This is often +seen in the SELECT clause of an SQL query (e.g. `SELECT a+b FROM my_table`). + +To create dynamic columns you can pass in a Map. A column will be returned +for each entry in the map. The key provides the name of the column. The value is +an SQL string used to specify how the column is calculated. + +For example, an SQL query might state `SELECT a + b AS combined, c`. The equivalent +input to this method would be: + +#### Parameters + +* **columns**: `string` \| `string`[] \| `Record`<`string`, `string`> \| `Map`<`string`, `string`> + +#### Returns + +`this` + +#### Example + +```ts +new Map([["combined", "a + b"], ["c", "c"]]) + +Columns will always be returned in the order given, even if that order is different than +the order used when adding the data. + +Note that you can pass in a `Record` (e.g. an object literal). This method +uses `Object.entries` which should preserve the insertion order of the object. However, +object insertion order is easy to get wrong and `Map` is more foolproof. +``` + +#### Inherited from + +`StandardQueryBase.select` + +*** + +### toArray() + +```ts +toArray(options?): Promise +``` + +Collect the results as an array of objects. + +#### Parameters + +* **options?**: `Partial`<[`QueryExecutionOptions`](../interfaces/QueryExecutionOptions.md)> + +#### Returns + +`Promise`<`any`[]> + +#### Inherited from + +`StandardQueryBase.toArray` + +*** + +### toArrow() + +```ts +toArrow(options?): Promise> +``` + +Collect the results as an Arrow + +#### Parameters + +* **options?**: `Partial`<[`QueryExecutionOptions`](../interfaces/QueryExecutionOptions.md)> + +#### Returns + +`Promise`<`Table`<`any`>> + +#### See + +ArrowTable. + +#### Inherited from + +`StandardQueryBase.toArrow` + +*** + +### useLsm() + +```ts +useLsm(enable): this +``` + +Control MemWAL read routing for this query. + +By default (unset), when the table carries a MemWAL write spec (see +[Table#setLsmWriteSpec](Table.md#setlsmwritespec)), reads are routed through the LSM scanner so +they also return data written via the `mergeInsert` LSM path that has not yet +been compacted into the base table (the active/frozen in-memory memtables and +the flushed generations), deduplicated by primary key; a table without a spec +reads the base table. + +#### Parameters + +* **enable**: `boolean` + `true` forces the LSM scanner and errors if the table has no + MemWAL write spec. `false` bypasses the MemWAL and reads the base table only, + even when a spec is present. + Note: the LSM scanner does not support every query shape (e.g. reranking, + hybrid search, `orderBy`). On a MemWAL table those shapes error unless + `useLsm(false)` is set, because a base-only read would silently exclude + un-compacted MemWAL data. + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.useLsm` + +*** + +### where() + +```ts +where(predicate): this +``` + +A filter statement to be applied to this query. + +The filter should be supplied as an SQL query string. For example: + +#### Parameters + +* **predicate**: `string` + +#### Returns + +`this` + +#### Example + +```ts +x > 10 +y > 0 AND y < 100 +x > 5 OR y = 'test' + +Filtering performance can often be improved by creating a scalar index +on the filter column(s). + +Calling this multiple times combines the filters with a logical AND rather +than replacing the previous filter. +``` + +#### Inherited from + +`StandardQueryBase.where` + +*** + +### withRowId() + +```ts +withRowId(): this +``` + +Whether to return the row id in the results. + +This column can be used to match results between different queries. For +example, to match results from a full text search and a vector search in +order to perform hybrid search. + +#### Returns + +`this` + +#### Inherited from + +`StandardQueryBase.withRowId` diff --git a/docs/src/js/classes/Table.md b/docs/src/js/classes/Table.md index 06dc8479e..9a85d0d96 100644 --- a/docs/src/js/classes/Table.md +++ b/docs/src/js/classes/Table.md @@ -942,7 +942,7 @@ Get the schema of the table. abstract search( query, queryType?, - ftsColumns?): Query | VectorQuery + ftsColumns?): Query | VectorQuery | AutoQuery ``` Create a search query to find the nearest neighbors @@ -964,7 +964,7 @@ of the given query #### Returns -[`Query`](Query.md) \| [`VectorQuery`](VectorQuery.md) +[`Query`](Query.md) \| [`VectorQuery`](VectorQuery.md) \| [`AutoQuery`](AutoQuery.md) *** diff --git a/docs/src/js/globals.md b/docs/src/js/globals.md index e0635ab65..beb9cbeff 100644 --- a/docs/src/js/globals.md +++ b/docs/src/js/globals.md @@ -18,6 +18,7 @@ ## Classes +- [AutoQuery](classes/AutoQuery.md) - [BooleanQuery](classes/BooleanQuery.md) - [BoostQuery](classes/BoostQuery.md) - [BranchContents](classes/BranchContents.md) diff --git a/docs/src/js/namespaces/embedding/functions/getRegistry.md b/docs/src/js/namespaces/embedding/functions/getRegistry.md index 331dbd60b..b149a4a88 100644 --- a/docs/src/js/namespaces/embedding/functions/getRegistry.md +++ b/docs/src/js/namespaces/embedding/functions/getRegistry.md @@ -10,16 +10,12 @@ function getRegistry(): EmbeddingFunctionRegistry ``` -Utility function to get the global instance of the registry +Get the global embedding function registry. + +LanceDB built-in providers are initialized when this public API is first +used, so importing the root package does not change automatic search +selection for tables without embedding metadata. ## Returns [`EmbeddingFunctionRegistry`](../classes/EmbeddingFunctionRegistry.md) - -`EmbeddingFunctionRegistry` The global instance of the registry - -## Example - -```ts -const registry = getRegistry(); -const openai = registry.get("openai").create(); diff --git a/nodejs/__test__/embedding_registry.test.ts b/nodejs/__test__/embedding_registry.test.ts new file mode 100644 index 000000000..83933399a --- /dev/null +++ b/nodejs/__test__/embedding_registry.test.ts @@ -0,0 +1,95 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +import { execFileSync } from "node:child_process"; +import { resolve } from "node:path"; + +import type { OpenAIEmbeddingFunction } from "../lancedb/embedding/openai"; +import type { EmbeddingFunctionRegistry } from "../lancedb/embedding/registry"; + +type EmbeddingModule = typeof import("../lancedb/embedding"); +type OpenAIModule = typeof import("../lancedb/embedding/openai"); +type RegistryModule = typeof import("../lancedb/embedding/registry"); + +describe("embedding function registry", () => { + const registries: EmbeddingFunctionRegistry[] = []; + + afterEach(() => { + for (const registry of registries) { + registry.reset(); + } + registries.length = 0; + }); + + it("defers built-in providers until the public registry API is used", () => { + jest.isolateModules(() => { + const embedding = require("../lancedb/embedding") as EmbeddingModule; + const { getRegistry: getInternalRegistry } = + require("../lancedb/embedding/registry") as RegistryModule; + const registry = getInternalRegistry(); + registries.push(registry); + + expect(registry.length()).toBe(0); + expect(embedding.getRegistry()).toBe(registry); + expect(registry.get("openai")).toBeDefined(); + expect(registry.get("huggingface")).toBeDefined(); + }); + }); + + it("preserves automatic FTS search in a fresh process", () => { + execFileSync( + process.execPath, + [resolve(__dirname, "fixtures", "auto_fts_search.cjs")], + { stdio: "pipe" }, + ); + }); + + it("shares registrations across duplicated provider module graphs", () => { + let registeringRegistry: EmbeddingFunctionRegistry | undefined; + let latestOpenAIConstructor: typeof OpenAIEmbeddingFunction | undefined; + + jest.isolateModules(() => { + require("../lancedb/embedding/openai"); + const { getRegistry } = + require("../lancedb/embedding/registry") as RegistryModule; + registeringRegistry = getRegistry(); + registries.push(registeringRegistry); + expect(registeringRegistry.get("openai")).toBeDefined(); + }); + + expect(() => { + jest.isolateModules(() => { + const { OpenAIEmbeddingFunction } = + require("../lancedb/embedding/openai") as OpenAIModule; + latestOpenAIConstructor = OpenAIEmbeddingFunction; + const { getRegistry } = + require("../lancedb/embedding/registry") as RegistryModule; + registries.push(getRegistry()); + }); + }).not.toThrow(); + + const previousApiKey = process.env.OPENAI_API_KEY; + process.env.OPENAI_API_KEY = "test"; + try { + const latestOpenAI = registeringRegistry! + .get("openai")! + .create(); + expect(latestOpenAI).toBeInstanceOf(latestOpenAIConstructor!); + } finally { + if (previousApiKey === undefined) { + delete process.env.OPENAI_API_KEY; + } else { + process.env.OPENAI_API_KEY = previousApiKey; + } + } + + jest.isolateModules(() => { + const { getRegistry } = + require("../lancedb/embedding") as EmbeddingModule; + const publicRegistry = getRegistry(); + registries.push(publicRegistry); + expect(publicRegistry).toBe(registeringRegistry); + expect(publicRegistry.get("openai")).toBeDefined(); + }); + }); +}); diff --git a/nodejs/__test__/fixtures/auto_fts_search.cjs b/nodejs/__test__/fixtures/auto_fts_search.cjs new file mode 100644 index 000000000..b5ab060b3 --- /dev/null +++ b/nodejs/__test__/fixtures/auto_fts_search.cjs @@ -0,0 +1,33 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +const assert = require("node:assert/strict"); +const tmp = require("tmp"); +const { connect, embedding, Index } = require("../../dist"); +const { getRegistry } = require("../../dist/embedding/registry"); + +async function main() { + assert.equal(typeof embedding.getRegistry, "function"); + assert.equal(getRegistry().length(), 0); + assert.equal(embedding.getRegistry(), getRegistry()); + assert.equal(getRegistry().length(), 2); + + const dir = tmp.dirSync({ unsafeCleanup: true }); + let db; + try { + db = await connect(dir.name); + const table = await db.createTable("docs", [{ text: "hello world" }]); + await table.createIndex("text", { config: Index.fts() }); + + const rows = await table.search("hello").toArray(); + assert.equal(rows[0].text, "hello world"); + } finally { + db?.close(); + dir.removeCallback(); + } +} + +main().catch((error) => { + console.error(error); + process.exitCode = 1; +}); diff --git a/nodejs/__test__/table.test.ts b/nodejs/__test__/table.test.ts index 0f6ed3615..0aee5caf0 100644 --- a/nodejs/__test__/table.test.ts +++ b/nodejs/__test__/table.test.ts @@ -11,10 +11,13 @@ import * as arrow17 from "apache-arrow-17"; import * as arrow18 from "apache-arrow-18"; import { + AutoQuery, Connection, MatchQuery, PhraseQuery, + Query, Table, + VectorQuery, connect, tokenize, } from "../lancedb"; @@ -1777,6 +1780,194 @@ describe("Read consistency interval", () => { }); }); +describe("automatic search schema consistency", () => { + let tmpDir: tmp.DirResult; + + class SchemaRefreshEmbedding extends EmbeddingFunction { + ndims() { + return 2; + } + + embeddingDataType() { + return new Float32(); + } + + async computeSourceEmbeddings(data: string[]) { + return data.map((value) => [value.length, 1]); + } + + async computeQueryEmbeddings(value: string) { + return [value.length, 1]; + } + } + + function embeddingSchema() { + const func = new SchemaRefreshEmbedding(); + return LanceSchema({ + text: func.sourceField(new Utf8()), + vector: func.vectorField(), + }); + } + + beforeEach(() => { + getRegistry().reset(); + register("schema-refresh")(SchemaRefreshEmbedding); + tmpDir = tmp.dirSync({ unsafeCleanup: true }); + }); + + afterEach(() => { + getRegistry().reset(); + tmpDir.removeCallback(); + }); + + it("uses the schema refreshed from another connection", async () => { + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + const stale = await first.createTable("docs", [{ text: "before" }], { + schema: embeddingSchema(), + }); + const replacement = await second.createTable( + "docs", + [{ text: "after hello" }], + { mode: "overwrite" }, + ); + await replacement.createIndex("text", { config: Index.fts() }); + + const search = stale.search("hello"); + expect(search).toBeInstanceOf(AutoQuery); + expect(search).not.toBeInstanceOf(Query); + expect(search).not.toBeInstanceOf(VectorQuery); + expect("nprobes" in search).toBe(false); + + const rows = await search.toArray(); + expect(rows[0].text).toBe("after hello"); + expect((await stale.schema()).metadata.has("embedding_functions")).toBe( + false, + ); + } finally { + first.close(); + second.close(); + } + }); + + it("tracks embedding metadata across checkout and restore", async () => { + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + await first.createTable("docs", [{ text: "before" }], { + schema: embeddingSchema(), + }); + const table = await second.createTable( + "docs", + [{ text: "after hello" }], + { mode: "overwrite" }, + ); + await table.createIndex("text", { config: Index.fts() }); + + await table.checkout(1); + expect((await table.search("before").toArray())[0].text).toBe("before"); + + await table.checkoutLatest(); + expect((await table.search("hello").toArray())[0].text).toBe( + "after hello", + ); + + await table.checkout(1); + await table.restore(); + expect((await table.search("before").toArray())[0].text).toBe("before"); + } finally { + first.close(); + second.close(); + } + }); + + it("pins automatic search while computing an embedding", async () => { + let markStarted!: () => void; + let releaseEmbedding!: () => void; + const started = new Promise((resolve) => { + markStarted = resolve; + }); + const released = new Promise((resolve) => { + releaseEmbedding = resolve; + }); + + class BlockingEmbedding extends SchemaRefreshEmbedding { + async computeQueryEmbeddings(value: string) { + markStarted(); + await released; + return [value.length, 1]; + } + } + + register("schema-refresh-blocking")(BlockingEmbedding); + const func = new BlockingEmbedding(); + const schema = LanceSchema({ + text: func.sourceField(new Utf8()), + vector: func.vectorField(), + }); + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + const table = await first.createTable( + "docs", + [{ text: "hello before" }], + { schema }, + ); + const pending = table.search("hello").toArray(); + await started; + + const replacement = await second.createTable( + "docs", + [{ text: "hello after" }], + { mode: "overwrite" }, + ); + await replacement.createIndex("text", { config: Index.fts() }); + releaseEmbedding(); + + expect((await pending)[0].text).toBe("hello before"); + } finally { + releaseEmbedding(); + first.close(); + second.close(); + } + }); + + it("refreshes a reused automatic search for every execution", async () => { + const first = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + const second = await connect(tmpDir.name, { readConsistencyInterval: 0 }); + + try { + const table = await first.createTable("docs", [ + { text: "hello before", marker: "before" }, + ]); + await table.createIndex("text", { config: Index.fts() }); + const search = table.search("hello").select(["text"]); + + const before = (await search.toArray())[0]; + expect(before.text).toBe("hello before"); + expect(before.marker).toBeUndefined(); + + const replacement = await second.createTable( + "docs", + [{ text: "hello after", marker: "after" }], + { mode: "overwrite" }, + ); + await replacement.createIndex("text", { config: Index.fts() }); + + const after = (await search.toArray())[0]; + expect(after.text).toBe("hello after"); + expect(after.marker).toBeUndefined(); + } finally { + first.close(); + second.close(); + } + }); +}); + describe("schema evolution", function () { let tmpDir: tmp.DirResult; beforeEach(() => { diff --git a/nodejs/lancedb/embedding/index.ts b/nodejs/lancedb/embedding/index.ts index d0ffaec0d..d748e245a 100644 --- a/nodejs/lancedb/embedding/index.ts +++ b/nodejs/lancedb/embedding/index.ts @@ -4,7 +4,15 @@ import { Field, Schema } from "../arrow"; import { sanitizeType } from "../sanitize"; import { EmbeddingFunction } from "./embedding_function"; -import { EmbeddingFunctionConfig, getRegistry } from "./registry"; +import { + EmbeddingFunctionConfig, + EmbeddingFunctionRegistry, + getRegistry as getGlobalRegistry, + registerBuiltIn, +} from "./registry"; + +type OpenAIModule = typeof import("./openai"); +type TransformersModule = typeof import("./transformers"); export { FieldOptions, @@ -14,7 +22,39 @@ export { EmbeddingFunctionConstructor, } from "./embedding_function"; -export * from "./registry"; +export { + EmbeddingFunctionRegistry, + parseEmbeddingMetadata, + register, +} from "./registry"; +export type { + CreateReturnType, + EmbeddingFunctionConfig, + EmbeddingFunctionCreate, + EmbeddingMetadataEntry, + ResolvedEmbeddingFunctionConfig, +} from "./registry"; + +function initializeBuiltInProviders() { + const { OpenAIEmbeddingFunction } = require("./openai") as OpenAIModule; + const { TransformersEmbeddingFunction } = + require("./transformers") as TransformersModule; + + registerBuiltIn("openai", OpenAIEmbeddingFunction); + registerBuiltIn("huggingface", TransformersEmbeddingFunction); +} + +/** + * Get the global embedding function registry. + * + * LanceDB built-in providers are initialized when this public API is first + * used, so importing the root package does not change automatic search + * selection for tables without embedding metadata. + */ +export function getRegistry(): EmbeddingFunctionRegistry { + initializeBuiltInProviders(); + return getGlobalRegistry(); +} /** * Create a schema with embedding functions. diff --git a/nodejs/lancedb/embedding/openai.ts b/nodejs/lancedb/embedding/openai.ts index 5771cfeb5..2218d44bc 100644 --- a/nodejs/lancedb/embedding/openai.ts +++ b/nodejs/lancedb/embedding/openai.ts @@ -5,14 +5,13 @@ import type OpenAI from "openai"; import type { EmbeddingCreateParams } from "openai/resources/index"; import { Float, Float32 } from "../arrow"; import { EmbeddingFunction } from "./embedding_function"; -import { register } from "./registry"; +import { registerBuiltIn } from "./registry"; export type OpenAIOptions = { apiKey: string; model: EmbeddingCreateParams["model"]; }; -@register("openai") export class OpenAIEmbeddingFunction extends EmbeddingFunction< string, Partial @@ -100,3 +99,5 @@ export class OpenAIEmbeddingFunction extends EmbeddingFunction< return response.data[0].embedding; } } + +registerBuiltIn("openai", OpenAIEmbeddingFunction); diff --git a/nodejs/lancedb/embedding/registry.ts b/nodejs/lancedb/embedding/registry.ts index 5f32f683c..c9ee9135e 100644 --- a/nodejs/lancedb/embedding/registry.ts +++ b/nodejs/lancedb/embedding/registry.ts @@ -7,6 +7,10 @@ import { } from "./embedding_function"; import "reflect-metadata"; +const builtInFunctionsKey = Symbol.for( + "@lancedb/lancedb::embedding-built-in-functions::v1", +); + export type CreateReturnType = T extends { init: () => Promise } ? Promise : T; @@ -59,6 +63,15 @@ export class EmbeddingFunctionRegistry { }; } + /** @ignore */ + setBuiltIn< + T extends EmbeddingFunctionConstructor = EmbeddingFunctionConstructor, + >(name: string, ctor: T): T { + this.#functions.set(name, ctor); + Reflect.defineMetadata("lancedb::embedding::name", name, ctor); + return ctor; + } + get>( name: string, ): EmbeddingFunctionCreate | undefined; @@ -96,6 +109,7 @@ export class EmbeddingFunctionRegistry { */ reset(this: EmbeddingFunctionRegistry) { this.#functions.clear(); + getBuiltInFunctions(this).clear(); } /** @@ -183,12 +197,56 @@ export class EmbeddingFunctionRegistry { } } -const _REGISTRY = new EmbeddingFunctionRegistry(); +function getBuiltInFunctions(registry: EmbeddingFunctionRegistry): Set { + const registryWithBuiltIns = registry as EmbeddingFunctionRegistry & { + [key: symbol]: Set | undefined; + }; + let builtInFunctions = registryWithBuiltIns[builtInFunctionsKey]; + if (builtInFunctions === undefined) { + builtInFunctions = new Set(); + registryWithBuiltIns[builtInFunctionsKey] = builtInFunctions; + } + return builtInFunctions; +} + +// Server bundlers can load the side-effect embedding entry points and the public +// embedding API from separate module graphs. Keep their registry shared. +const registryKey = Symbol.for( + "@lancedb/lancedb::embedding-function-registry::v1", +); +const registryGlobal = globalThis as typeof globalThis & { + [key: symbol]: EmbeddingFunctionRegistry | undefined; +}; + +function getGlobalRegistry(): EmbeddingFunctionRegistry { + const existingRegistry = registryGlobal[registryKey]; + if (existingRegistry !== undefined) { + return existingRegistry; + } + const registry = new EmbeddingFunctionRegistry(); + registryGlobal[registryKey] = registry; + return registry; +} + +const _REGISTRY = getGlobalRegistry(); export function register(name?: string) { return _REGISTRY.register(name); } +/** @ignore */ +export function registerBuiltIn< + T extends EmbeddingFunctionConstructor = EmbeddingFunctionConstructor, +>(name: string, ctor: T): T { + const builtInFunctions = getBuiltInFunctions(_REGISTRY); + if (builtInFunctions.has(name)) { + return _REGISTRY.setBuiltIn(name, ctor); + } + _REGISTRY.register(name)(ctor); + builtInFunctions.add(name); + return ctor; +} + /** * Utility function to get the global instance of the registry * @returns `EmbeddingFunctionRegistry` The global instance of the registry diff --git a/nodejs/lancedb/embedding/transformers.ts b/nodejs/lancedb/embedding/transformers.ts index 161575285..06043ea7c 100644 --- a/nodejs/lancedb/embedding/transformers.ts +++ b/nodejs/lancedb/embedding/transformers.ts @@ -3,7 +3,7 @@ import { Float, Float32 } from "../arrow"; import { EmbeddingFunction } from "./embedding_function"; -import { register } from "./registry"; +import { registerBuiltIn } from "./registry"; export type XenovaTransformerOptions = { /** The wasm compatible model to use */ @@ -31,7 +31,6 @@ export type XenovaTransformerOptions = { }; }; -@register("huggingface") export class TransformersEmbeddingFunction extends EmbeddingFunction< string, Partial @@ -158,6 +157,8 @@ export class TransformersEmbeddingFunction extends EmbeddingFunction< } } +registerBuiltIn("huggingface", TransformersEmbeddingFunction); + const tensorDiv = ( src: import("@huggingface/transformers").Tensor, divBy: number, diff --git a/nodejs/lancedb/index.ts b/nodejs/lancedb/index.ts index ebc8cda8d..34d7ce4d9 100644 --- a/nodejs/lancedb/index.ts +++ b/nodejs/lancedb/index.ts @@ -103,6 +103,7 @@ export { } from "./native.js"; export { + AutoQuery, ExecutableQuery, Query, QueryBase, diff --git a/nodejs/lancedb/query.ts b/nodejs/lancedb/query.ts index 843a1276f..3b9b286a0 100644 --- a/nodejs/lancedb/query.ts +++ b/nodejs/lancedb/query.ts @@ -111,13 +111,15 @@ export class QueryBase< NativeQueryType extends NativeQuery | NativeVectorQuery | NativeTakeQuery, > implements AsyncIterable { + protected inner!: NativeQueryType | Promise; + /** * @hidden */ - protected constructor( - protected inner: NativeQueryType | Promise, - ) { - // intentionally empty + protected constructor(inner?: NativeQueryType | Promise) { + if (inner !== undefined) { + this.inner = inner; + } } // call a function on the inner (either a promise or the actual object) @@ -135,6 +137,15 @@ export class QueryBase< } } + /** + * Return the native query used by the next terminal operation. + * + * @hidden + */ + protected async getInner(): Promise { + return this.inner; + } + /** * Return only the specified columns. * @@ -207,16 +218,11 @@ export class QueryBase< /** * @hidden */ - protected nativeExecute( + protected async nativeExecute( options?: Partial, ): Promise { - if (this.inner instanceof Promise) { - return this.inner.then((inner) => - inner.execute(options?.maxBatchLength, options?.timeoutMs), - ); - } else { - return this.inner.execute(options?.maxBatchLength, options?.timeoutMs); - } + const inner = await this.getInner(); + return inner.execute(options?.maxBatchLength, options?.timeoutMs); } /** @@ -245,12 +251,7 @@ export class QueryBase< /** Collect the results as an Arrow @see {@link ArrowTable}. */ async toArrow(options?: Partial): Promise { const batches = []; - let inner; - if (this.inner instanceof Promise) { - inner = await this.inner; - } else { - inner = this.inner; - } + const inner = await this.getInner(); for await (const batch of new RecordBatchIterable(inner, options)) { batches.push(batch); } @@ -279,11 +280,8 @@ export class QueryBase< * @returns A Promise that resolves to a string containing the query execution plan explanation. */ async explainPlan(verbose = false): Promise { - if (this.inner instanceof Promise) { - return this.inner.then((inner) => inner.explainPlan(verbose)); - } else { - return this.inner.explainPlan(verbose); - } + const inner = await this.getInner(); + return inner.explainPlan(verbose); } /** @@ -321,13 +319,8 @@ export class QueryBase< distributedMetrics?: AnalyzePlanDistributedMetrics, ): Promise { const distributedMetricsMode = distributedMetrics ?? "aggregate"; - if (this.inner instanceof Promise) { - return this.inner.then((inner) => - inner.analyzePlan(distributedMetricsMode), - ); - } else { - return this.inner.analyzePlan(distributedMetricsMode); - } + const inner = await this.getInner(); + return inner.analyzePlan(distributedMetricsMode); } /** @@ -339,12 +332,8 @@ export class QueryBase< * @returns An Arrow Schema describing the output columns. */ async outputSchema(): Promise { - let schemaBuffer: Buffer; - if (this.inner instanceof Promise) { - schemaBuffer = await this.inner.then((inner) => inner.outputSchema()); - } else { - schemaBuffer = await this.inner.outputSchema(); - } + const inner = await this.getInner(); + const schemaBuffer = await inner.outputSchema(); const schema = tableFromIPC(schemaBuffer).schema; return schema; } @@ -356,7 +345,7 @@ export class StandardQueryBase< extends QueryBase implements ExecutableQuery { - constructor(inner: NativeQueryType | Promise) { + constructor(inner?: NativeQueryType | Promise) { super(inner); } @@ -788,6 +777,51 @@ export class TakeQuery extends QueryBase { } } +/** + * A builder for automatic string searches. + * + * Automatic search determines whether to use full-text or vector search from + * the table revision selected for each execution. This builder exposes the + * common operations supported by both query families. + * + * @hideconstructor + */ +export class AutoQuery extends StandardQueryBase< + NativeQuery | NativeVectorQuery +> { + private readonly calls: Array< + (inner: NativeQuery | NativeVectorQuery) => void + > = []; + + /** @hidden */ + constructor( + private readonly createInner: () => Promise< + NativeQuery | NativeVectorQuery + >, + ) { + super(); + } + + /** @hidden */ + protected override doCall( + fn: (inner: NativeQuery | NativeVectorQuery) => void, + ) { + this.calls.push(fn); + } + + /** @hidden */ + protected override async getInner(): Promise< + NativeQuery | NativeVectorQuery + > { + const calls = [...this.calls]; + const inner = await this.createInner(); + for (const call of calls) { + call(inner); + } + return inner; + } +} + /** A builder for LanceDB queries. * * @see {@link Table#query}, {@link Table#search} @@ -802,6 +836,37 @@ export class Query extends StandardQueryBase { super(tbl.query()); } + /** @hidden */ + static autoSearch( + tbl: () => Promise, + query: string, + vector: (tbl: NativeTable) => Promise | undefined>, + columns?: string[], + ): AutoQuery { + const nativeQuery = async () => { + const snapshot = await Promise.resolve(tbl()); + const resolved = await vector(snapshot); + const inner = snapshot.query(); + if (resolved === undefined) { + inner.fullTextSearch({ + query, + columns: columns ?? null, + }); + return inner; + } + + const raw = Array.isArray(resolved) + ? null + : extractVectorBuffer(resolved); + if (raw) { + return inner.nearestToRaw(raw.data, raw.dtype); + } + return inner.nearestTo(Float32Array.from(resolved as number[])); + }; + + return new AutoQuery(nativeQuery); + } + /** * Find the nearest vectors to the given query vector. * diff --git a/nodejs/lancedb/table.ts b/nodejs/lancedb/table.ts index 28603cca9..a4fc76ef1 100644 --- a/nodejs/lancedb/table.ts +++ b/nodejs/lancedb/table.ts @@ -43,6 +43,7 @@ import { Table as _NativeTable, } from "./native"; import { + AutoQuery, FullTextQuery, Query, TakeQuery, @@ -523,7 +524,7 @@ export abstract class Table { query: string | IntoVector | MultiVector | FullTextQuery, queryType?: string, ftsColumns?: string | string[], - ): VectorQuery | Query; + ): VectorQuery | Query | AutoQuery; /** * Search the table with a given query vector. * @@ -975,10 +976,11 @@ export class LocalTable extends Table { return this.inner.display(); } - private async getEmbeddingFunctions(): Promise< - Map - > { - const schema = await this.schema(); + private async getEmbeddingFunctions( + inner: _NativeTable = this.inner, + ): Promise> { + const schemaBuf = await inner.schema(); + const schema = tableFromIPC(schemaBuf).schema; const registry = getRegistry(); return registry.parseFunctions(schema.metadata); } @@ -1160,7 +1162,7 @@ export class LocalTable extends Table { query: string | IntoVector | MultiVector | FullTextQuery, queryType: string = "auto", ftsColumns?: string | string[], - ): VectorQuery | Query { + ): VectorQuery | Query | AutoQuery { if (typeof query !== "string" && !instanceOfFullTextQuery(query)) { if (queryType === "fts") { throw new Error("Cannot perform full text search on a vector query"); @@ -1175,17 +1177,35 @@ export class LocalTable extends Table { }); } - // The query type is auto or vector - // fall back to full text search if no embedding functions are defined and the query is a string - if ( - queryType === "auto" && - (getRegistry().length() === 0 || instanceOfFullTextQuery(query)) - ) { + if (queryType === "auto" && typeof query !== "string") { return this.query().fullTextSearch(query, { columns: ftsColumns, }); } + if (queryType === "auto" && typeof query === "string") { + const vector = async (snapshot: _NativeTable) => { + const functions = await this.getEmbeddingFunctions(snapshot); + // TODO: Support multiple embedding functions + const embeddingFunc: EmbeddingFunctionConfig | undefined = functions + .values() + .next().value; + if (embeddingFunc === undefined) { + return undefined; + } + return await embeddingFunc.function.computeQueryEmbeddings(query); + }; + + const columns = + typeof ftsColumns === "string" ? [ftsColumns] : ftsColumns; + return Query.autoSearch( + () => this.inner.checkoutCurrent(), + query, + vector, + columns, + ); + } + const queryPromise = this.getEmbeddingFunctions().then( async (functions) => { // TODO: Support multiple embedding functions diff --git a/nodejs/src/table.rs b/nodejs/src/table.rs index 9d60b2056..ff16ac042 100644 --- a/nodejs/src/table.rs +++ b/nodejs/src/table.rs @@ -554,6 +554,12 @@ impl Table { .default_error() } + #[napi(catch_unwind)] + pub async fn checkout_current(&self) -> napi::Result { + let table = self.inner_ref()?.checkout_current().await.default_error()?; + Ok(Self::new(table)) + } + #[napi(catch_unwind)] pub async fn checkout(&self, version: i64) -> napi::Result<()> { self.inner_ref()? diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 44c92310f..77d5983eb 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -1725,6 +1725,22 @@ impl BaseTable for RemoteTable { async fn version(&self) -> Result { self.describe().await.map(|desc| desc.version) } + + async fn checkout_current(&self) -> Result> { + let description = self.describe().await?; + let TableDescription { + version, + schema, + location, + } = description; + let schema = Arc::new(arrow_schema::Schema::try_from(schema)?); + let snapshot = self.with_branch(self.branch.clone()); + *snapshot.version.write().await = Some(version); + *snapshot.location.write().await = location; + snapshot.schema_cache.seed(schema); + Ok(Arc::new(snapshot)) + } + async fn checkout(&self, version: u64) -> Result<()> { // Validate the version exists. The describe is sent without freshness // headers so a stale `min_version` from a previous write doesn't ride @@ -8739,6 +8755,28 @@ mod tests { } } + /// A pinned snapshot should reuse the version and schema returned by its + /// initial describe instead of issuing two more describe requests. + #[tokio::test] + async fn test_checkout_current_seeds_schema_from_single_describe() { + let describe_calls = Arc::new(AtomicUsize::new(0)); + let calls = describe_calls.clone(); + let table = Table::new_with_handler("my_table", move |request| { + assert_eq!(request.url().path(), "/v1/table/my_table/describe/"); + calls.fetch_add(1, Ordering::SeqCst); + http::Response::builder() + .status(200) + .body( + r#"{"version":42,"schema":{"fields":[{"name":"a","type":{"type":"int32"},"nullable":false}]}}"#, + ) + .unwrap() + }); + + let snapshot = table.checkout_current().await.unwrap(); + assert_eq!(snapshot.schema().await.unwrap().fields().len(), 1); + assert_eq!(describe_calls.load(Ordering::SeqCst), 1); + } + /// Test that schema cache is invalidated after checkout #[tokio::test] async fn test_schema_cache_invalidation_on_checkout() { diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index efc36d260..5007a61d9 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -785,6 +785,12 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync { async fn drop_columns(&self, columns: &[&str]) -> Result; /// Get the version of the table. async fn version(&self) -> Result; + /// Return a new table handle pinned to the exact revision currently visible. + async fn checkout_current(&self) -> Result> { + Err(Error::NotSupported { + message: "checkout_current is not supported on this table type".into(), + }) + } /// Checkout a specific version of the table. async fn checkout(&self, version: u64) -> Result<()>; /// Checkout a table version referenced by a tag. @@ -1944,6 +1950,20 @@ impl Table { self.inner.version().await } + /// Return a new table handle pinned to the exact revision currently visible. + /// + /// This is used when asynchronous preparation must remain consistent with + /// the revision used for a later read. + #[doc(hidden)] + pub async fn checkout_current(&self) -> Result { + let inner = self.inner.checkout_current().await?; + Ok(Self { + inner, + database: self.database.clone(), + embedding_registry: self.embedding_registry.clone(), + }) + } + /// Checks out a specific version of the Table /// /// Any read operation on the table will now access the data at the checked out version. @@ -3043,6 +3063,18 @@ impl BaseTable for NativeTable { Ok(self.dataset.get().await?.version().version) } + async fn checkout_current(&self) -> Result> { + let current = self.dataset.get().await?; + let dataset = dataset::DatasetConsistencyWrapper::new_time_travel( + current.as_ref().clone(), + self.read_consistency_interval, + ); + Ok(Arc::new(Self { + dataset, + ..self.clone() + })) + } + async fn checkout(&self, version: u64) -> Result<()> { self.dataset.as_time_travel(version).await } diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 9b81786cd..d35286c84 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -697,6 +697,7 @@ mod tests { use super::*; use crate::query::{QueryExecutionOptions, QueryRequest}; + use crate::table::BaseTable; fn fixed_size_list_array(values: Vec, dimension: i32) -> FixedSizeListArray { FixedSizeListArray::try_new_from_values(Float32Array::from(values), dimension).unwrap() @@ -889,10 +890,56 @@ mod tests { async fn query_table(&self, _request: NsQueryTableRequest) -> lance::Result { self.query_table_calls.fetch_add(1, Ordering::SeqCst); - panic!("approx_mode queries must not be pushed down to namespace query_table"); + panic!("query must not be pushed down to namespace query_table"); } } + #[tokio::test] + async fn test_execute_query_pinned_snapshot_with_namespace_pushdown_runs_locally() { + use crate::connect; + use arrow_array::{Int32Array, RecordBatch}; + use arrow_schema::{DataType, Field, Schema}; + + let conn = connect("memory://").execute().await.unwrap(); + let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)])); + let batch = RecordBatch::try_new( + schema, + vec![Arc::new(Int32Array::from(vec![1, 2, 3, 4, 5]))], + ) + .unwrap(); + let table = conn + .create_table("test_pinned_namespace_fallback", vec![batch]) + .execute() + .await + .unwrap(); + + let namespace_client = Arc::new(CountingNamespaceClient::default()); + let mut native_table = table.as_native().unwrap().clone(); + native_table.namespace_client = Some(namespace_client.clone()); + native_table + .pushdown_operations + .insert(NamespaceClientPushdownOperation::QueryTable); + + let snapshot = native_table.checkout_current().await.unwrap(); + let snapshot = snapshot.as_any().downcast_ref::().unwrap(); + assert!(snapshot.dataset.time_travel_version().is_some()); + + let query = AnyQuery::Query(QueryRequest { + filter: Some(QueryFilter::Sql("id > 3".to_string())), + ..Default::default() + }); + let stream = execute_query(snapshot, &query, QueryExecutionOptions::default()) + .await + .unwrap(); + let batches = stream.try_collect::>().await.unwrap(); + + assert_eq!( + batches.iter().map(|batch| batch.num_rows()).sum::(), + 2 + ); + assert_eq!(namespace_client.query_table_calls.load(Ordering::SeqCst), 0); + } + #[tokio::test] async fn test_execute_query_approx_mode_with_namespace_pushdown_runs_locally() { use crate::connect; From 302b21aa94317dc74cbe6601212242f54e592d8a Mon Sep 17 00:00:00 2001 From: "lancedb-gatefixer[bot]" <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Wed, 26 Aug 2026 06:42:29 +0800 Subject: [PATCH 19/23] test(node): cover nested PDF metadata queries (#3827) ## Summary - add an end-to-end Node regression matching LangChain PDFLoader metadata - verify create/query round trips rich nested `loc` and `pdf.info` fields against the currently configured Apache Arrow peer ## Root cause LanceDB v0.14 delegated nested object inference to Apache Arrow. Nested strings were dictionary-encoded with colliding dictionary IDs, so serializing query results as an IPC file failed with a dictionary-replacement error. Current `main` recursively infers nested fields and avoids those invalid dictionaries, but the reported LangChain path had no end-to-end regression coverage. ## Validation - `pnpm build` - `pnpm lint` - `pnpm run docs` - `pnpm test --runInBand` (678 passed, 5 skipped) Fixes #1963 --------- Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Co-authored-by: Xuanwo --- nodejs/__test__/table.test.ts | 50 +++++++++++++++++++++++++++++++++++ 1 file changed, 50 insertions(+) diff --git a/nodejs/__test__/table.test.ts b/nodejs/__test__/table.test.ts index 0aee5caf0..aaa144f16 100644 --- a/nodejs/__test__/table.test.ts +++ b/nodejs/__test__/table.test.ts @@ -685,6 +685,56 @@ describe.each([arrow15, arrow16, arrow17, arrow18])( }, ); +// https://github.com/lancedb/lancedb/issues/1963 +it("should query documents with LangChain PDF metadata", async () => { + const tmpDir = tmp.dirSync({ unsafeCleanup: true }); + try { + const db = await connect(tmpDir.name); + const documents = [ + { + text: "first page", + vector: [1, 0], + source: "first.pdf", + loc: { pageNumber: 1, lines: { from: 1, to: 12 } }, + pdf: { + version: "1.10.100", + info: { + format: "PDF 1.7", + producer: "pdf.js", + creator: "Writer", + }, + totalPages: 2, + }, + }, + { + text: "second page", + vector: [0, 1], + source: "second.pdf", + loc: { pageNumber: 2, lines: { from: 13, to: 24 } }, + pdf: { + version: "1.10.100", + info: { + format: "PDF 1.7", + producer: "pdf.js", + creator: "Writer", + }, + totalPages: 2, + }, + }, + ]; + const documentsTable = await db.createTable("documents", documents); + + const results = await documentsTable.query().toArray(); + + expect(results).toHaveLength(2); + expect(results[0].source).toBe("first.pdf"); + expect(results[0].pdf.info.producer).toBe("pdf.js"); + expect(results[1].loc.pageNumber).toBe(2); + } finally { + tmpDir.removeCallback(); + } +}); + describe("merge insert", () => { let tmpDir: tmp.DirResult; let table: Table; From 8083232dd57c556a8b45bd82e8fa0d85e5be26bb Mon Sep 17 00:00:00 2001 From: LanceDB Robot Date: Tue, 25 Aug 2026 16:49:17 -0700 Subject: [PATCH 20/23] chore: update lance dependency to v12.0.0-beta.1 (#4055) Updates the Lance Rust workspace dependencies and Java lance-core dependency to [v12.0.0-beta.1](https://github.com/lance-format/lance/releases/tag/v12.0.0-beta.1). Includes compatibility updates for the renamed shard-manifest API and paginated object-store wrappers. --- Cargo.lock | 84 +++++++++---------- Cargo.toml | 28 +++---- java/pom.xml | 2 +- rust/lancedb/src/io/object_store.rs | 10 ++- .../src/io/object_store/io_tracking.rs | 10 ++- rust/lancedb/src/table.rs | 16 ++++ rust/lancedb/src/table/query/lsm.rs | 2 +- 7 files changed, 92 insertions(+), 60 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 33e43f9fa..63ee13e78 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" [[package]] name = "fsst" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-array", "rand 0.9.5", @@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a" [[package]] name = "lance" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arc-swap", "arrow", @@ -4888,8 +4888,8 @@ dependencies = [ [[package]] name = "lance-arrow" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" 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-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" 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-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-array", "arrow-schema", @@ -4934,8 +4934,8 @@ dependencies = [ [[package]] name = "lance-bitpacking" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrayref", "crunchy", @@ -4945,8 +4945,8 @@ dependencies = [ [[package]] name = "lance-core" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-array", "arrow-buffer", @@ -4983,8 +4983,8 @@ dependencies = [ [[package]] name = "lance-datafusion" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow", "arrow-array", @@ -5013,8 +5013,8 @@ dependencies = [ [[package]] name = "lance-datagen" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow", "arrow-array", @@ -5031,8 +5031,8 @@ dependencies = [ [[package]] name = "lance-derive" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "proc-macro2", "quote", @@ -5041,8 +5041,8 @@ dependencies = [ [[package]] name = "lance-encoding" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-arith", "arrow-array", @@ -5075,8 +5075,8 @@ dependencies = [ [[package]] name = "lance-file" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-arith", "arrow-array", @@ -5107,8 +5107,8 @@ dependencies = [ [[package]] name = "lance-index" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arc-swap", "arrow", @@ -5172,8 +5172,8 @@ dependencies = [ [[package]] name = "lance-index-core" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-array", "arrow-schema", @@ -5195,8 +5195,8 @@ dependencies = [ [[package]] name = "lance-io" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow", "arrow-array", @@ -5236,8 +5236,8 @@ dependencies = [ [[package]] name = "lance-linalg" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-array", "arrow-schema", @@ -5251,8 +5251,8 @@ dependencies = [ [[package]] name = "lance-namespace" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow", "async-trait", @@ -5264,8 +5264,8 @@ dependencies = [ [[package]] name = "lance-namespace-impls" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow", "arrow-ipc", @@ -5318,8 +5318,8 @@ dependencies = [ [[package]] name = "lance-select" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-array", "arrow-buffer", @@ -5333,8 +5333,8 @@ dependencies = [ [[package]] name = "lance-table" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow", "arrow-array", @@ -5374,8 +5374,8 @@ dependencies = [ [[package]] name = "lance-testing" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "arrow-array", "arrow-schema", @@ -5388,8 +5388,8 @@ dependencies = [ [[package]] name = "lance-tokenizer" -version = "11.0.0-beta.22" -source = "git+https://github.com/lance-format/lance.git?tag=v11.0.0-beta.22#ea3cb4d799c468232735e9bcb43959487aca5c20" +version = "12.0.0-beta.1" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" dependencies = [ "frostem", "icu_segmenter", diff --git a/Cargo.toml b/Cargo.toml index 716b14d7a..3f2248e4b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,20 +13,20 @@ categories = ["database-implementations"] rust-version = "1.91.0" [workspace.dependencies] -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" } +lance = { "version" = "=12.0.0-beta.1", default-features = false, "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-core = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-datagen = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-file = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-io = { "version" = "=12.0.0-beta.1", default-features = false, "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-index = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-linalg = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-namespace = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-namespace-impls = { "version" = "=12.0.0-beta.1", default-features = false, "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-table = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-testing = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-datafusion = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-encoding = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance-arrow = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "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 diff --git a/java/pom.xml b/java/pom.xml index 4d83153ab..78aa17bb7 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -28,7 +28,7 @@ UTF-8 15.0.0 - 11.0.0-beta.22 + 12.0.0-beta.1 false 2.30.0 1.7 diff --git a/rust/lancedb/src/io/object_store.rs b/rust/lancedb/src/io/object_store.rs index d594bd857..c4a9a4f7e 100644 --- a/rust/lancedb/src/io/object_store.rs +++ b/rust/lancedb/src/io/object_store.rs @@ -10,7 +10,7 @@ use lance::io::WrappingObjectStore; use object_store::{ CopyOptions, Error, GetOptions, GetResult, ListResult, MultipartUpload, ObjectMeta, ObjectStore, ObjectStoreExt, PutMultipartOptions, PutOptions, PutPayload, PutResult, Result, - UploadPart, path::Path, + UploadPart, list::PaginatedListStore, path::Path, }; use async_trait::async_trait; @@ -187,6 +187,14 @@ impl WrappingObjectStore for MirroringObjectStoreWrapper { secondary: self.secondary.clone(), }) } + + fn wrap_paginated( + &self, + _store_prefix: &str, + original: Arc, + ) -> Option> { + Some(original) + } } // windows pathing can't be simply concatenated diff --git a/rust/lancedb/src/io/object_store/io_tracking.rs b/rust/lancedb/src/io/object_store/io_tracking.rs index bd4f8f54a..7f9750216 100644 --- a/rust/lancedb/src/io/object_store/io_tracking.rs +++ b/rust/lancedb/src/io/object_store/io_tracking.rs @@ -12,7 +12,7 @@ use lance::io::WrappingObjectStore; use object_store::{ CopyOptions, GetOptions, GetResult, ListResult, MultipartUpload, ObjectMeta, ObjectStore, PutMultipartOptions, PutOptions, PutPayload, PutResult, RenameOptions, Result as OSResult, - UploadPart, path::Path, + UploadPart, list::PaginatedListStore, path::Path, }; #[derive(Debug, Default)] @@ -57,6 +57,14 @@ impl WrappingObjectStore for IoStatsHolder { stats: self.0.clone(), }) } + + fn wrap_paginated( + &self, + _store_prefix: &str, + original: Arc, + ) -> Option> { + Some(original) + } } impl IoTrackingStore { diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index 5007a61d9..e544607f4 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -4118,6 +4118,14 @@ mod tests { parent_list_calls: self.parent_list_calls.clone(), }) } + + fn wrap_paginated( + &self, + _store_prefix: &str, + _original: Arc, + ) -> Option> { + None + } } #[tokio::test] @@ -4221,6 +4229,14 @@ mod tests { self.called.store(true, Ordering::Relaxed); original } + + fn wrap_paginated( + &self, + _store_prefix: &str, + original: Arc, + ) -> Option> { + Some(original) + } } #[tokio::test] diff --git a/rust/lancedb/src/table/query/lsm.rs b/rust/lancedb/src/table/query/lsm.rs index 255d649b1..5c340cc35 100644 --- a/rust/lancedb/src/table/query/lsm.rs +++ b/rust/lancedb/src/table/query/lsm.rs @@ -298,7 +298,7 @@ async fn build_read_context( for shard_id in shard_ids { let manifest_store = ShardManifestStore::new(store.clone(), &base_path, shard_id, scan_batch_size); - if let Some(manifest) = manifest_store.read_latest().await? { + if let Some(manifest) = manifest_store.latest().await? { snapshots.push(snapshot_from_manifest(shard_id, &manifest, &exclude)); } } From 9b825c5f298f45001afd1837bd360682c7282ecc Mon Sep 17 00:00:00 2001 From: "lancedb-gatefixer[bot]" <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Wed, 26 Aug 2026 10:08:48 +0800 Subject: [PATCH 21/23] fix(node): route auto search using table embeddings (#3832) ## Summary - Resolve automatic string-search routing from the active table schema whenever the query executes. - Defer embedding-provider construction while leaving explicit vector and FTS routes unchanged. - Cover unrelated global registrations and metadata transitions across repeated executions of one query builder. ## Root cause LocalTable.search used the number of globally registered embedding providers to choose between vector and full-text search. A provider registered for any other table therefore sent a plain FTS table down the vector path. A wrapper-lifetime metadata snapshot avoided that contamination but became stale after time travel or read-consistency refreshes. The query now records fluent builder operations and creates the appropriate native vector or FTS query from the active schema on each execution. ## Validation - pnpm build - pnpm tsc - pnpm lint - pnpm run docs - pnpm test --runInBand (681 passed, 5 skipped) Fixes #1557 --------- Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Co-authored-by: Xuanwo --- nodejs/__test__/table.test.ts | 343 +++++++++++++++++++++++++++++- nodejs/lancedb/query.ts | 224 ++++++++++--------- nodejs/lancedb/table.ts | 39 ++-- nodejs/src/table.rs | 7 + rust/lancedb/src/remote/table.rs | 14 ++ rust/lancedb/src/table.rs | 28 +++ rust/lancedb/src/table/dataset.rs | 69 +++++- rust/lancedb/src/table/merge.rs | 38 ++++ rust/lancedb/src/table/query.rs | 31 +++ 9 files changed, 672 insertions(+), 121 deletions(-) diff --git a/nodejs/__test__/table.test.ts b/nodejs/__test__/table.test.ts index aaa144f16..dae640850 100644 --- a/nodejs/__test__/table.test.ts +++ b/nodejs/__test__/table.test.ts @@ -2585,7 +2585,24 @@ describe.each([arrow15, arrow16, arrow17, arrow18])( ); }); - test("full text search if no embedding function provided", async () => { + test("full text search if only an unrelated embedding function is registered", async () => { + register("unused")( + class extends EmbeddingFunction { + ndims() { + return 3; + } + embeddingDataType() { + return new Float32(); + } + async computeQueryEmbeddings(_data: string) { + return [1, 2, 3]; + } + async computeSourceEmbeddings(data: string[]) { + return data.map(() => [1, 2, 3]); + } + }, + ); + const db = await connect(tmpDir.name); const data = [ { text: "hello world", vector: [0.1, 0.2, 0.3] }, @@ -2607,6 +2624,306 @@ describe.each([arrow15, arrow16, arrow17, arrow18])( expect(results2[0].text).toBe(data[1].text); }); + test("auto search stays consistent with the active revision", async () => { + let initCalls = 0; + let queryCalls = 0; + let markStarted!: () => void; + const started = new Promise((resolve) => { + markStarted = resolve; + }); + let releaseEmbedding!: () => void; + const embeddingReleased = new Promise((resolve) => { + releaseEmbedding = resolve; + }); + + @register("refresh-test") + class TestEmbedding extends EmbeddingFunction { + async init() { + initCalls += 1; + } + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings(value: string) { + queryCalls += 1; + if (value === "blocked") { + markStarted(); + await embeddingReleased; + } + return value === "greetings" ? [0.1] : [0.2]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map((value) => + value === "hello world" ? [0.1] : [0.2], + ); + } + } + + const writer = await connect(tmpDir.name); + await writer.createTable("test", [{ text: "plain", vector: [0.0] }]); + const reader = await connect(tmpDir.name, { + readConsistencyInterval: 0, + }); + const tracked = await reader.openTable("test"); + type SnapshotCountingNative = { + querySnapshot: () => Promise; + }; + const native = (tracked as unknown as { inner: SnapshotCountingNative }) + .inner; + const querySnapshot = native.querySnapshot.bind(native); + let snapshotCalls = 0; + native.querySnapshot = async () => { + snapshotCalls += 1; + return await querySnapshot(); + }; + const autoQuery = tracked.search("greetings").select(["text"]).limit(1); + + const func = new TestEmbedding(); + const schema = LanceSchema({ + text: func.sourceField(new arrow.Utf8()), + vector: func.vectorField(), + }); + const data = [{ text: "hello world" }, { text: "goodbye world" }]; + await writer.createTable("test", data, { mode: "overwrite", schema }); + const baselineInitCalls = initCalls; + + expect( + (await tracked.schema()).metadata.get("embedding_functions"), + ).toBeDefined(); + const results = await autoQuery.toArray(); + expect(results[0].text).toBe(data[0].text); + expect(initCalls).toBe(baselineInitCalls + 1); + expect(queryCalls).toBe(1); + expect(snapshotCalls).toBe(1); + + const repeatedResults = await autoQuery.toArray(); + expect(repeatedResults[0].text).toBe(data[0].text); + expect(initCalls).toBe(baselineInitCalls + 1); + expect(queryCalls).toBe(1); + expect(snapshotCalls).toBe(2); + + const pending = tracked + .search("blocked") + .select(["text"]) + .limit(1) + .toArray(); + await started; + + const ftsData = [ + { text: "greetings from full text", vector: [0.0] }, + { text: "blocked from full text", vector: [0.0] }, + ]; + const ftsTable = await writer.createTable("test", ftsData, { + mode: "overwrite", + }); + await ftsTable.createIndex("text", { config: Index.fts() }); + releaseEmbedding(); + + const pendingResults = await pending; + expect(pendingResults[0].text).toBe(data[1].text); + + expect( + (await tracked.schema()).metadata.get("embedding_functions"), + ).toBeUndefined(); + const ftsResults = await autoQuery.toArray(); + expect(ftsResults[0].text).toBe(ftsData[0].text); + }); + + test("auto search keeps newer preparation during a revision race", async () => { + let aCalls = 0; + let bCalls = 0; + let markAStarted!: () => void; + const aStarted = new Promise((resolve) => { + markAStarted = resolve; + }); + let releaseA!: () => void; + const aReleased = new Promise((resolve) => { + releaseA = resolve; + }); + let markBStarted!: () => void; + const bStarted = new Promise((resolve) => { + markBStarted = resolve; + }); + let releaseB!: () => void; + const bReleased = new Promise((resolve) => { + releaseB = resolve; + }); + + @register("race-a") + class EmbeddingA extends EmbeddingFunction { + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings() { + aCalls += 1; + markAStarted(); + await aReleased; + return [0.1]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map(() => [0.1]); + } + } + + @register("race-b") + class EmbeddingB extends EmbeddingFunction { + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings() { + bCalls += 1; + markBStarted(); + await bReleased; + return [0.2]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map(() => [0.2]); + } + } + + const writer = await connect(tmpDir.name); + const embeddingA = new EmbeddingA(); + const schemaA = LanceSchema({ + text: embeddingA.sourceField(new arrow.Utf8()), + vector: embeddingA.vectorField(), + }); + await writer.createTable("race", [{ text: "revision a" }], { + schema: schemaA, + }); + const reader = await connect(tmpDir.name, { + readConsistencyInterval: 0, + }); + const tracked = await reader.openTable("race"); + const query = tracked.search("query"); + + const first = query.toArray(); + await aStarted; + + const embeddingB = new EmbeddingB(); + const schemaB = LanceSchema({ + text: embeddingB.sourceField(new arrow.Utf8()), + vector: embeddingB.vectorField(), + }); + await writer.createTable("race", [{ text: "revision b" }], { + mode: "overwrite", + schema: schemaB, + }); + const second = query.toArray(); + await bStarted; + + releaseA(); + releaseB(); + await Promise.all([first, second]); + expect(aCalls).toBe(1); + expect(bCalls).toBe(1); + }); + + test("stale FTS routing keeps newer vector preparation", async () => { + let vectorCalls = 0; + let markVectorStarted!: () => void; + const vectorStarted = new Promise((resolve) => { + markVectorStarted = resolve; + }); + let releaseVector!: () => void; + const vectorReleased = new Promise((resolve) => { + releaseVector = resolve; + }); + + @register("stale-fts-race") + class RaceEmbedding extends EmbeddingFunction { + ndims() { + return 1; + } + embeddingDataType() { + return new arrow.Float32(); + } + async computeQueryEmbeddings() { + vectorCalls += 1; + markVectorStarted(); + await vectorReleased; + return [0.1]; + } + async computeSourceEmbeddings(values: string[]) { + return values.map(() => [0.1]); + } + } + + const writer = await connect(tmpDir.name); + const ftsTable = await writer.createTable("stale_fts", [ + { text: "hello", vector: [0.0] }, + ]); + await ftsTable.createIndex("text", { config: Index.fts() }); + + const reader = await connect(tmpDir.name, { + readConsistencyInterval: 0, + }); + const tracked = await reader.openTable("stale_fts"); + type Snapshot = { + schema: () => Promise; + }; + type NativeWithSnapshot = { + querySnapshot: () => Promise; + }; + const native = (tracked as unknown as { inner: NativeWithSnapshot }) + .inner; + const querySnapshot = native.querySnapshot.bind(native); + let snapshotCalls = 0; + let markStaleSchemaStarted!: () => void; + const staleSchemaStarted = new Promise((resolve) => { + markStaleSchemaStarted = resolve; + }); + let releaseStaleSchema!: () => void; + const staleSchemaReleased = new Promise((resolve) => { + releaseStaleSchema = resolve; + }); + native.querySnapshot = async () => { + const snapshot = await querySnapshot(); + snapshotCalls += 1; + if (snapshotCalls === 1) { + const schema = snapshot.schema.bind(snapshot); + snapshot.schema = async () => { + markStaleSchemaStarted(); + await staleSchemaReleased; + return await schema(); + }; + } + return snapshot; + }; + + const query = tracked.search("hello"); + const staleFtsExecution = query.toArray(); + await staleSchemaStarted; + + const embedding = new RaceEmbedding(); + const vectorSchema = LanceSchema({ + text: embedding.sourceField(new arrow.Utf8()), + vector: embedding.vectorField(), + }); + await writer.createTable("stale_fts", [{ text: "hello" }], { + mode: "overwrite", + schema: vectorSchema, + }); + + const vectorExecution = query.toArray(); + await vectorStarted; + releaseStaleSchema(); + await staleFtsExecution; + releaseVector(); + await vectorExecution; + + await query.toArray(); + expect(vectorCalls).toBe(1); + }); + test("tokenizes FTS queries by column or index name", async () => { const db = await connect(tmpDir.name); const data = [ @@ -3157,6 +3474,30 @@ describe("column name options", () => { expect(results[1].query_index).toBe(1); }); + test("observes promised additional vectors while the query is pending", async () => { + const initialVector = new Promise(() => undefined); + const query = table.query().nearestTo(initialVector); + const unhandled: unknown[] = []; + const onUnhandled = (reason: unknown) => unhandled.push(reason); + process.on("unhandledRejection", onUnhandled); + + try { + query.addQueryVector(Promise.reject(new Error("extra vector failed"))); + await new Promise((resolve) => setImmediate(resolve)); + expect(unhandled).toEqual([]); + + const rejectedQuery = table + .query() + .nearestTo([0.1, 0.2]) + .addQueryVector(Promise.reject(new Error("consumed vector failed"))); + await expect(rejectedQuery.toArray()).rejects.toThrow( + "consumed vector failed", + ); + } finally { + process.off("unhandledRejection", onUnhandled); + } + }); + test("index and search multivectors", async () => { const db = await connect(tmpDir.name); const data = []; diff --git a/nodejs/lancedb/query.ts b/nodejs/lancedb/query.ts index 3b9b286a0..f1d31eae1 100644 --- a/nodejs/lancedb/query.ts +++ b/nodejs/lancedb/query.ts @@ -100,6 +100,29 @@ export interface FullTextSearchOptions { columns?: string | string[]; } +function nearestToNative( + inner: NativeQuery, + vector: Awaited, +): NativeVectorQuery { + const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector); + if (raw) { + return inner.nearestToRaw(raw.data, raw.dtype); + } + return inner.nearestTo(Float32Array.from(vector as number[])); +} + +function addQueryVectorToNative( + inner: NativeVectorQuery, + vector: Awaited, +) { + const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector); + if (raw) { + inner.addQueryVectorRaw(raw.data, raw.dtype); + } else { + inner.addQueryVector(Float32Array.from(vector as number[])); + } +} + /** Common methods supported by all query types * * @see {@link Query} @@ -499,6 +522,13 @@ export class VectorQuery extends StandardQueryBase { super(inner); } + /** + * @hidden + */ + protected doVectorCall(fn: (inner: NativeVectorQuery) => void) { + super.doCall(fn); + } + /** * Set the number of partitions to search (probe) * @@ -526,7 +556,7 @@ export class VectorQuery extends StandardQueryBase { * the minimum and maximum to the same value. */ nprobes(nprobes: number): VectorQuery { - super.doCall((inner) => inner.nprobes(nprobes)); + this.doVectorCall((inner) => inner.nprobes(nprobes)); return this; } @@ -540,7 +570,7 @@ export class VectorQuery extends StandardQueryBase { * but will also increase latency. */ minimumNprobes(minimumNprobes: number): VectorQuery { - super.doCall((inner) => inner.minimumNprobes(minimumNprobes)); + this.doVectorCall((inner) => inner.minimumNprobes(minimumNprobes)); return this; } @@ -554,7 +584,7 @@ export class VectorQuery extends StandardQueryBase { * potential false negatives. */ maximumNprobes(maximumNprobes: number): VectorQuery { - super.doCall((inner) => inner.maximumNprobes(maximumNprobes)); + this.doVectorCall((inner) => inner.maximumNprobes(maximumNprobes)); return this; } @@ -567,7 +597,7 @@ export class VectorQuery extends StandardQueryBase { * `undefined` means no lower or upper bound. */ distanceRange(lowerBound?: number, upperBound?: number): VectorQuery { - super.doCall((inner) => inner.distanceRange(lowerBound, upperBound)); + this.doVectorCall((inner) => inner.distanceRange(lowerBound, upperBound)); return this; } @@ -581,7 +611,7 @@ export class VectorQuery extends StandardQueryBase { * also increase the latency of your query. The default value is 1.5*limit. */ ef(ef: number): VectorQuery { - super.doCall((inner) => inner.ef(ef)); + this.doVectorCall((inner) => inner.ef(ef)); return this; } @@ -595,7 +625,7 @@ export class VectorQuery extends StandardQueryBase { * whose data type is a fixed-size-list of floats. */ column(column: string): VectorQuery { - super.doCall((inner) => inner.column(column)); + this.doVectorCall((inner) => inner.column(column)); return this; } @@ -616,7 +646,7 @@ export class VectorQuery extends StandardQueryBase { distanceType( distanceType: Required["distanceType"], ): VectorQuery { - super.doCall((inner) => inner.distanceType(distanceType)); + this.doVectorCall((inner) => inner.distanceType(distanceType)); return this; } @@ -650,7 +680,7 @@ export class VectorQuery extends StandardQueryBase { * distance between the query vector and the actual uncompressed vector. */ refineFactor(refineFactor: number): VectorQuery { - super.doCall((inner) => inner.refineFactor(refineFactor)); + this.doVectorCall((inner) => inner.refineFactor(refineFactor)); return this; } @@ -675,7 +705,7 @@ export class VectorQuery extends StandardQueryBase { * factor can often help restore some of the results lost by post filtering. */ postfilter(): VectorQuery { - super.doCall((inner) => inner.postfilter()); + this.doVectorCall((inner) => inner.postfilter()); return this; } @@ -689,7 +719,7 @@ export class VectorQuery extends StandardQueryBase { * calculate your recall to select an appropriate value for nprobes. */ bypassVectorIndex(): VectorQuery { - super.doCall((inner) => inner.bypassVectorIndex()); + this.doVectorCall((inner) => inner.bypassVectorIndex()); return this; } @@ -705,35 +735,31 @@ export class VectorQuery extends StandardQueryBase { */ addQueryVector(vector: IntoVector): VectorQuery { if (vector instanceof Promise) { + // Observe the promise as soon as it is accepted. The existing native + // query may still be pending, and delaying observation until it resolves + // can otherwise surface a fast rejection as unhandled. + const settledVector = vector.then( + (value) => ({ status: "fulfilled" as const, value }), + (reason) => ({ status: "rejected" as const, reason }), + ); const res = (async () => { - try { - const v = await vector; - // biome-ignore lint/suspicious/noExplicitAny: we need to get the `inner`, but js has no package scoping - const value: any = this.addQueryVector(v); - const inner = value.inner as - | NativeVectorQuery - | Promise; - return inner; - } catch (e) { - return Promise.reject(e); + const inner = await this.getInner(); + const outcome = await settledVector; + if (outcome.status === "rejected") { + throw outcome.reason; } + addQueryVectorToNative(inner, outcome.value); + return inner; })(); return new VectorQuery(res); } else { - super.doCall((inner) => { - const raw = Array.isArray(vector) ? null : extractVectorBuffer(vector); - if (raw) { - inner.addQueryVectorRaw(raw.data, raw.dtype); - } else { - inner.addQueryVector(Float32Array.from(vector as number[])); - } - }); + this.doVectorCall((inner) => addQueryVectorToNative(inner, vector)); return this; } } rerank(reranker: Reranker): VectorQuery { - super.doCall((inner) => + this.doVectorCall((inner) => inner.rerank(async (args) => { const vecResults = await fromBufferToRecordBatch(args.vecResults); const ftsResults = await fromBufferToRecordBatch(args.ftsResults); @@ -752,6 +778,71 @@ export class VectorQuery extends StandardQueryBase { } } +/** + * Create a string query whose vector/FTS routing is resolved against the active + * table schema when the query executes. + * + * @hidden + */ +export function createAutoQuery( + table: NativeTable, + query: string, + columns: string[] | null, + getVector: (metadata: string) => Promise>, +): AutoQuery { + type RouteSnapshot = { + table: NativeTable; + embeddingMetadata: string | undefined; + }; + type CachedPreparation = { + metadata: string; + vector: Promise>; + }; + + let cachedPreparation: CachedPreparation | undefined; + + const snapshotRoute = async (): Promise => { + const snapshot = await table.querySnapshot(); + const schema = tableFromIPC(await snapshot.schema()).schema; + return { + table: snapshot, + embeddingMetadata: schema.metadata.get("embedding_functions"), + }; + }; + + const createInner = async (): Promise => { + const route = await snapshotRoute(); + if (route.embeddingMetadata === undefined) { + const inner = route.table.query(); + inner.fullTextSearch({ query, columns }); + return inner; + } + + const metadata = route.embeddingMetadata; + if (cachedPreparation?.metadata !== metadata) { + cachedPreparation = { + metadata, + vector: Promise.resolve().then(() => getVector(metadata)), + }; + } + + const preparation = cachedPreparation; + let vector: Awaited; + try { + vector = await preparation.vector; + } catch (error) { + if (cachedPreparation === preparation) { + cachedPreparation = undefined; + } + throw error; + } + + return nearestToNative(route.table.query(), vector); + }; + + return new AutoQuery(createInner); +} + /** * A query that returns a subset of the rows in the table. * @@ -836,37 +927,6 @@ export class Query extends StandardQueryBase { super(tbl.query()); } - /** @hidden */ - static autoSearch( - tbl: () => Promise, - query: string, - vector: (tbl: NativeTable) => Promise | undefined>, - columns?: string[], - ): AutoQuery { - const nativeQuery = async () => { - const snapshot = await Promise.resolve(tbl()); - const resolved = await vector(snapshot); - const inner = snapshot.query(); - if (resolved === undefined) { - inner.fullTextSearch({ - query, - columns: columns ?? null, - }); - return inner; - } - - const raw = Array.isArray(resolved) - ? null - : extractVectorBuffer(resolved); - if (raw) { - return inner.nearestToRaw(raw.data, raw.dtype); - } - return inner.nearestTo(Float32Array.from(resolved as number[])); - }; - - return new AutoQuery(nativeQuery); - } - /** * Find the nearest vectors to the given query vector. * @@ -905,45 +965,19 @@ export class Query extends StandardQueryBase { * a default `limit` of 10 will be used. @see {@link Query#limit} */ nearestTo(vector: IntoVector): VectorQuery { - const callNearestTo = ( - inner: NativeQuery, - resolved: Float32Array | Float64Array | Uint8Array | number[], - ): NativeVectorQuery => { - const raw = Array.isArray(resolved) - ? null - : extractVectorBuffer(resolved); - if (raw) { - return inner.nearestToRaw(raw.data, raw.dtype); - } - return inner.nearestTo(Float32Array.from(resolved as number[])); - }; - - if (this.inner instanceof Promise) { - const nativeQuery = this.inner.then(async (inner) => { - const resolved = vector instanceof Promise ? await vector : vector; - return callNearestTo(inner, resolved); - }); + const inner = this.inner; + if (inner instanceof Promise) { + const nativeQuery = inner.then(async (resolvedInner) => + nearestToNative(resolvedInner, await vector), + ); return new VectorQuery(nativeQuery); } if (vector instanceof Promise) { - const res = (async () => { - try { - const v = await vector; - // biome-ignore lint/suspicious/noExplicitAny: we need to get the `inner`, but js has no package scoping - const value: any = this.nearestTo(v); - const inner = value.inner as - | NativeVectorQuery - | Promise; - return inner; - } catch (e) { - return Promise.reject(e); - } - })(); - return new VectorQuery(res); - } else { - const vectorQuery = callNearestTo(this.inner, vector); - return new VectorQuery(vectorQuery); + return new VectorQuery( + vector.then((resolvedVector) => nearestToNative(inner, resolvedVector)), + ); } + return new VectorQuery(nearestToNative(inner, vector)); } nearestToText(query: string | FullTextQuery, columns?: string[]): Query { diff --git a/nodejs/lancedb/table.ts b/nodejs/lancedb/table.ts index a4fc76ef1..82f23e2f2 100644 --- a/nodejs/lancedb/table.ts +++ b/nodejs/lancedb/table.ts @@ -48,6 +48,7 @@ import { Query, TakeQuery, VectorQuery, + createAutoQuery, instanceOfFullTextQuery, } from "./query"; import { sanitizeType } from "./sanitize"; @@ -1177,33 +1178,29 @@ export class LocalTable extends Table { }); } - if (queryType === "auto" && typeof query !== "string") { - return this.query().fullTextSearch(query, { - columns: ftsColumns, - }); - } + if (queryType === "auto") { + if (instanceOfFullTextQuery(query)) { + return this.query().fullTextSearch(query, { + columns: ftsColumns, + }); + } - if (queryType === "auto" && typeof query === "string") { - const vector = async (snapshot: _NativeTable) => { - const functions = await this.getEmbeddingFunctions(snapshot); + const columns = + typeof ftsColumns === "string" ? [ftsColumns] : (ftsColumns ?? null); + return createAutoQuery(this.inner, query, columns, async (metadata) => { + const functions = await getRegistry().parseFunctions( + new Map([["embedding_functions", metadata]]), + ); // TODO: Support multiple embedding functions const embeddingFunc: EmbeddingFunctionConfig | undefined = functions .values() .next().value; - if (embeddingFunc === undefined) { - return undefined; - } + // The route only calls this callback when embedding metadata exists. + // parseFunctions either yields a provider or reports malformed metadata. + if (!embeddingFunc) + throw new Error("Invalid embedding function metadata"); return await embeddingFunc.function.computeQueryEmbeddings(query); - }; - - const columns = - typeof ftsColumns === "string" ? [ftsColumns] : ftsColumns; - return Query.autoSearch( - () => this.inner.checkoutCurrent(), - query, - vector, - columns, - ); + }); } const queryPromise = this.getEmbeddingFunctions().then( diff --git a/nodejs/src/table.rs b/nodejs/src/table.rs index ff16ac042..db74d38fa 100644 --- a/nodejs/src/table.rs +++ b/nodejs/src/table.rs @@ -278,6 +278,13 @@ impl Table { Ok(Query::new(self.inner_ref()?.query())) } + /// Return a read-only table handle pinned to the current query revision. + #[napi(catch_unwind)] + pub async fn query_snapshot(&self) -> napi::Result { + let snapshot = self.inner_ref()?.query_snapshot().await.default_error()?; + Ok(Self::new(snapshot)) + } + #[napi(catch_unwind)] pub fn take_offsets(&self, offsets: Vec) -> napi::Result { Ok(TakeQuery::new( diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 77d5983eb..b116083c8 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -1722,6 +1722,20 @@ impl BaseTable for RemoteTable { fn id(&self) -> &str { &self.identifier } + async fn query_snapshot(&self) -> Result> { + let description = self.describe().await?; + let TableDescription { + version, + schema, + location, + } = description; + let schema = Arc::new(arrow_schema::Schema::try_from(schema)?); + let snapshot = self.with_branch(self.branch.clone()); + *snapshot.version.write().await = Some(version); + *snapshot.location.write().await = location; + snapshot.schema_cache.seed(schema); + Ok(Arc::new(snapshot)) + } async fn version(&self) -> Result { self.describe().await.map(|desc| desc.version) } diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index e544607f4..6b5f687c7 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -560,6 +560,13 @@ pub trait BaseTable: std::fmt::Display + std::fmt::Debug + Send + Sync { fn id(&self) -> &str; /// Get the arrow [Schema] of the table. async fn schema(&self) -> Result; + /// Create a read-only handle pinned to the table's current active revision. + /// + /// The returned handle is independent from later refreshes or checkouts on + /// this handle. This is used by bindings that must prepare client-side + /// query state from the same revision that the query will execute against. + #[doc(hidden)] + async fn query_snapshot(&self) -> Result>; /// Count the number of rows in this table. async fn count_rows(&self, filter: Option) -> Result; /// Create a physical plan for the query. @@ -1139,6 +1146,16 @@ impl Table { self.inner.schema().await } + /// Create a read-only handle pinned to the current active revision. + #[doc(hidden)] + pub async fn query_snapshot(&self) -> Result { + Ok(Self { + inner: self.inner.query_snapshot().await?, + database: self.database.clone(), + embedding_registry: self.embedding_registry.clone(), + }) + } + /// Count the number of rows in this dataset. /// /// # Arguments @@ -3059,6 +3076,17 @@ impl BaseTable for NativeTable { &self.id } + async fn query_snapshot(&self) -> Result> { + let snapshot = self.dataset.new_query_snapshot().await?; + let mut table = self.with_dataset(snapshot); + // QueryTable requests do not carry a revision. A pinned snapshot must + // execute locally until the namespace API can accept that revision. + table + .pushdown_operations + .remove(&NamespaceClientPushdownOperation::QueryTable); + Ok(Arc::new(table)) + } + async fn version(&self) -> Result { Ok(self.dataset.get().await?.version().version) } diff --git a/rust/lancedb/src/table/dataset.rs b/rust/lancedb/src/table/dataset.rs index e37c3fc9f..5e3733b85 100644 --- a/rust/lancedb/src/table/dataset.rs +++ b/rust/lancedb/src/table/dataset.rs @@ -32,6 +32,10 @@ struct DatasetState { /// `Some(version)` = pinned to a specific version (time travel), /// `None` = tracking latest. pinned_version: Option, + /// Whether the pin is an internal query snapshot rather than user-visible + /// time travel. Query snapshots remain read-only but preserve MemWAL read + /// semantics. + query_snapshot: bool, } #[derive(Debug, Clone)] @@ -70,6 +74,7 @@ impl DatasetConsistencyWrapper { state: Arc::new(Mutex::new(DatasetState { dataset, pinned_version: None, + query_snapshot: false, })), consistency, shard_writer: Arc::new(ShardWriterCache::default()), @@ -93,6 +98,36 @@ impl DatasetConsistencyWrapper { wrapper } + /// Create an independent read-only wrapper pinned to the current dataset + /// while retaining this wrapper's live MemWAL read context. + pub async fn new_query_snapshot(&self) -> Result { + // Apply the configured consistency policy before taking the snapshot. + // The returned dataset is intentionally discarded: a checkout may race + // after this await, so the dataset and its pin provenance must instead + // be cloned together from one authoritative state sample below. + self.get().await?; + + let (dataset, query_snapshot) = { + let state = self.state.lock()?; + // Preserve user time travel so the MemWAL safety guard still sees + // it. Latest and already-internal snapshots remain internal pins. + ( + state.dataset.clone(), + state.query_snapshot || state.pinned_version.is_none(), + ) + }; + let version = dataset.version().version; + Ok(Self { + state: Arc::new(Mutex::new(DatasetState { + dataset, + pinned_version: Some(version), + query_snapshot, + })), + consistency: ConsistencyMode::Lazy, + shard_writer: self.shard_writer.clone(), + }) + } + /// The MemWAL `ShardWriter` cache co-located with this dataset. pub(crate) fn shard_writer(&self) -> &Arc { &self.shard_writer @@ -169,6 +204,7 @@ impl DatasetConsistencyWrapper { let mut state = self.state.lock()?; state.dataset = Arc::new(new_dataset); state.pinned_version = None; + state.query_snapshot = false; drop(state); if let ConsistencyMode::Eventual(bg_cache) = &self.consistency { bg_cache.invalidate(); @@ -202,10 +238,10 @@ impl DatasetConsistencyWrapper { /// Returns the version, if in time travel mode, or None otherwise. pub fn time_travel_version(&self) -> Option { - self.state - .lock() - .unwrap_or_else(|e| e.into_inner()) - .pinned_version + let state = self.state.lock().unwrap_or_else(|e| e.into_inner()); + (!state.query_snapshot) + .then_some(state.pinned_version) + .flatten() } /// Convert into a wrapper in latest version mode. @@ -225,6 +261,7 @@ impl DatasetConsistencyWrapper { if state.pinned_version.is_some() { state.dataset = Arc::new(new_dataset); state.pinned_version = None; + state.query_snapshot = false; } drop(state); if let ConsistencyMode::Eventual(bg_cache) = &self.consistency { @@ -260,6 +297,7 @@ impl DatasetConsistencyWrapper { let mut state = self.state.lock()?; state.dataset = Arc::new(new_dataset); state.pinned_version = Some(version_value); + state.query_snapshot = false; Ok(()) } @@ -461,6 +499,29 @@ mod tests { assert_eq!(wrapper.time_travel_version(), Some(1)); } + #[tokio::test] + async fn test_query_snapshot_samples_dataset_and_pin_together() { + let dir = tempfile::tempdir().unwrap(); + let uri = dir.path().to_str().unwrap(); + let ds = create_test_dataset(uri).await; + + let wrapper = DatasetConsistencyWrapper::new_latest(ds, None); + wrapper.as_time_travel(1u64).await.unwrap(); + let stale_time_travel_dataset = wrapper.get().await.unwrap(); + + append_to_dataset(uri).await; + wrapper.as_latest().await.unwrap(); + + let snapshot = wrapper.new_query_snapshot().await.unwrap(); + let snapshot_dataset = snapshot.get().await.unwrap(); + assert_eq!(snapshot_dataset.version().version, 2); + assert_ne!( + snapshot_dataset.version().version, + stale_time_travel_dataset.version().version + ); + assert_eq!(snapshot.time_travel_version(), None); + } + #[tokio::test] async fn test_as_latest_from_time_travel() { let dir = tempfile::tempdir().unwrap(); diff --git a/rust/lancedb/src/table/merge.rs b/rust/lancedb/src/table/merge.rs index b9dd5732b..3227e3edf 100644 --- a/rust/lancedb/src/table/merge.rs +++ b/rust/lancedb/src/table/merge.rs @@ -1056,6 +1056,44 @@ mod lsm_tests { ); } + #[tokio::test] + async fn query_snapshot_preserves_lsm_read_semantics() { + let dir = tempdir().unwrap(); + let table = id_value_table(&dir).await; + table + .set_lsm_write_spec(LsmWriteSpec::unsharded()) + .await + .unwrap(); + lsm_upsert(&table, vec![4, 5]).await; + + let snapshot = table.query_snapshot().await.unwrap(); + let rows = collect_id_value(snapshot.query().execute().await.unwrap()).await; + assert_eq!( + rows.iter().map(|(id, _)| *id).collect::>(), + vec![1, 2, 3, 4, 5] + ); + } + + #[tokio::test] + async fn query_snapshot_preserves_time_travel_lsm_guard() { + let dir = tempdir().unwrap(); + let table = id_value_table(&dir).await; + table + .set_lsm_write_spec(LsmWriteSpec::unsharded()) + .await + .unwrap(); + lsm_upsert(&table, vec![4]).await; + + let version = table.version().await.unwrap(); + table.checkout(version).await.unwrap(); + let direct_error = table.query().execute().await.err().unwrap(); + assert!(matches!(direct_error, Error::NotSupported { .. })); + + let snapshot = table.query_snapshot().await.unwrap(); + let snapshot_error = snapshot.query().execute().await.err().unwrap(); + assert!(matches!(snapshot_error, Error::NotSupported { .. })); + } + #[tokio::test] async fn lsm_read_dedup_newest_wins() { let dir = tempdir().unwrap(); diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index d35286c84..629cb4e6f 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -1056,6 +1056,37 @@ mod tests { assert_eq!(namespace_client.query_table_calls.load(Ordering::SeqCst), 0); } + #[tokio::test] + async fn test_query_snapshot_disables_namespace_pushdown() { + use crate::connect; + use crate::table::BaseTable; + use arrow_array::{Int32Array, RecordBatch}; + use arrow_schema::{DataType, Field, Schema}; + + let conn = connect("memory://").execute().await.unwrap(); + let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)])); + let batch = + RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 2, 3]))]).unwrap(); + let table = conn + .create_table("test_snapshot_namespace_fallback", vec![batch]) + .execute() + .await + .unwrap(); + let mut native_table = table.as_native().unwrap().clone(); + native_table.namespace_client = Some(Arc::new(CountingNamespaceClient::default())); + native_table + .pushdown_operations + .insert(NamespaceClientPushdownOperation::QueryTable); + + let snapshot = BaseTable::query_snapshot(&native_table).await.unwrap(); + let snapshot = snapshot.as_any().downcast_ref::().unwrap(); + assert!( + !can_execute_namespace_query(snapshot, &AnyQuery::Query(QueryRequest::default()),) + .await + .unwrap() + ); + } + #[tokio::test] async fn test_create_plan_multivector_structure() { use arrow_array::{Float32Array, RecordBatch}; From 21530432a06ddd3814ec8b0e617abe3eda041fc3 Mon Sep 17 00:00:00 2001 From: LanceDB Robot Date: Tue, 25 Aug 2026 19:56:26 -0700 Subject: [PATCH 22/23] chore: update lance dependency to v12.0.0-beta.2 (#4056) Updates the Rust workspace Lance crates and Java lance-core dependency to [v12.0.0-beta.2](https://github.com/lance-format/lance/releases/tag/v12.0.0-beta.2). No compatibility fixes were required; full-workspace Clippy passes with all features and warnings denied. --- Cargo.lock | 84 ++++++++++++++++++++++++++-------------------------- Cargo.toml | 28 +++++++++--------- java/pom.xml | 2 +- 3 files changed, 57 insertions(+), 57 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 63ee13e78..e27f9f271 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3455,8 +3455,8 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" [[package]] name = "fsst" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "rand 0.9.5", @@ -4815,8 +4815,8 @@ checksum = "e037a2e1d8d5fdbd49b16a4ea09d5d6401c1f29eca5ff29d03d3824dba16256a" [[package]] name = "lance" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arc-swap", "arrow", @@ -4888,8 +4888,8 @@ dependencies = [ [[package]] name = "lance-arrow" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" 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=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" 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=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -4934,8 +4934,8 @@ dependencies = [ [[package]] name = "lance-bitpacking" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrayref", "crunchy", @@ -4945,8 +4945,8 @@ dependencies = [ [[package]] name = "lance-core" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-buffer", @@ -4983,8 +4983,8 @@ dependencies = [ [[package]] name = "lance-datafusion" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5013,8 +5013,8 @@ dependencies = [ [[package]] name = "lance-datagen" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5031,8 +5031,8 @@ dependencies = [ [[package]] name = "lance-derive" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "proc-macro2", "quote", @@ -5041,8 +5041,8 @@ dependencies = [ [[package]] name = "lance-encoding" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-arith", "arrow-array", @@ -5075,8 +5075,8 @@ dependencies = [ [[package]] name = "lance-file" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-arith", "arrow-array", @@ -5107,8 +5107,8 @@ dependencies = [ [[package]] name = "lance-index" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arc-swap", "arrow", @@ -5172,8 +5172,8 @@ dependencies = [ [[package]] name = "lance-index-core" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -5195,8 +5195,8 @@ dependencies = [ [[package]] name = "lance-io" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5236,8 +5236,8 @@ dependencies = [ [[package]] name = "lance-linalg" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -5251,8 +5251,8 @@ dependencies = [ [[package]] name = "lance-namespace" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "async-trait", @@ -5264,8 +5264,8 @@ dependencies = [ [[package]] name = "lance-namespace-impls" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-ipc", @@ -5318,8 +5318,8 @@ dependencies = [ [[package]] name = "lance-select" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-buffer", @@ -5333,8 +5333,8 @@ dependencies = [ [[package]] name = "lance-table" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow", "arrow-array", @@ -5374,8 +5374,8 @@ dependencies = [ [[package]] name = "lance-testing" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "arrow-array", "arrow-schema", @@ -5388,8 +5388,8 @@ dependencies = [ [[package]] name = "lance-tokenizer" -version = "12.0.0-beta.1" -source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.1#e37cee4397abc6a71df3cdfa0e637c274f820aa6" +version = "12.0.0-beta.2" +source = "git+https://github.com/lance-format/lance.git?tag=v12.0.0-beta.2#dafa4642658d996b3e31dde91e02f72db7860d7e" dependencies = [ "frostem", "icu_segmenter", diff --git a/Cargo.toml b/Cargo.toml index 3f2248e4b..a16f0412c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,20 +13,20 @@ categories = ["database-implementations"] rust-version = "1.91.0" [workspace.dependencies] -lance = { "version" = "=12.0.0-beta.1", default-features = false, "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-core = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-datagen = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-file = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-io = { "version" = "=12.0.0-beta.1", default-features = false, "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-index = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-linalg = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-namespace = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-namespace-impls = { "version" = "=12.0.0-beta.1", default-features = false, "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-table = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-testing = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-datafusion = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-encoding = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } -lance-arrow = { "version" = "=12.0.0-beta.1", "tag" = "v12.0.0-beta.1", "git" = "https://github.com/lance-format/lance.git" } +lance = { "version" = "=12.0.0-beta.2", default-features = false, "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-core = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-datagen = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-file = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-io = { "version" = "=12.0.0-beta.2", default-features = false, "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-index = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-linalg = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-namespace = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-namespace-impls = { "version" = "=12.0.0-beta.2", default-features = false, "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-table = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-testing = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-datafusion = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-encoding = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" } +lance-arrow = { "version" = "=12.0.0-beta.2", "tag" = "v12.0.0-beta.2", "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 diff --git a/java/pom.xml b/java/pom.xml index 78aa17bb7..a2cec19c0 100644 --- a/java/pom.xml +++ b/java/pom.xml @@ -28,7 +28,7 @@ UTF-8 15.0.0 - 12.0.0-beta.1 + 12.0.0-beta.2 false 2.30.0 1.7 From 391cac903438fdf67c7e3675757b4e4898ddb150 Mon Sep 17 00:00:00 2001 From: Jack Ye Date: Tue, 25 Aug 2026 21:54:36 -0700 Subject: [PATCH 23/23] fix(remote): centralize timeline consistency (#4053) Centralizes remote table freshness fencing and response-version tracking in the default transport path. Covers schema and blob bypass paths, keeps explicit time-travel and cross-timeline operations unfenced, and advances freshness after refresh and index job completion. --- rust/lancedb/src/job.rs | 4 + rust/lancedb/src/remote/table.rs | 1354 ++++++++++++++++++----- rust/lancedb/src/remote/table/blobs.rs | 77 +- rust/lancedb/src/remote/table/insert.rs | 83 +- 4 files changed, 1218 insertions(+), 300 deletions(-) diff --git a/rust/lancedb/src/job.rs b/rust/lancedb/src/job.rs index 94d1ba2b6..22f1a0450 100644 --- a/rust/lancedb/src/job.rs +++ b/rust/lancedb/src/job.rs @@ -47,6 +47,10 @@ impl TerminalResult { } } + pub(crate) fn value(&self) -> Option<&Value> { + self.value.as_ref() + } + fn decode(self) -> Result { let value = self.value.ok_or_else(|| match &self.request_id { Some(request_id) => Error::Http { diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index b116083c8..b1a02a260 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -86,22 +86,31 @@ const METRIC_TYPE_KEY: &str = "metric_type"; const INDEX_TYPE_KEY: &str = "index_type"; const SCHEMA_CACHE_TTL: Duration = Duration::from_secs(30); const SCHEMA_CACHE_REFRESH_WINDOW: Duration = Duration::from_secs(5); +const SCHEMA_SELECTOR_CHANGED: &str = "table selector changed while fetching schema"; /// Per-table state driving the freshness headers (`x-lancedb-min-version`, -/// `x-lancedb-min-timestamp`, and `x-lancedb-min-read-version`) sent on read +/// `x-lancedb-min-timestamp`, and `x-lancedb-min-read-version`) sent on table /// requests. #[derive(Debug, Default, Clone, Copy)] struct FreshnessState { + /// Identifies the handle timeline that produced this state. Explicit + /// checkout operations advance the generation so responses from older + /// in-flight requests cannot repopulate the new timeline's constraints. + generation: u64, + /// Exact-version, tag, and snapshot handles must not carry latest-timeline + /// constraints. Their request body already selects the precise version. + pinned: bool, /// Provides read-your-write within a single handle: writes that return a /// version update this, and reads send it as `x-lancedb-min-version`. min_version: Option, - /// Highest dataset version observed in a *read* response on this handle. - /// Reads send it as `x-lancedb-min-read-version` so a load-balanced query + /// Highest committed dataset version advertised by a successful table + /// response on this handle. Later requests send it as + /// `x-lancedb-min-read-version` so a load-balanced query /// node whose cache is behind this version must refresh before serving, /// giving monotonic reads across nodes regardless of which one the load - /// balancer routes to. Sourced only from reads (always committed dataset - /// versions), never from writes (which may return WAL entry ids), so it is - /// unaffected by the WAL/version mismatch that retired `min_version`. + /// balancer routes to. Unlike write result bodies, this is sourced only + /// from the server's committed dataset-version response header or other + /// typed dataset-version fields, so WAL entry ids cannot enter it. min_read_version: Option, /// Wall-clock time captured at the last [`BaseTable::checkout_latest`] /// call. Subsequent reads send @@ -118,14 +127,22 @@ struct FreshnessState { checkout_baseline: Option, } -/// Snapshot of the headers that should be attached to a single read request. +/// Snapshot of the headers that should be attached to a single table request. #[derive(Debug, Default, Clone, Copy)] struct FreshnessHeaders { + generation: u64, min_version: Option, min_timestamp: Option, min_read_version: Option, } +#[derive(Debug, Default, Clone, Copy)] +struct ReadSnapshot { + version: Option, + freshness_state: FreshnessState, + freshness: FreshnessHeaders, +} + impl FreshnessHeaders { fn apply(self, mut request: RequestBuilder) -> RequestBuilder { if let Some(v) = self.min_version { @@ -140,6 +157,61 @@ impl FreshnessHeaders { } request } + + fn observe_version(self, freshness: &Mutex, version: u64) { + track_read_version_for_generation(freshness, self.generation, version); + } + + fn observe_headers( + self, + freshness: &Mutex, + headers: &reqwest::header::HeaderMap, + ) { + if let Some(version) = headers + .get(&VERSION_HEADER) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.parse::().ok()) + { + self.observe_version(freshness, version); + } + } + + fn update_if_current( + self, + freshness: &Mutex, + update: impl FnOnce(&mut FreshnessState), + ) { + let mut state = freshness.lock().unwrap(); + if state.generation == self.generation { + update(&mut state); + } + } + + fn is_current(self, freshness: &Mutex) -> bool { + freshness.lock().unwrap().generation == self.generation + } +} + +fn track_read_version(freshness: &Mutex, version: u64) { + if version == 0 { + return; + } + let mut state = freshness.lock().unwrap(); + state.min_read_version = Some(state.min_read_version.map_or(version, |v| v.max(version))); +} + +fn track_read_version_for_generation( + freshness: &Mutex, + generation: u64, + version: u64, +) { + if version == 0 { + return; + } + let mut state = freshness.lock().unwrap(); + if state.generation == generation { + state.min_read_version = Some(state.min_read_version.map_or(version, |v| v.max(version))); + } } /// A backfill job whose successful wait establishes a read-freshness @@ -150,6 +222,8 @@ struct FreshnessJob { inner: RemoteJob, freshness: Arc>, version: Arc>>, + track_refresh_result: bool, + freshness_request: FreshnessHeaders, } #[async_trait] @@ -166,7 +240,31 @@ impl crate::job::JobHandle for FreshnessJob { let result = crate::job::JobHandle::wait(&self.inner).await?; let version = self.version.read().await; if version.is_none() { - self.freshness.lock().unwrap().checkout_baseline = Some(SystemTime::now()); + let result_version = self + .track_refresh_result + .then(|| result.value()) + .flatten() + .and_then(|value| { + serde_json::from_value::(value.clone()) + .ok() + }) + .map(|result| { + result + .published_version + .map_or(result.source_version, |version| { + version.max(result.source_version) + }) + }) + .filter(|version| *version != 0); + if let Some(version) = result_version { + self.freshness_request + .observe_version(&self.freshness, version); + } else { + self.freshness_request + .update_if_current(&self.freshness, |state| { + state.checkout_baseline = Some(SystemTime::now()); + }); + } } Ok(result) } @@ -193,6 +291,38 @@ fn compute_min_timestamp( } } +fn freshness_headers_snapshot( + freshness: &Mutex, + interval: Option, +) -> FreshnessHeaders { + freshness_state_snapshot(freshness, interval).1 +} + +fn freshness_state_snapshot( + freshness: &Mutex, + interval: Option, +) -> (FreshnessState, FreshnessHeaders) { + let state = *freshness.lock().unwrap(); + if state.pinned { + return ( + state, + FreshnessHeaders { + generation: state.generation, + ..FreshnessHeaders::default() + }, + ); + } + ( + state, + FreshnessHeaders { + generation: state.generation, + min_version: state.min_version, + min_timestamp: compute_min_timestamp(&state, interval, SystemTime::now()), + min_read_version: state.min_read_version, + }, + ) +} + /// Normalize a branch selector: trim whitespace and treat `""` or `"main"` as /// the (absent) main branch, matching the server's convention. fn normalize_branch(branch: Option) -> Option { @@ -210,8 +340,9 @@ impl Tags for RemoteTags<'_, S> { async fn list(&self) -> Result> { let request = self .inner - .post_read(&format!("/v1/table/{}/tags/list/", self.inner.identifier)); - let (request_id, response) = self.inner.send(request, true).await?; + .client + .post(&format!("/v1/table/{}/tags/list/", self.inner.identifier)); + let (request_id, response) = self.inner.send_unfenced(request, true).await?; let response = self .inner .check_table_response(&request_id, response) @@ -241,12 +372,12 @@ impl Tags for RemoteTags<'_, S> { } async fn get_version(&self, tag: &str) -> Result { - let request = self.inner.post_read(&format!( + let request = self.inner.client.post(&format!( "/v1/table/{}/tags/version/", self.inner.identifier )); self.inner - .resolve_tag_version_with_request(tag, request) + .resolve_tag_version_with_request(tag, request, false) .await } @@ -276,7 +407,7 @@ impl Tags for RemoteTags<'_, S> { .post(&format!("/v1/table/{}/tags/delete/", self.inner.identifier)) .json(&serde_json::json!({ "tag": tag })); - let (request_id, response) = self.inner.send(request, true).await?; + let (request_id, response) = self.inner.send_unfenced(request, true).await?; self.inner .check_table_response(&request_id, response) .await?; @@ -468,6 +599,7 @@ impl RemoteTable { let Ok(description) = serde_json::from_str::(describe_body) else { return; }; + self.track_read_version(description.version); if let Ok(schema) = arrow_schema::Schema::try_from(description.schema) { self.schema_cache.seed(Arc::new(schema)); } @@ -510,23 +642,50 @@ impl RemoteTable { } async fn describe(&self) -> Result { - let version = self.current_version().await; - self.describe_version(version).await + self.describe_read_snapshot(self.snapshot_read_state().await) + .await } - async fn describe_version(&self, version: Option) -> Result { - let request = self.post_read(&format!("/v1/table/{}/describe/", self.identifier)); - self.describe_with_request(request, version).await + async fn describe_read_snapshot( + &self, + read_snapshot: ReadSnapshot, + ) -> Result { + let request = self + .client + .post(&format!("/v1/table/{}/describe/", self.identifier)); + self.describe_with_request( + request, + read_snapshot.version, + Some(read_snapshot.freshness), + ) + .await + } + + async fn schema_read_snapshot(&self, read_snapshot: ReadSnapshot) -> Result { + if read_snapshot.freshness.is_current(&self.freshness) + && let Some(schema) = self.schema_cache.try_get() + && read_snapshot.freshness.is_current(&self.freshness) + { + return Ok(schema); + } + + let description = self.describe_read_snapshot(read_snapshot).await?; + Ok(Arc::new(description.schema.try_into()?)) } async fn resolve_tag_version_with_request( &self, tag: &str, request: RequestBuilder, + fenced: bool, ) -> Result { let request = request.json(&serde_json::json!({ "tag": tag })); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = if fenced { + self.send(request, true).await? + } else { + self.send_unfenced(request, true).await? + }; let response = self.check_table_response(&request_id, response).await?; match response.text().await { @@ -565,7 +724,7 @@ impl RemoteTable { .client .post(&format!("/v1/table/{}/tags/version/", self.identifier)) .json(&serde_json::json!({ "tag": tag })); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self.send_unfenced(request, true).await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; let value: serde_json::Value = serde_json::from_str(&body).map_err(|e| Error::Http { @@ -592,21 +751,34 @@ impl RemoteTable { &self, request: RequestBuilder, version: Option, + freshness_request: Option, ) -> Result { let mut body = serde_json::json!({ "version": version }); self.apply_branch_body(&mut body); let request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = if let Some(freshness_request) = freshness_request { + self.send_with_freshness(request, true, freshness_request) + .await? + } else { + self.send_unfenced(request, true).await? + }; let response = self.check_table_response(&request_id, response).await?; match response.text().await { - Ok(body) => serde_json::from_str(&body).map_err(|e| Error::Http { - source: format!("Failed to parse table description: {}", e).into(), - request_id, - status_code: None, - }), + Ok(body) => { + let description: TableDescription = + serde_json::from_str(&body).map_err(|e| Error::Http { + source: format!("Failed to parse table description: {}", e).into(), + request_id, + status_code: None, + })?; + if let Some(freshness_request) = freshness_request { + freshness_request.observe_version(&self.freshness, description.version); + } + Ok(description) + } Err(err) => { let status_code = err.status(); Err(Error::Http { @@ -619,14 +791,41 @@ impl RemoteTable { } async fn send(&self, req: RequestBuilder, with_retry: bool) -> Result<(String, Response)> { + let freshness_request = self.snapshot_freshness_headers(); + self.send_with_freshness(req, with_retry, freshness_request) + .await + } + + async fn send_with_freshness( + &self, + req: RequestBuilder, + with_retry: bool, + freshness_request: FreshnessHeaders, + ) -> Result<(String, Response)> { + let req = freshness_request.apply(req); let res = if with_retry { self.client.send_with_retry(req, None, true).await? } else { self.client.send(req).await? }; + if res.1.status().is_success() { + freshness_request.observe_headers(&self.freshness, res.1.headers()); + } Ok(res) } + async fn send_unfenced( + &self, + req: RequestBuilder, + with_retry: bool, + ) -> Result<(String, Response)> { + if with_retry { + self.client.send_with_retry(req, None, true).await + } else { + self.client.send(req).await + } + } + pub(super) async fn handle_table_not_found( table_name: &str, response: reqwest::Response, @@ -1003,24 +1202,32 @@ impl RemoteTable { } } - async fn current_version(&self) -> Option { - let read_guard = self.version.read().await; - *read_guard + async fn snapshot_read_state(&self) -> ReadSnapshot { + let version = self.version.read().await; + let (freshness_state, freshness) = + freshness_state_snapshot(&self.freshness, self.client.read_consistency_interval); + ReadSnapshot { + version: *version, + freshness_state, + freshness, + } } - /// Snapshot the freshness headers to attach to a single read request. + /// Snapshot the freshness headers to attach to a single table request. /// Computed at call time so that retries reuse the same snapshot. fn snapshot_freshness_headers(&self) -> FreshnessHeaders { - let state = *self.freshness.lock().unwrap(); - FreshnessHeaders { - min_version: state.min_version, - min_timestamp: compute_min_timestamp( - &state, - self.client.read_consistency_interval, - SystemTime::now(), - ), - min_read_version: state.min_read_version, - } + freshness_headers_snapshot(&self.freshness, self.client.read_consistency_interval) + } + + fn reset_freshness(&self, checkout_baseline: Option, pinned: bool) { + let mut state = self.freshness.lock().unwrap(); + let generation = state.generation.wrapping_add(1); + *state = FreshnessState { + generation, + pinned, + checkout_baseline, + ..FreshnessState::default() + }; } /// Send an LSM operator request with the transport retry layer **off**. @@ -1035,46 +1242,25 @@ impl RemoteTable { Ok((request_id, response)) } - /// Build a POST request and attach the read-freshness headers - /// (`x-lancedb-min-version`, `x-lancedb-min-timestamp`). - fn post_read(&self, uri: &str) -> RequestBuilder { - self.snapshot_freshness_headers() - .apply(self.client.post(uri)) - } - /// Record a version returned by a write so subsequent reads can request at /// least that version via `x-lancedb-min-version`. A returned `0` from a /// backward-compatible old server is ignored. - fn track_write_version(&self, version: u64) { + fn track_write_version(&self, freshness_request: FreshnessHeaders, version: u64) { if version == 0 { return; } - let mut state = self.freshness.lock().unwrap(); - state.min_version = Some(state.min_version.map_or(version, |v| v.max(version))); + freshness_request.update_if_current(&self.freshness, |state| { + state.min_version = Some(state.min_version.map_or(version, |v| v.max(version))); + }); } - /// Record a dataset version observed in a *read* response so subsequent - /// reads request at least this version via `x-lancedb-min-read-version`, + /// Record a committed dataset version observed in a table response so + /// subsequent requests ask for at least this version via + /// `x-lancedb-min-read-version`, /// giving monotonic reads across load-balanced query nodes. A returned `0` /// (or absent header from an old server) is ignored. fn track_read_version(&self, version: u64) { - if version == 0 { - return; - } - let mut state = self.freshness.lock().unwrap(); - state.min_read_version = Some(state.min_read_version.map_or(version, |v| v.max(version))); - } - - /// Parse the `x-lancedb-version` response header (the dataset version a read - /// reflects) and fold it into the read-version watermark. - fn track_read_version_from_headers(&self, headers: &reqwest::header::HeaderMap) { - if let Some(version) = headers - .get(&VERSION_HEADER) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.parse::().ok()) - { - self.track_read_version(version); - } + track_read_version(&self.freshness, version); } async fn execute_query( @@ -1082,7 +1268,9 @@ impl RemoteTable { query: &AnyQuery, options: &QueryExecutionOptions, ) -> Result>>> { - let mut request = self.post_read(&format!("/v1/table/{}/query/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/query/", self.identifier)); if let Some(timeout) = options.timeout { // Also send to server, so it can abort the query if it takes too long. @@ -1092,15 +1280,17 @@ impl RemoteTable { } } - let query_bodies = self.prepare_query_bodies(query).await?; + let read_snapshot = self.snapshot_read_state().await; + let query_bodies = self.prepare_query_bodies(query, read_snapshot.version)?; let requests: Vec = query_bodies .into_iter() .map(|body| request.try_clone().unwrap().json(&body)) .collect(); let futures = requests.into_iter().map(|req| async move { - let (request_id, response) = self.send(req, true).await?; - self.track_read_version_from_headers(response.headers()); + let (request_id, response) = self + .send_with_freshness(req, true, read_snapshot.freshness) + .await?; self.read_arrow_response(&request_id, response).await }); let streams = futures::future::try_join_all(futures); @@ -1125,8 +1315,11 @@ impl RemoteTable { } } - async fn prepare_query_bodies(&self, query: &AnyQuery) -> Result> { - let version = self.current_version().await; + fn prepare_query_bodies( + &self, + query: &AnyQuery, + version: Option, + ) -> Result> { let mut base_body = serde_json::json!({ "version": version }); self.apply_branch_body(&mut base_body); @@ -1230,14 +1423,15 @@ async fn fetch_schema( client: &RestfulLanceDbClient, identifier: &str, table_name: &str, - version: Option, + read_snapshot: ReadSnapshot, branch: Option, - freshness_headers: FreshnessHeaders, + freshness: Arc>, ) -> Result { - let mut body = serde_json::json!({ "version": version }); + let mut body = serde_json::json!({ "version": read_snapshot.version }); if let Some(branch) = &branch { body["branch"] = serde_json::Value::String(branch.clone()); } + let freshness_headers = read_snapshot.freshness; let request = freshness_headers .apply(client.post(&format!("/v1/table/{}/describe/", identifier))) .json(&body); @@ -1257,6 +1451,7 @@ async fn fetch_schema( } let response = client.check_response(&request_id, response).await?; + freshness_headers.observe_headers(&freshness, response.headers()); let body = response.text().await.map_err(|e| { let status_code = e.status(); Error::Http { @@ -1271,6 +1466,12 @@ async fn fetch_schema( request_id, status_code: None, })?; + freshness_headers.observe_version(&freshness, description.version); + if !freshness_headers.is_current(&freshness) { + return Err(Error::Runtime { + message: SCHEMA_SELECTOR_CHANGED.to_string(), + }); + } let arrow_schema: arrow_schema::Schema = description.schema.try_into()?; Ok(Arc::new(arrow_schema)) @@ -1401,18 +1602,25 @@ impl RemoteTable { use crate::remote::retry::RetryCounter; let _guard = output.tracker.as_ref().map(|t| t.track_task()); + let freshness_request = self.snapshot_freshness_headers(); - let mut insert: Arc = Arc::new(RemoteWriteExec::new( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - output.plan, - WriteOp::Insert { - overwrite: output.overwrite, - }, - output.tracker.clone(), - self.branch.clone(), - )); + let mut insert: Arc = Arc::new( + RemoteWriteExec::new( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + output.plan, + WriteOp::Insert { + overwrite: output.overwrite, + }, + output.tracker.clone(), + self.branch.clone(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + ); let mut retry_counter = RetryCounter::new(&self.client.retry_config, uuid::Uuid::new_v4().to_string()); @@ -1431,7 +1639,7 @@ impl RemoteTable { if output.overwrite { self.invalidate_schema_cache(); } - self.track_write_version(add_result.version); + self.track_write_version(freshness_request, add_result.version); return Ok(add_result); } @@ -1457,6 +1665,7 @@ impl RemoteTable { RetryCounter::new(&self.client.retry_config, uuid::Uuid::new_v4().to_string()); loop { + let freshness_request = self.snapshot_freshness_headers(); let upload_id = self.create_multipart_write().await?; let result = self @@ -1469,7 +1678,7 @@ impl RemoteTable { if output.overwrite { self.invalidate_schema_cache(); } - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); return Ok(result); } Err(e) => { @@ -1525,18 +1734,24 @@ impl RemoteTable { )?, ) as Arc; - let insert = Arc::new(RemoteWriteExec::new_multipart( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - plan, - output.overwrite, - upload_id.to_string(), - output.tracker.clone(), - self.branch.clone(), - self.client.max_bytes_per_request(), - self.client.max_request_duration(), - )); + let insert = Arc::new( + RemoteWriteExec::new_multipart( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + plan, + output.overwrite, + upload_id.to_string(), + output.tracker.clone(), + self.branch.clone(), + self.client.max_bytes_per_request(), + self.client.max_request_duration(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + ); let task_ctx = Arc::new(datafusion_execution::TaskContext::default()); let tracker = output.tracker.clone(); @@ -1599,6 +1814,38 @@ where } impl RemoteTable { + async fn index_stats_read_snapshot( + &self, + index_name: &str, + read_snapshot: ReadSnapshot, + ) -> Result> { + let encoded_name = urlencoding::encode(index_name); + let mut body = serde_json::json!({ "version": read_snapshot.version }); + self.apply_branch_body(&mut body); + let request = self + .client + .post(&format!( + "/v1/table/{}/index/{encoded_name}/stats/", + self.identifier + )) + .json(&body); + + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; + if response.status() == StatusCode::NOT_FOUND { + return Ok(None); + } + let response = self.check_table_response(&request_id, response).await?; + let body = response.text().await.err_to_http(request_id.clone())?; + let stats = serde_json::from_str(&body).map_err(|e| Error::Http { + source: format!("Failed to parse index statistics: {}", e).into(), + request_id, + status_code: None, + })?; + Ok(Some(stats)) + } + /// Parse the response from `/index/list/` into `IndexConfig` entries. /// /// When the server returns `index_type` inline, all enriched fields are @@ -1610,6 +1857,7 @@ impl RemoteTable { body: &str, request_id: &str, schema: &SchemaRef, + read_snapshot: ReadSnapshot, ) -> Result> { use crate::index::IndexType; @@ -1678,7 +1926,10 @@ impl RemoteTable { })) } else { // Legacy response: fetch index type via stats endpoint. - match self.index_stats(&entry.index_name).await { + match self + .index_stats_read_snapshot(&entry.index_name, read_snapshot) + .await + { Ok(Some(stats)) => Ok(Some(IndexConfig { name: entry.index_name, index_type: stats.index_type, @@ -1762,7 +2013,7 @@ impl BaseTable for RemoteTable { let request = self .client .post(&format!("/v1/table/{}/describe/", self.identifier)); - self.describe_with_request(request, Some(version)) + self.describe_with_request(request, Some(version), None) .await .map_err(|e| match e { // try to map the error to a more user-friendly error telling them @@ -1775,58 +2026,52 @@ impl BaseTable for RemoteTable { })?; let mut write_guard = self.version.write().await; + // Commit the selector and its freshness mode while holding the selector + // write lock, with no cancellation point between the two updates. + self.reset_freshness(None, true); *write_guard = Some(version); - drop(write_guard); - - // Explicit time-travel: drop any read-your-write / freshness - // constraints so the user sees exactly the requested version. - *self.freshness.lock().unwrap() = FreshnessState::default(); - - // Invalidate schema cache since we're switching versions self.invalidate_schema_cache(); + drop(write_guard); Ok(()) } async fn checkout_latest(&self) -> Result<()> { let mut write_guard = self.version.write().await; - *write_guard = None; - drop(write_guard); - // Drop any per-handle read/write tracking; subsequent reads use the // baseline timestamp captured now to guarantee freshness. - *self.freshness.lock().unwrap() = FreshnessState { - min_version: None, - checkout_baseline: Some(SystemTime::now()), - min_read_version: None, - }; - - // Invalidate schema cache since we're switching versions + self.reset_freshness(Some(SystemTime::now()), false); + *write_guard = None; self.invalidate_schema_cache(); + drop(write_guard); Ok(()) } async fn snapshot_at_current_version(&self) -> Result>> { // 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 { + let read_snapshot = self.snapshot_read_state().await; + let version = match read_snapshot.version { Some(version) => version, - None => self.describe().await?.version, + None => self.describe_read_snapshot(read_snapshot).await?.version, }; let snapshot = self.with_branch(self.branch.clone()); *snapshot.version.write().await = Some(version); + snapshot.reset_freshness(None, true); Ok(Some(Arc::new(snapshot))) } async fn restore(&self) -> Result<()> { let mut request = self .client .post(&format!("/v1/table/{}/restore/", self.identifier)); - let version = self.current_version().await; - let mut body = serde_json::json!({ "version": version }); + let read_snapshot = self.snapshot_read_state().await; + let mut body = serde_json::json!({ "version": read_snapshot.version }); self.apply_branch_body(&mut body); request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; self.check_table_response(&request_id, response).await?; self.checkout_latest().await?; Ok(()) @@ -1834,7 +2079,8 @@ impl BaseTable for RemoteTable { async fn list_versions(&self) -> Result> { let request = self.apply_branch_query( - self.post_read(&format!("/v1/table/{}/version/list/", self.identifier)), + self.client + .post(&format!("/v1/table/{}/version/list/", self.identifier)), ); let (request_id, response) = self.send(request, true).await?; let response = self.check_table_response(&request_id, response).await?; @@ -1898,31 +2144,42 @@ impl BaseTable for RemoteTable { } async fn schema(&self) -> Result { - if let Some(schema) = self.schema_cache.try_get() { - return Ok(schema); - } + loop { + let read_snapshot = self.snapshot_read_state().await; + if let Some(schema) = self.schema_cache.try_get() { + return Ok(schema); + } - let version = self.current_version().await; - let client = self.client.clone(); - let identifier = self.identifier.clone(); - let table_name = self.name.clone(); - let branch = self.branch.clone(); - let freshness_headers = self.snapshot_freshness_headers(); + let client = self.client.clone(); + let identifier = self.identifier.clone(); + let table_name = self.name.clone(); + let branch = self.branch.clone(); + let freshness = self.freshness.clone(); - self.schema_cache - .get(move || async move { - fetch_schema( - &client, - &identifier, - &table_name, - version, - branch, - freshness_headers, - ) + match self + .schema_cache + .get(move || async move { + fetch_schema( + &client, + &identifier, + &table_name, + read_snapshot, + branch, + freshness, + ) + .await + }) .await - }) - .await - .map_err(unwrap_shared_error) + { + Ok(schema) => return Ok(schema), + Err(error) + if matches!( + &*error, + Error::Runtime { message } if message == SCHEMA_SELECTOR_CHANGED + ) => {} + Err(error) => return Err(unwrap_shared_error(error)), + } + } } async fn create_branch( @@ -1964,7 +2221,7 @@ impl BaseTable for RemoteTable { // Send without retry so the expected 409 (branch already exists) is // surfaced as a response we can map, rather than being retried. - let (request_id, response) = self.send(request, false).await?; + let (request_id, response) = self.send_unfenced(request, false).await?; match response.status() { StatusCode::CONFLICT => { return Err(Error::TableAlreadyExists { @@ -2025,8 +2282,10 @@ impl BaseTable for RemoteTable { async fn list_branches(&self) -> Result> { use lance::dataset::refs::BranchContents; - let request = self.post_read(&format!("/v1/table/{}/branches/list/", self.identifier)); - let (request_id, response) = self.send(request, true).await?; + let request = self + .client + .post(&format!("/v1/table/{}/branches/list/", self.identifier)); + let (request_id, response) = self.send_unfenced(request, true).await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2080,7 +2339,7 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/branches/diff/", self.identifier)) .json(&serde_json::json!({ "from_branch": from_branch })); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self.send_unfenced(request, true).await?; if response.status() == StatusCode::NOT_FOUND { return Err(Error::TableNotFound { name: format!("{} (branch: {})", self.name, from_branch), @@ -2106,6 +2365,12 @@ impl BaseTable for RemoteTable { message: "Branch name cannot be empty.".into(), }); } + let read_snapshot = self.snapshot_read_state().await; + let target_freshness = if self.branch.is_none() && read_snapshot.version.is_none() { + Some(read_snapshot.freshness) + } else { + None + }; let request = self .client .post(&format!( @@ -2117,7 +2382,7 @@ impl BaseTable for RemoteTable { "dry_run": dry_run, })); // No retry. HTTP 409 is CherryPickStatus::Failed with a body, not a transport error. - let (request_id, response) = self.send(request, false).await?; + let (request_id, response) = self.send_unfenced(request, false).await?; let status = response.status(); if status == StatusCode::NOT_FOUND { return Err(Error::TableNotFound { @@ -2135,7 +2400,7 @@ impl BaseTable for RemoteTable { }); } let body = response.text().await.err_to_http(request_id.clone())?; - serde_json::from_str(&body).map_err(|err| Error::Http { + let result: CherryPickResult = serde_json::from_str(&body).map_err(|err| Error::Http { source: format!( "Failed to parse cherry_pick response: {}, body: {}", err, body @@ -2143,7 +2408,15 @@ impl BaseTable for RemoteTable { .into(), request_id, status_code: Some(status), - }) + })?; + if !dry_run + && status == StatusCode::OK + && result.status == crate::table::CherryPickStatus::CherryPicked + && let (Some(freshness), Some(version)) = (target_freshness, result.main_version_after) + { + freshness.observe_version(&self.freshness, version); + } + Ok(result) } fn current_branch(&self) -> Option { @@ -2151,23 +2424,28 @@ impl BaseTable for RemoteTable { } async fn count_rows(&self, filter: Option) -> Result { - let mut request = self.post_read(&format!("/v1/table/{}/count_rows/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/count_rows/", self.identifier)); - let version = self.current_version().await; + let read_snapshot = self.snapshot_read_state().await; let mut body = if let Some(filter) = filter { let filter_sql = match filter { Filter::Sql(sql) => sql.clone(), Filter::Datafusion(expr) => expr_to_sql_string(&expr)?, }; - serde_json::json!({ "predicate": filter_sql, "version": version }) + serde_json::json!({ "predicate": filter_sql, "version": read_snapshot.version }) } else { - serde_json::json!({ "version": version }) + serde_json::json!({ "version": read_snapshot.version }) }; self.apply_branch_body(&mut body); request = request.json(&body); - let (request_id, response) = match self.send(request, true).await { + let (request_id, response) = match self + .send_with_freshness(request, true, read_snapshot.freshness) + .await + { Ok((id, resp)) => { // check_table_response now handles error-based invalidation let response = self.check_table_response(&id, resp).await?; @@ -2179,7 +2457,6 @@ impl BaseTable for RemoteTable { } }; - self.track_read_version_from_headers(response.headers()); let body = response.text().await.err_to_http(request_id.clone())?; serde_json::from_str(&body).map_err(|e| Error::Http { @@ -2312,9 +2589,12 @@ impl BaseTable for RemoteTable { } async fn explain_plan(&self, query: &AnyQuery, verbose: bool) -> Result { - let base_request = self.post_read(&format!("/v1/table/{}/explain_plan/", self.identifier)); + let base_request = self + .client + .post(&format!("/v1/table/{}/explain_plan/", self.identifier)); - let query_bodies = self.prepare_query_bodies(query).await?; + let read_snapshot = self.snapshot_read_state().await; + let query_bodies = self.prepare_query_bodies(query, read_snapshot.version)?; let requests: Vec = query_bodies .into_iter() .map(|query_body| { @@ -2328,7 +2608,9 @@ impl BaseTable for RemoteTable { .collect::>(); let futures = requests.into_iter().map(|req| async move { - let (request_id, response) = self.send(req, true).await?; + let (request_id, response) = self + .send_with_freshness(req, true, read_snapshot.freshness) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2359,7 +2641,9 @@ impl BaseTable for RemoteTable { query: &AnyQuery, options: QueryExecutionOptions, ) -> Result { - let mut request = self.post_read(&format!("/v1/table/{}/analyze_plan/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/analyze_plan/", self.identifier)); if options.analyze_plan_distributed_metrics != AnalyzePlanDistributedMetrics::Aggregate { request = request.query(&[( @@ -2368,14 +2652,17 @@ impl BaseTable for RemoteTable { )]); } - let query_bodies = self.prepare_query_bodies(query).await?; + let read_snapshot = self.snapshot_read_state().await; + let query_bodies = self.prepare_query_bodies(query, read_snapshot.version)?; let requests: Vec = query_bodies .into_iter() .map(|body| request.try_clone().unwrap().json(&body)) .collect(); let futures = requests.into_iter().map(|req| async move { - let (request_id, response) = self.send(req, true).await?; + let (request_id, response) = self + .send_with_freshness(req, true, read_snapshot.freshness) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2419,7 +2706,10 @@ impl BaseTable for RemoteTable { self.apply_branch_body(&mut body); let request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2438,7 +2728,7 @@ impl BaseTable for RemoteTable { status_code: None, })?; - self.track_write_version(update_response.version); + self.track_write_version(freshness_request, update_response.version); Ok(update_response) } @@ -2454,7 +2744,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/delete/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; if body.trim().is_empty() { @@ -2470,7 +2763,7 @@ impl BaseTable for RemoteTable { request_id, status_code: None, })?; - self.track_write_version(delete_response.version); + self.track_write_version(freshness_request, delete_response.version); Ok(delete_response) } @@ -2480,7 +2773,13 @@ impl BaseTable for RemoteTable { async fn create_index_async(&self, index: IndexBuilder) -> Result { Ok(match self.submit_create_index(index).await? { - Some(job_id) => Job::new(Box::new(RemoteJob::new(self.client.clone(), job_id))), + Some(job_id) => Job::new(Box::new(FreshnessJob { + inner: RemoteJob::new(self.client.clone(), job_id), + freshness: self.freshness.clone(), + version: self.version.clone(), + track_refresh_result: false, + freshness_request: self.snapshot_freshness_headers(), + })), None => Job::new_done(), }) } @@ -2520,16 +2819,23 @@ impl BaseTable for RemoteTable { let rescannable = source.rescannable(); let input: Arc = Arc::new(crate::table::datafusion::scannable_exec::ScannableExec::new(source, None)); + let freshness_request = self.snapshot_freshness_headers(); - let mut merge: Arc = Arc::new(RemoteWriteExec::new( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - input, - WriteOp::MergeInsert { query, timeout }, - None, - self.branch.clone(), - )); + let mut merge: Arc = Arc::new( + RemoteWriteExec::new( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + input, + WriteOp::MergeInsert { query, timeout }, + None, + self.branch.clone(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + ); let mut retry_counter = crate::remote::retry::RetryCounter::new( &self.client.retry_config, @@ -2547,7 +2853,7 @@ impl BaseTable for RemoteTable { .and_then(|m| m.merge_result()) .unwrap_or_default(); - self.track_write_version(merge_result.version); + self.track_write_version(freshness_request, merge_result.version); return Ok(merge_result); } Err(err) if rescannable && self.is_retryable_write_error(&err) => { @@ -2586,7 +2892,8 @@ impl BaseTable for RemoteTable { async fn get_lsm_stats(&self, include_generation_rows: bool) -> Result> { // Read-semantics POST, like `get_lsm_write_spec`. let request = self - .post_read(&format!("/v1/table/{}/get_lsm_stats/", self.identifier)) + .client + .post(&format!("/v1/table/{}/get_lsm_stats/", self.identifier)) .json(&serde_json::json!({ "include_generation_rows": include_generation_rows, })); @@ -2660,7 +2967,7 @@ impl BaseTable for RemoteTable { // re-encodes it into the same sophon-owned shape the set endpoint // accepts — no lance/lancedb types cross the wire. `lsm_write_spec` is // null when the LSM write path is not enabled for the table. - let request = self.post_read(&format!( + let request = self.client.post(&format!( "/v1/table/{}/get_lsm_write_spec/", self.identifier )); @@ -2726,18 +3033,17 @@ impl BaseTable for RemoteTable { let request = self .client .post(&format!("/v1/table/{}/tags/version/", self.identifier)); - let version = self.resolve_tag_version_with_request(tag, request).await?; + let version = self + .resolve_tag_version_with_request(tag, request, false) + .await?; let mut write_guard = self.version.write().await; + // Commit the selector and its freshness mode while holding the selector + // write lock, with no cancellation point between the two updates. + self.reset_freshness(None, true); *write_guard = Some(version); - drop(write_guard); - - // Explicit time-travel: drop any read-your-write / freshness - // constraints so the user sees exactly the tagged version. - *self.freshness.lock().unwrap() = FreshnessState::default(); - - // Invalidate schema cache since we're switching versions self.invalidate_schema_cache(); + drop(write_guard); Ok(()) } @@ -2774,7 +3080,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/add_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2791,7 +3100,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -2827,7 +3136,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/add_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2843,7 +3155,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -2886,7 +3198,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/add_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2901,7 +3216,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -2923,7 +3238,8 @@ impl BaseTable for RemoteTable { let mut body = serde_json::json!({ "column": column }); self.apply_branch_body(&mut body); let request = self - .post_read(&format!("/v1/table/{}/backfill_column", self.identifier)) + .client + .post(&format!("/v1/table/{}/backfill_column", self.identifier)) .json(&body); let (request_id, response) = self.send(request, true).await?; let response = self.check_table_response(&request_id, response).await?; @@ -2943,6 +3259,8 @@ impl BaseTable for RemoteTable { inner: RemoteJob::new(self.client.clone(), response.job_id), freshness: self.freshness.clone(), version: self.version.clone(), + track_refresh_result: true, + freshness_request: self.snapshot_freshness_headers(), }))) } @@ -2974,7 +3292,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/alter_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -2990,7 +3311,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -3009,7 +3330,10 @@ impl BaseTable for RemoteTable { self.identifier )) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -3021,7 +3345,7 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } @@ -3033,7 +3357,10 @@ impl BaseTable for RemoteTable { .client .post(&format!("/v1/table/{}/drop_columns/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let freshness_request = self.snapshot_freshness_headers(); + let (request_id, response) = self + .send_with_freshness(request, true, freshness_request) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; @@ -3049,55 +3376,34 @@ impl BaseTable for RemoteTable { })?; self.invalidate_schema_cache(); - self.track_write_version(result.version); + self.track_write_version(freshness_request, result.version); Ok(result) } async fn list_indices(&self) -> Result> { - let mut request = self.post_read(&format!("/v1/table/{}/index/list/", self.identifier)); - let version = self.current_version().await; - let mut body = serde_json::json!({ "version": version }); + let mut request = self + .client + .post(&format!("/v1/table/{}/index/list/", self.identifier)); + let read_snapshot = self.snapshot_read_state().await; + let mut body = serde_json::json!({ "version": read_snapshot.version }); self.apply_branch_body(&mut body); request = request.json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; let response = self.check_table_response(&request_id, response).await?; let body = response.text().await.err_to_http(request_id.clone())?; - let schema = self.schema().await?; + let schema = self.schema_read_snapshot(read_snapshot).await?; - self.parse_index_list_response(&body, &request_id, &schema) + self.parse_index_list_response(&body, &request_id, &schema, read_snapshot) .await } async fn index_stats(&self, index_name: &str) -> Result> { - let encoded_name = urlencoding::encode(index_name); - let mut request = self.post_read(&format!( - "/v1/table/{}/index/{encoded_name}/stats/", - self.identifier - )); - let version = self.current_version().await; - let mut body = serde_json::json!({ "version": version }); - self.apply_branch_body(&mut body); - request = request.json(&body); - - let (request_id, response) = self.send(request, true).await?; - - if response.status() == StatusCode::NOT_FOUND { - return Ok(None); - } - - let response = self.check_table_response(&request_id, response).await?; - - let body = response.text().await.err_to_http(request_id.clone())?; - - let stats = serde_json::from_str(&body).map_err(|e| Error::Http { - source: format!("Failed to parse index statistics: {}", e).into(), - request_id, - status_code: None, - })?; - - Ok(Some(stats)) + self.index_stats_read_snapshot(index_name, self.snapshot_read_state().await) + .await } async fn drop_index(&self, index_name: &str) -> Result<()> { @@ -3187,7 +3493,9 @@ impl BaseTable for RemoteTable { } async fn stats(&self) -> Result { - let mut request = self.post_read(&format!("/v1/table/{}/stats/", self.identifier)); + let mut request = self + .client + .post(&format!("/v1/table/{}/stats/", self.identifier)); if let Some(branch) = &self.branch { request = request.json(&serde_json::json!({ "branch": branch })); } @@ -3209,15 +3517,21 @@ impl BaseTable for RemoteTable { write_params: lance::dataset::WriteParams, ) -> Result> { let overwrite = matches!(write_params.mode, lance::dataset::WriteMode::Overwrite); - Ok(Arc::new(insert::RemoteWriteExec::new( - self.name.clone(), - self.identifier.clone(), - self.client.clone(), - input, - WriteOp::Insert { overwrite }, - None, - self.branch.clone(), - ))) + Ok(Arc::new( + insert::RemoteWriteExec::new( + self.name.clone(), + self.identifier.clone(), + self.client.clone(), + input, + WriteOp::Insert { overwrite }, + None, + self.branch.clone(), + ) + .with_freshness( + self.freshness.clone(), + self.client.read_consistency_interval, + ), + )) } } @@ -7113,12 +7427,12 @@ mod tests { } /// The gate's reproducer: after a successful wait, a same-handle read - /// must carry a freshness baseline so a stale server cache cannot serve - /// the pre-backfill snapshot. + /// must carry the exact published version so a stale server cache cannot + /// serve the pre-backfill snapshot. #[tokio::test] async fn test_backfill_wait_establishes_read_freshness() { - let saw_min_timestamp = Arc::new(std::sync::atomic::AtomicBool::new(false)); - let saw = saw_min_timestamp.clone(); + let saw_published_version = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let saw = saw_published_version.clone(); let table = Table::new_with_handler("my_table", move |request| match request.url().path() { "/v1/table/my_table/backfill_column" => http::Response::builder() @@ -7131,7 +7445,11 @@ mod tests { .unwrap(), "/v1/table/my_table/count_rows/" => { saw.store( - request.headers().contains_key("x-lancedb-min-timestamp"), + request + .headers() + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()) + == Some("8"), std::sync::atomic::Ordering::SeqCst, ); http::Response::builder() @@ -7148,8 +7466,8 @@ mod tests { assert_eq!(result.published_version, Some(8)); table.count_rows(None).await.unwrap(); assert!( - saw_min_timestamp.load(std::sync::atomic::Ordering::SeqCst), - "read after wait carried no freshness baseline" + saw_published_version.load(std::sync::atomic::Ordering::SeqCst), + "read after wait did not carry the published version" ); } @@ -7324,12 +7642,11 @@ mod tests { } /// checkout_latest keeps the handle on latest, so a completed backfill - /// must still establish its post-fill baseline -- strictly later than the - /// checkout's own, or a pre-fill cache could still serve. + /// must retain the checkout timestamp and add its exact published version. #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn test_checkout_latest_during_submission_keeps_the_fence() { - let seen_min_timestamp = Arc::new(std::sync::Mutex::new(None::)); - let saw = seen_min_timestamp.clone(); + let seen_headers = Arc::new(std::sync::Mutex::new(None::)); + let saw = seen_headers.clone(); let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); @@ -7353,10 +7670,7 @@ mod tests { .body(refresh_done("j-11")) .unwrap(), "/v1/table/my_table/count_rows/" => { - *saw.lock().unwrap() = request - .headers() - .get("x-lancedb-min-timestamp") - .map(|v| v.to_str().unwrap().to_string()); + *saw.lock().unwrap() = Some(request.headers().clone()); http::Response::builder() .status(200) .body("1".to_string()) @@ -7377,25 +7691,18 @@ mod tests { .await .unwrap(); table.checkout_latest().await.unwrap(); - let after_checkout = SystemTime::now(); - // Real separation between the checkout baseline and completion. - tokio::time::sleep(std::time::Duration::from_millis(50)).await; release_tx.send(()).unwrap(); let job = submit.await.unwrap().unwrap(); job.wait().await.unwrap(); table.count_rows(None).await.unwrap(); - let header = seen_min_timestamp - .lock() - .unwrap() - .clone() - .expect("no baseline"); - let sent: SystemTime = chrono::DateTime::parse_from_rfc3339(&header) - .unwrap() - .into(); - assert!( - sent > after_checkout, - "baseline {header} did not advance past the checkout" + let headers = seen_headers.lock().unwrap().clone().expect("no request"); + assert!(headers.contains_key("x-lancedb-min-timestamp")); + assert_eq!( + headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("8") ); } @@ -8853,6 +9160,50 @@ mod tests { assert_ne!(Arc::as_ptr(&schema3), Arc::as_ptr(&schema1)); } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_schema_fetch_does_not_cross_checkout_generation() { + let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); + let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); + let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); + let arrived_tx = Arc::new(std::sync::Mutex::new(arrived_tx)); + let table = Table::new_with_handler("my_table", move |request| { + let body = request_body_json(&request); + let pinned = body["version"].as_u64() == Some(5); + if !pinned { + arrived_tx.lock().unwrap().send(()).unwrap(); + release_rx + .lock() + .unwrap() + .recv_timeout(Duration::from_secs(10)) + .unwrap(); + } + let field = if pinned { "pinned" } else { "latest" }; + http::Response::builder() + .status(200) + .body(format!( + r#"{{"version":5,"schema":{{"fields":[{{"name":"{field}","type":{{"type":"int32"}},"nullable":false}}]}}}}"# + )) + .unwrap() + }); + + let schema_fetch = tokio::spawn({ + let table = table.clone(); + async move { table.schema().await } + }); + tokio::task::spawn_blocking(move || { + arrived_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + table.checkout(5).await.unwrap(); + release_tx.send(()).unwrap(); + + let schema = schema_fetch.await.unwrap().unwrap(); + assert!(schema.field_with_name("pinned").is_ok()); + let cached = table.schema().await.unwrap(); + assert!(cached.field_with_name("pinned").is_ok()); + } + /// Test that schema cache is invalidated after checkout_latest #[tokio::test] async fn test_schema_cache_invalidation_on_checkout_latest() { @@ -10087,6 +10438,7 @@ mod tests { min_version: None, checkout_baseline: Some(baseline), min_read_version: None, + ..FreshnessState::default() }; assert_eq!(compute_min_timestamp(&state, None, now), Some(baseline)); @@ -10112,6 +10464,7 @@ mod tests { min_version: None, checkout_baseline: Some(baseline), min_read_version: None, + ..FreshnessState::default() }; assert_eq!( compute_min_timestamp(&state, Some(Duration::from_secs(10)), now), @@ -10124,6 +10477,7 @@ mod tests { min_version: None, checkout_baseline: Some(recent_baseline), min_read_version: None, + ..FreshnessState::default() }; assert_eq!( compute_min_timestamp(&state, Some(Duration::from_secs(60)), now), @@ -10200,6 +10554,110 @@ mod tests { assert!(!headers.contains_key("x-lancedb-min-version")); } + #[tokio::test] + async fn test_checkout_disables_read_consistency_interval() { + let (handler, captured) = capturing_handler(|path| match path { + "/v1/table/my_table/describe/" => r#"{"version":5,"schema":{"fields":[]}}"#.to_string(), + "/v1/table/my_table/count_rows/" => "42".to_string(), + _ => panic!("unexpected path: {}", path), + }); + let table = + Table::new_with_handler_and_interval("my_table", handler, Some(Duration::from_secs(0))); + + table.checkout(5).await.unwrap(); + table.count_rows(None).await.unwrap(); + + let headers = captured.lock().unwrap().clone().unwrap(); + assert!(!headers.contains_key("x-lancedb-min-timestamp")); + assert!(!headers.contains_key("x-lancedb-min-version")); + assert!(!headers.contains_key("x-lancedb-min-read-version")); + } + + #[tokio::test] + async fn test_read_snapshot_keeps_selector_and_freshness_generation_bound() { + let table = RemoteTable::new_mock_with_consistency_interval( + "my_table".to_string(), + |_| { + http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + }, + Some(Duration::ZERO), + ); + + let latest = table.snapshot_read_state().await; + table.checkout(5).await.unwrap(); + + let latest_request = latest + .freshness + .apply( + table + .client + .post("/v1/table/my_table/count_rows/") + .json(&serde_json::json!({ "version": latest.version })), + ) + .build() + .unwrap(); + assert!(request_body_json(&latest_request)["version"].is_null()); + assert!(latest_request.headers().contains_key(MIN_TIMESTAMP_HEADER)); + + let pinned = table.snapshot_read_state().await; + let pinned_request = pinned + .freshness + .apply( + table + .client + .post("/v1/table/my_table/count_rows/") + .json(&serde_json::json!({ "version": pinned.version })), + ) + .build() + .unwrap(); + assert_eq!(request_body_json(&pinned_request)["version"], 5); + assert!(!pinned_request.headers().contains_key(MIN_TIMESTAMP_HEADER)); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_cancelled_checkout_keeps_latest_freshness_enabled() { + let (described_tx, described_rx) = std::sync::mpsc::channel::<()>(); + let table = Arc::new(RemoteTable::new_mock_with_consistency_interval( + "my_table".to_string(), + move |_| { + described_tx.send(()).unwrap(); + http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + }, + Some(Duration::from_secs(0)), + )); + + let version_guard = table.version.write().await; + let checkout = tokio::spawn({ + let table = table.clone(); + async move { table.checkout(5).await } + }); + tokio::task::spawn_blocking(move || { + described_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + for _ in 0..100 { + if table.freshness.lock().unwrap().pinned { + break; + } + tokio::task::yield_now().await; + } + assert!(!checkout.is_finished()); + + checkout.abort(); + assert!(checkout.await.unwrap_err().is_cancelled()); + drop(version_guard); + + assert_eq!(*table.version.read().await, None); + assert!(table.snapshot_freshness_headers().min_timestamp.is_some()); + } + #[tokio::test] async fn test_freshness_positive_interval_sends_now_minus_interval() { let (handler, captured) = capturing_handler(|_| "42".to_string()); @@ -10268,6 +10726,62 @@ mod tests { ); } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_inflight_write_result_cannot_cross_checkout_generation() { + let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); + let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); + let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); + let arrived_tx = Arc::new(std::sync::Mutex::new(arrived_tx)); + let count_headers = Arc::new(std::sync::Mutex::new(None)); + let captured = count_headers.clone(); + let table = + Table::new_with_handler("my_table", move |request| match request.url().path() { + "/v1/table/my_table/update/" => { + arrived_tx.lock().unwrap().send(()).unwrap(); + release_rx + .lock() + .unwrap() + .recv_timeout(Duration::from_secs(10)) + .unwrap(); + http::Response::builder() + .status(200) + .body(r#"{"rows_updated":1,"version":100}"#.to_string()) + .unwrap() + } + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap(), + "/v1/table/my_table/count_rows/" => { + *captured.lock().unwrap() = Some(request.headers().clone()); + http::Response::builder() + .status(200) + .body("1".to_string()) + .unwrap() + } + path => panic!("unexpected path: {path}"), + }); + + let update = tokio::spawn({ + let table = table.clone(); + async move { table.update().column("a", "a + 1").execute().await } + }); + tokio::task::spawn_blocking(move || { + arrived_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + table.checkout(5).await.unwrap(); + release_tx.send(()).unwrap(); + update.await.unwrap().unwrap(); + table.count_rows(None).await.unwrap(); + + let headers = count_headers.lock().unwrap(); + let headers = headers.as_ref().unwrap(); + assert!(!headers.contains_key("x-lancedb-min-version")); + assert!(!headers.contains_key("x-lancedb-min-read-version")); + } + /// A handler that records every request's headers and answers each read with /// an `x-lancedb-version` response header taken from `versions` (by call /// index, saturating at the last entry). An empty string means "no header". @@ -10314,6 +10828,164 @@ mod tests { ); } + #[tokio::test] + async fn test_schema_response_advances_read_watermark() { + let requests = Arc::new(std::sync::Mutex::new(Vec::new())); + let captured = requests.clone(); + let table = Table::new_with_handler("my_table", move |request| { + captured + .lock() + .unwrap() + .push((request.url().path().to_string(), request.headers().clone())); + match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .header("x-lancedb-version", "100") + .body(r#"{"version":100,"schema":{"fields":[]}}"#.to_string()) + .unwrap(), + "/v1/table/my_table/count_rows/" => http::Response::builder() + .status(200) + .body("42".to_string()) + .unwrap(), + path => panic!("unexpected path: {path}"), + } + }); + + table.schema().await.unwrap(); + assert_eq!(table.count_rows(None).await.unwrap(), 42); + + let requests = requests.lock().unwrap(); + let count_headers = &requests[1].1; + assert_eq!( + count_headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("100") + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn test_inflight_schema_response_cannot_cross_checkout_generation() { + let (release_tx, release_rx) = std::sync::mpsc::channel::<()>(); + let release_rx = Arc::new(std::sync::Mutex::new(release_rx)); + let (arrived_tx, arrived_rx) = std::sync::mpsc::channel::<()>(); + let arrived_tx = Arc::new(std::sync::Mutex::new(arrived_tx)); + let count_headers = Arc::new(std::sync::Mutex::new(None)); + let captured = count_headers.clone(); + let table = + Table::new_with_handler("my_table", move |request| match request.url().path() { + "/v1/table/my_table/describe/" => { + let body = request_body_json(&request); + if body["version"].is_null() { + arrived_tx.lock().unwrap().send(()).unwrap(); + release_rx + .lock() + .unwrap() + .recv_timeout(Duration::from_secs(10)) + .unwrap(); + http::Response::builder() + .status(200) + .header("x-lancedb-version", "100") + .body(r#"{"version":100,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + } else { + http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.to_string()) + .unwrap() + } + } + "/v1/table/my_table/count_rows/" => { + *captured.lock().unwrap() = Some(request.headers().clone()); + http::Response::builder() + .status(200) + .body("1".to_string()) + .unwrap() + } + path => panic!("unexpected path: {path}"), + }); + + let schema = tokio::spawn({ + let table = table.clone(); + async move { table.schema().await } + }); + tokio::task::spawn_blocking(move || { + arrived_rx.recv_timeout(Duration::from_secs(10)).unwrap() + }) + .await + .unwrap(); + table.checkout(5).await.unwrap(); + release_tx.send(()).unwrap(); + schema.await.unwrap().unwrap(); + table.count_rows(None).await.unwrap(); + + assert!( + !count_headers + .lock() + .unwrap() + .as_ref() + .unwrap() + .contains_key("x-lancedb-min-read-version") + ); + } + + #[tokio::test] + async fn test_streaming_write_uses_and_advances_read_watermark() { + let data = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])), + vec![Arc::new(Int32Array::from(vec![1, 2, 3]))], + ) + .unwrap(); + let describe_body = serde_json::to_string(&json!({ + "version": 7, + "schema": JsonSchema::try_from(data.schema().as_ref()).unwrap(), + })) + .unwrap(); + let requests = Arc::new(std::sync::Mutex::new(Vec::new())); + let captured = requests.clone(); + let table = Table::new_with_handler("my_table", move |request| { + captured + .lock() + .unwrap() + .push((request.url().path().to_string(), request.headers().clone())); + match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(describe_body.clone()) + .unwrap(), + "/v1/table/my_table/insert/" => http::Response::builder() + .status(200) + .header("x-lancedb-version", "8") + .body(r#"{"version":8}"#.to_string()) + .unwrap(), + "/v1/table/my_table/count_rows/" => http::Response::builder() + .status(200) + .body("3".to_string()) + .unwrap(), + path => panic!("unexpected path: {path}"), + } + }); + + assert_eq!(table.add(data).execute().await.unwrap().version, 8); + assert_eq!(table.count_rows(None).await.unwrap(), 3); + + let requests = requests.lock().unwrap(); + let insert_headers = &requests[1].1; + assert_eq!( + insert_headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("7") + ); + let count_headers = &requests[2].1; + assert_eq!( + count_headers + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("8") + ); + } + #[tokio::test] async fn test_read_version_watermark_keeps_max() { // Server reports 100 then a stale 50; the watermark must not regress. @@ -10729,6 +11401,86 @@ mod tests { ); } + #[tokio::test] + async fn test_main_only_metadata_is_unfenced_from_branch_timeline() { + let requests = Arc::new(std::sync::Mutex::new(HashMap::new())); + let captured = requests.clone(); + let saw_delete_response_floor = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let saw_high_floor = saw_delete_response_floor.clone(); + let table = RemoteTable::new_mock( + "my_table".to_string(), + move |request| { + let path = request.url().path().to_string(); + captured + .lock() + .unwrap() + .insert(path.clone(), request.headers().clone()); + match path.as_str() { + "/v1/table/my_table/count_rows/" => { + saw_high_floor.store( + request + .headers() + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()) + == Some("100"), + std::sync::atomic::Ordering::SeqCst, + ); + http::Response::builder() + .status(200) + .header("x-lancedb-version", "2") + .body("1".to_string()) + .unwrap() + } + "/v1/table/my_table/tags/list/" => http::Response::builder() + .status(200) + .body("{}".to_string()) + .unwrap(), + "/v1/table/my_table/tags/version/" => http::Response::builder() + .status(200) + .body(r#"{"version":1}"#.to_string()) + .unwrap(), + "/v1/table/my_table/tags/delete/" => http::Response::builder() + .status(200) + .header("x-lancedb-version", "100") + .body("{}".to_string()) + .unwrap(), + "/v1/table/my_table/branches/list/" => http::Response::builder() + .status(200) + .body(r#"{"branches":{}}"#.to_string()) + .unwrap(), + path => panic!("unexpected path: {path}"), + } + }, + None, + ); + let branch = table.with_branch(Some("exp".to_string())); + + branch.count_rows(None).await.unwrap(); + let mut tags = branch.tags().await.unwrap(); + tags.list().await.unwrap(); + tags.get_version("v1").await.unwrap(); + tags.delete("v1").await.unwrap(); + branch.list_branches().await.unwrap(); + branch.count_rows(None).await.unwrap(); + + let requests = requests.lock().unwrap(); + for path in [ + "/v1/table/my_table/tags/list/", + "/v1/table/my_table/tags/version/", + "/v1/table/my_table/tags/delete/", + "/v1/table/my_table/branches/list/", + ] { + assert!( + !requests[path].contains_key("x-lancedb-min-read-version"), + "{path} inherited the branch timeline" + ); + } + assert!( + !saw_delete_response_floor.load(std::sync::atomic::Ordering::SeqCst), + "tag deletion contaminated the branch timeline" + ); + } + #[tokio::test] async fn test_delete_branch() { let table = Table::new_with_handler("my_table", |request| { @@ -10820,6 +11572,50 @@ mod tests { assert!(result.main_version_after.is_none()); } + #[tokio::test] + async fn test_successful_cherry_pick_advances_main_read_watermark() { + let count_headers = Arc::new(std::sync::Mutex::new(None)); + let captured = count_headers.clone(); + let table = + Table::new_with_handler("my_table", move |request| match request.url().path() { + "/v1/table/my_table/branches/cherry_pick/" => { + let response = serde_json::json!({ + "status": "cherryPicked", + "diff": serde_json::from_str::(sample_branch_diff_json()) + .unwrap(), + "preview": { "promotedColumns": ["tag"] }, + "mainVersionAfter": 2 + }); + http::Response::builder() + .status(200) + .body(response.to_string()) + .unwrap() + } + "/v1/table/my_table/count_rows/" => { + *captured.lock().unwrap() = Some(request.headers().clone()); + http::Response::builder() + .status(200) + .body("1".to_string()) + .unwrap() + } + path => panic!("unexpected path: {path}"), + }); + + let result = table.cherry_pick("exp", false).await.unwrap(); + assert_eq!(result.status, crate::table::CherryPickStatus::CherryPicked); + table.count_rows(None).await.unwrap(); + assert_eq!( + count_headers + .lock() + .unwrap() + .as_ref() + .unwrap() + .get("x-lancedb-min-read-version") + .and_then(|value| value.to_str().ok()), + Some("2") + ); + } + #[tokio::test] async fn test_cherry_pick_failed_returns_ok_with_body() { let table = Table::new_with_handler("my_table", |request| { diff --git a/rust/lancedb/src/remote/table/blobs.rs b/rust/lancedb/src/remote/table/blobs.rs index 387d3c6dc..597b41782 100644 --- a/rust/lancedb/src/remote/table/blobs.rs +++ b/rust/lancedb/src/remote/table/blobs.rs @@ -6,6 +6,7 @@ use std::ops::Range; use std::sync::Arc; use std::sync::atomic::{AtomicBool, Ordering}; +use std::time::Duration; use arrow_array::{Array, LargeBinaryArray}; use arrow_schema::DataType; @@ -20,7 +21,7 @@ use crate::error::Result; use crate::remote::client::{HttpSend, RequestResultExt, RestfulLanceDbClient}; use crate::table::BaseTable; -use super::{FreshnessHeaders, RemoteTable}; +use super::{FreshnessHeaders, FreshnessState, RemoteTable, freshness_headers_snapshot}; #[derive(Debug, Clone, Copy)] enum RangeRequestMode { @@ -43,7 +44,10 @@ struct TableBlobRangeRequester { path: String, version: Option, branch: Option, - freshness: FreshnessHeaders, + freshness: Arc>, + parent_freshness: Arc>, + parent_freshness_request: FreshnessHeaders, + read_consistency_interval: Option, } #[async_trait::async_trait] @@ -53,8 +57,9 @@ impl BlobRangeRequester for TableBlobRangeRequester { range_header: &str, mode: RangeRequestMode, ) -> Result<(String, Response)> { - let mut request = self - .freshness + let freshness_request = + freshness_headers_snapshot(&self.freshness, self.read_consistency_interval); + let mut request = freshness_request .apply(self.client.get(&self.path)) .header(header::RANGE, range_header); if let Some(version) = self.version { @@ -71,6 +76,9 @@ impl BlobRangeRequester for TableBlobRangeRequester { return Ok((request_id, response)); } let response = self.client.check_response(&request_id, response).await?; + freshness_request.observe_headers(&self.freshness, response.headers()); + self.parent_freshness_request + .observe_headers(&self.parent_freshness, response.headers()); Ok((request_id, response)) } } @@ -361,18 +369,21 @@ impl RemoteTable { message: "fetch_blobs is not supported on this LanceDB Cloud server".into(), }); } - let version = self.current_version().await; + let read_snapshot = self.snapshot_read_state().await; let mut body = serde_json::json!({ - "version": version, + "version": read_snapshot.version, "column": column, "row_ids": row_ids, }); self.apply_branch_body(&mut body); let request = self - .post_read(&format!("/v1/table/{}/fetch_blobs/", self.identifier)) + .client + .post(&format!("/v1/table/{}/fetch_blobs/", self.identifier)) .json(&body); - let (request_id, response) = self.send(request, true).await?; + let (request_id, response) = self + .send_with_freshness(request, true, read_snapshot.freshness) + .await?; let mut stream = self.read_arrow_response(&request_id, response).await?; let mut blob_chunks: Vec> = Vec::new(); @@ -448,8 +459,7 @@ impl RemoteTable { }); } - let version = self.current_version().await; - let freshness = self.snapshot_freshness_headers(); + let read_snapshot = self.snapshot_read_state().await; let encoded_column = urlencoding::encode(column); let requesters = row_ids .iter() @@ -461,9 +471,12 @@ impl RemoteTable { let requester: Arc = Arc::new(TableBlobRangeRequester { client: self.client.clone(), path, - version, + version: read_snapshot.version, branch: self.branch.clone(), - freshness, + freshness: Arc::new(std::sync::Mutex::new(read_snapshot.freshness_state)), + parent_freshness: self.freshness.clone(), + parent_freshness_request: read_snapshot.freshness, + read_consistency_interval: self.client.read_consistency_interval, }); requester }) @@ -685,6 +698,46 @@ mod tests { assert!(requests.lock().unwrap().contains(&"bytes=5-11".to_string())); } + #[tokio::test] + async fn remote_blob_file_keeps_the_open_timeline_after_parent_checkout() { + let range_requests = Arc::new(StdMutex::new(Vec::new())); + let captured = range_requests.clone(); + let table = RemoteTable::new_mock( + "my_table".to_string(), + move |request| match request.url().path() { + "/v1/table/my_table/describe/" => http::Response::builder() + .status(200) + .body(r#"{"version":5,"schema":{"fields":[]}}"#.as_bytes().to_vec()) + .unwrap(), + "/v1/table/my_table/blob/image/10/bytes" => { + captured.lock().unwrap().push(( + request.url().query().unwrap_or_default().to_string(), + request.headers().clone(), + )); + range_response(&request, PAYLOAD) + } + path => panic!("unexpected path: {path}"), + }, + Some(Version::new(0, 5, 0)), + ); + + table.checkout(5).await.unwrap(); + let file = table + .fetch_blob_files_impl("image", &[10]) + .await + .unwrap() + .pop() + .flatten() + .unwrap(); + table.checkout_latest().await.unwrap(); + file.read_range(5..12).await.unwrap(); + + let requests = range_requests.lock().unwrap(); + let (query, headers) = requests.last().unwrap(); + assert!(query.contains("version=5")); + assert!(!headers.contains_key("x-lancedb-min-timestamp")); + } + #[tokio::test] async fn remote_blob_file_reuses_sequential_response_until_seek() { let requests = Arc::new(StdMutex::new(Vec::new())); diff --git a/rust/lancedb/src/remote/table/insert.rs b/rust/lancedb/src/remote/table/insert.rs index a4a28a9c6..4e0e0d666 100644 --- a/rust/lancedb/src/remote/table/insert.rs +++ b/rust/lancedb/src/remote/table/insert.rs @@ -24,7 +24,10 @@ use lance::io::exec::utils::InstrumentedRecordBatchStreamAdapter; use crate::Error; use crate::remote::ARROW_STREAM_CONTENT_TYPE; use crate::remote::client::{HttpSend, RestfulLanceDbClient, Sender}; -use crate::remote::table::{MergeInsertRequest, REQUEST_TIMEOUT_HEADER, RemoteTable}; +use crate::remote::table::{ + FreshnessHeaders, FreshnessState, MergeInsertRequest, REQUEST_TIMEOUT_HEADER, RemoteTable, + freshness_headers_snapshot, +}; use crate::table::datafusion::insert::COUNT_SCHEMA; use crate::table::write_progress::WriteProgressTracker; use crate::table::{AddResult, MergeResult}; @@ -54,6 +57,38 @@ pub enum WriteResult { Merge(MergeResult), } +#[derive(Debug, Clone, Default)] +struct WriteFreshness { + state: Option>>, + read_consistency_interval: Option, +} + +impl WriteFreshness { + fn prepare( + &self, + request: reqwest::RequestBuilder, + ) -> (reqwest::RequestBuilder, Option) { + match &self.state { + Some(state) => { + let freshness_request = + freshness_headers_snapshot(state, self.read_consistency_interval); + (freshness_request.apply(request), Some(freshness_request)) + } + None => (request, None), + } + } + + fn observe( + &self, + freshness_request: Option, + headers: &reqwest::header::HeaderMap, + ) { + if let (Some(state), Some(freshness_request)) = (&self.state, freshness_request) { + freshness_request.observe_headers(state, headers); + } + } +} + /// ExecutionPlan for streaming a write (add or merge_insert) to a remote /// LanceDB table. /// @@ -71,6 +106,7 @@ pub struct RemoteWriteExec { table_name: String, identifier: String, client: RestfulLanceDbClient, + freshness: WriteFreshness, input: Arc, op: WriteOp, properties: Arc, @@ -170,6 +206,7 @@ impl RemoteWriteExec { table_name, identifier, client, + freshness: WriteFreshness::default(), input, op, properties: Arc::new(properties), @@ -183,6 +220,18 @@ impl RemoteWriteExec { } } + pub(super) fn with_freshness( + mut self, + state: Arc>, + read_consistency_interval: Option, + ) -> Self { + self.freshness = WriteFreshness { + state: Some(state), + read_consistency_interval, + }; + self + } + /// Get the add result after execution, if this exec ran an insert. pub fn add_result(&self) -> Option { match self @@ -285,6 +334,7 @@ impl RemoteWriteExec { /// each threading the same handful of arguments. struct PartRequestCtx<'a, S: HttpSend> { client: &'a RestfulLanceDbClient, + freshness: &'a WriteFreshness, identifier: &'a str, table_name: &'a str, upload_id: &'a str, @@ -352,7 +402,11 @@ impl PartRequestCtx<'_, S> { } /// Build the `/insert` request for a single multipart part. - fn build_part_request(&self, part_id: &str, body: reqwest::Body) -> reqwest::RequestBuilder { + fn build_part_request( + &self, + part_id: &str, + body: reqwest::Body, + ) -> (reqwest::RequestBuilder, Option) { let mut request = self .client .post(&format!("/v1/table/{}/insert/", self.identifier)) @@ -368,12 +422,16 @@ impl PartRequestCtx<'_, S> { if let Some(b) = self.branch { request = request.query(&[("branch", b)]); } - request.body(body) + self.freshness.prepare(request.body(body)) } /// Send a single part's request and drain the response, mapping HTTP and /// table-not-found errors into `DataFusionError`. - async fn send_part_request(&self, request: reqwest::RequestBuilder) -> DataFusionResult<()> { + async fn send_part_request( + &self, + request: reqwest::RequestBuilder, + freshness_request: Option, + ) -> DataFusionResult<()> { let (request_id, response) = self .client .send(request) @@ -388,6 +446,8 @@ impl PartRequestCtx<'_, S> { .check_response(&request_id, response) .await .map_err(|e| DataFusionError::External(Box::new(e)))?; + self.freshness + .observe(freshness_request, response.headers()); response.bytes().await.map_err(|e| { DataFusionError::External(Box::new(Error::Http { source: Box::new(e), @@ -419,7 +479,7 @@ impl PartRequestCtx<'_, S> { let body = reqwest::Body::wrap_stream(chunk_rx); let part_id = uuid::Uuid::new_v4().to_string(); - let request = self.build_part_request(&part_id, body); + let (request, freshness_request) = self.build_part_request(&part_id, body); // Measured from just before the request is sent, matching the window the // client read timeout applies to the upload. @@ -495,7 +555,7 @@ impl PartRequestCtx<'_, S> { Ok::(input_ended) }; - let send = self.send_part_request(request); + let send = self.send_part_request(request, freshness_request); // `join!` rather than `tokio::spawn`: the producer borrows `input` (and // `schema`), so it cannot satisfy the `'static` bound a spawned task @@ -569,7 +629,7 @@ impl ExecutionPlan for RemoteWriteExec { // Building a fresh exec (with a new, empty `result`) is what makes the // outer rescannable retry loop work: `reset_state()` clears the captured // result so a re-execution starts clean. - Ok(Arc::new(Self::new_inner( + let mut exec = Self::new_inner( self.table_name.clone(), self.identifier.clone(), self.client.clone(), @@ -580,7 +640,9 @@ impl ExecutionPlan for RemoteWriteExec { self.branch.clone(), self.max_bytes_per_request, self.max_request_duration, - ))) + ); + exec.freshness = self.freshness.clone(); + Ok(Arc::new(exec)) } fn execute( @@ -613,6 +675,7 @@ impl ExecutionPlan for RemoteWriteExec { &self.metrics, )); let client = self.client.clone(); + let freshness = self.freshness.clone(); let identifier = self.identifier.clone(); let op = self.op.clone(); let result_slot = self.result.clone(); @@ -634,6 +697,7 @@ impl ExecutionPlan for RemoteWriteExec { let overwrite = matches!(op, WriteOp::Insert { overwrite: true }); let ctx = PartRequestCtx { client: &client, + freshness: &freshness, identifier: &identifier, table_name: &table_name, upload_id, @@ -688,7 +752,7 @@ impl ExecutionPlan for RemoteWriteExec { let (error_tx, mut error_rx) = tokio::sync::oneshot::channel(); let body = Self::stream_as_http_body(input_stream, error_tx, tracker)?; - let request = request.body(body); + let (request, freshness_request) = freshness.prepare(request.body(body)); let result: DataFusionResult<(String, _)> = async { let (request_id, response) = client @@ -708,6 +772,7 @@ impl ExecutionPlan for RemoteWriteExec { .check_response(&request_id, response) .await .map_err(|e| DataFusionError::External(Box::new(e)))?; + freshness.observe(freshness_request, response.headers()); Ok((request_id, response)) }