Merge remote-tracking branch 'origin/main' into gatekeeper/fix-2820-1

This commit is contained in:
Gatefixer
2026-08-24 17:41:48 +00:00
77 changed files with 6837 additions and 544 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "lancedb-python"
version = "0.38.0-beta.3"
version = "0.38.0-beta.6"
publish = false
edition.workspace = true
description = "Python bindings for LanceDB"
+9
View File
@@ -29,9 +29,15 @@ from .functions import (
FunctionRegistrationRequest as FunctionRegistrationRequest,
FunctionVersion as FunctionVersion,
PythonRuntimeSpec as PythonRuntimeSpec,
RefreshColumnResult as RefreshColumnResult,
UdfDefinition as UdfDefinition,
udf as udf,
)
from .materialized_view import (
AsyncMaterializedView,
MaterializedView,
MaterializedViewDefinition,
)
from .table import AsyncTable, Table
from .types import BaseTokenizerType
from ._lancedb import Session
@@ -506,6 +512,9 @@ async def connect_async(
__all__ = [
"AsyncMaterializedView",
"MaterializedView",
"MaterializedViewDefinition",
"connect",
"connect_async",
"tokenize",
+21 -10
View File
@@ -147,7 +147,7 @@ class Connection(object):
limit: Optional[int],
) -> list[str]: ... # Deprecated: Use list_tables instead
def job(self, job_id: str) -> Job: ...
async def create_function_async(self, request_json: str) -> FunctionJob: ...
async def create_function_async(self, request_json: str) -> Job: ...
async def get_function(self, name: str, version: str) -> str: ...
async def list_jobs(self) -> List[JobInfo]: ...
async def get_job(self, job_id: str) -> Optional[JobDescription]: ...
@@ -197,6 +197,15 @@ class Connection(object):
cur_namespace_path: Optional[List[str]] = None,
new_namespace_path: Optional[List[str]] = None,
) -> None: ...
async def create_materialized_view(
self,
name: str,
source: str,
projections: Optional[List[Tuple[str, str]]] = None,
filter: Optional[str] = None,
limit: Optional[int] = None,
) -> Table: ...
async def list_materialized_views(self) -> List[str]: ...
async def drop_table(
self, name: str, namespace_path: Optional[List[str]] = None
) -> None: ...
@@ -225,14 +234,7 @@ class Job:
@property
def id(self) -> Optional[str]: ...
async def status(self) -> str: ...
async def wait(self) -> None: ...
async def cancel(self) -> None: ...
class FunctionJob:
@property
def id(self) -> Optional[str]: ...
async def status(self) -> str: ...
async def wait(self) -> str: ...
async def wait(self) -> Optional[str]: ...
async def cancel(self) -> None: ...
class JobInfo:
@@ -355,6 +357,9 @@ class Table:
) -> AddColumnsResult: ...
async def refresh_column(self, column: str) -> RefreshColumnResult: ...
async def refresh_column_async(self, column: str) -> Job: ...
async def refresh_materialized_view(
self, full: bool = False, source_version: Optional[int] = None
) -> RefreshMaterializedViewResult: ...
async def add_columns_with_schema(self, schema: pa.Schema) -> AddColumnsResult: ...
async def alter_columns(
self, columns: list[dict[str, Any]]
@@ -420,7 +425,7 @@ class Branches:
async def checkout(self, name: str, version: Optional[int] = None) -> Table: ...
async def delete(self, name: str) -> None: ...
async def diff(self, from_branch: str) -> Dict[str, Any]: ...
async def merge(
async def cherry_pick(
self, from_branch: str, dry_run: bool = False
) -> Dict[str, Any]: ...
@@ -704,6 +709,12 @@ class RefreshColumnResult:
rows_filled: int
version: int
class RefreshMaterializedViewResult:
mode: str
rows_written: int
source_version: int
version: int
class AlterColumnsResult:
version: int
+168 -2
View File
@@ -46,7 +46,13 @@ from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
from . import __version__
from ._lancedb import connect as lancedb_connect # type: ignore
from .functions import FunctionVersion, UdfDefinition
from .job import AsyncJob, Job, _function_job
from .job import AsyncJob, Job, _typed_job
from .materialized_view import (
AsyncMaterializedView,
MaterializedView,
SelectArg,
normalize_select,
)
from .table import (
AsyncTable,
LanceTable,
@@ -510,6 +516,70 @@ class DBConnection(EnforceOverrides):
"""
raise NotImplementedError
def create_materialized_view(
self,
name: str,
source: str,
*,
select: SelectArg = None,
where: Optional[str] = None,
limit: Optional[int] = None,
) -> MaterializedView:
"""Define a materialized view named ``name`` over the table ``source``.
The view is created empty, with the query recorded in its schema
metadata; ``view.refresh()`` computes the rows. The view is a normal
table: it can be queried, indexed and searched, and it appears in
``table_names``. Local databases only.
The source table must have stable row ids (create it with the
``new_table_enable_stable_row_ids`` storage option): they keep the
view's provenance valid across source compactions, and cannot be
enabled after a table exists.
Parameters
----------
name: str
The name of the view.
source: str
The name of the source table, in this database.
select: list or dict, optional
The view's columns: column names, ``(alias, SQL expression)``
pairs, or a dict of the same. Omitting it selects every source
column, expanded against the source schema at creation time.
where: str, optional
SQL predicate; only matching source rows appear in the view.
limit: int, optional
Cap the view at this many rows, in materialization order.
Returns
-------
MaterializedView
"""
raise NotImplementedError(
"materialized views are not supported on this connection type"
)
def open_materialized_view(self, name: str) -> MaterializedView:
"""Open the materialized view named ``name``.
Raises ``ValueError`` if the table exists but is not a materialized
view.
"""
raise NotImplementedError(
"materialized views are not supported on this connection type"
)
def list_materialized_views(self) -> List[str]:
"""The names of the materialized views in this database.
Found by reading every table's schema, so this costs an open per
table.
"""
raise NotImplementedError(
"materialized views are not supported on this connection type"
)
def drop_table(self, name: str, namespace_path: Optional[List[str]] = None):
"""Drop a table from the database.
@@ -1136,6 +1206,58 @@ class LanceDBConnection(DBConnection):
tbl.checkout(version)
return tbl
@override
def create_materialized_view(
self,
name: str,
source: str,
*,
select: SelectArg = None,
where: Optional[str] = None,
limit: Optional[int] = None,
) -> MaterializedView:
"""Define a materialized view named ``name`` over the table ``source``.
See
[DBConnection.create_materialized_view][lancedb.DBConnection.create_materialized_view].
Examples
--------
>>> import lancedb
>>> db = lancedb.connect(
... "./.lancedb",
... storage_options={"new_table_enable_stable_row_ids": "true"},
... )
>>> data = [{"name": "ada", "age": 36}, {"name": "kid", "age": 7}]
>>> table = db.create_table("people", data)
>>> view = db.create_materialized_view(
... "adults",
... "people",
... select=["name", ("shout", "upper(name)")],
... where="age >= 18",
... )
>>> result = view.refresh()
>>> result.rows_written
1
"""
LOOP.run(
self._conn.create_materialized_view(
name, source, select=select, where=where, limit=limit
)
)
return MaterializedView(self.open_table(name))
@override
def open_materialized_view(self, name: str) -> MaterializedView:
"""Open the materialized view named ``name``."""
view = MaterializedView(self.open_table(name))
view.definition
return view
@override
def list_materialized_views(self) -> List[str]:
"""The names of the materialized views in this database."""
return LOOP.run(self._conn.list_materialized_views())
def clone_table(
self,
target_table_name: str,
@@ -1906,6 +2028,50 @@ class AsyncConnection(object):
await tbl.checkout(version)
return tbl
async def create_materialized_view(
self,
name: str,
source: str,
*,
select: SelectArg = None,
where: Optional[str] = None,
limit: Optional[int] = None,
) -> AsyncMaterializedView:
"""Define a materialized view named ``name`` over the table ``source``.
See
[DBConnection.create_materialized_view][lancedb.DBConnection.create_materialized_view].
"""
inner = await self._inner.create_materialized_view(
name,
source,
projections=normalize_select(select),
filter=where,
limit=limit,
)
return AsyncMaterializedView(AsyncTable(inner))
async def open_materialized_view(self, name: str) -> AsyncMaterializedView:
"""Open the materialized view named ``name``.
Raises ``ValueError`` if the table exists but is not a materialized
view.
"""
if self.uri.startswith("db://"):
raise NotImplementedError(
"materialized views are supported only on local databases"
)
view = AsyncMaterializedView(await self.open_table(name))
await view.definition()
return view
async def list_materialized_views(self) -> List[str]:
"""The names of the materialized views in this database.
Found by reading every table's schema, so this costs an open per
table.
"""
return await self._inner.list_materialized_views()
async def clone_table(
self,
target_table_name: str,
@@ -2071,7 +2237,7 @@ class AsyncConnection(object):
inner = await self._inner.create_function_async(
definition.registration_request.to_canonical_json()
)
return _function_job(inner)
return _typed_job(inner, FunctionVersion.from_json)
async def get_function(self, name: str, *, version: str) -> FunctionVersion:
"""Open one exact immutable Function version from the remote catalog."""
+8 -2
View File
@@ -1,10 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Canonical values exchanged with LanceDB Enterprise Function services.
"""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.
``RefreshColumnResult`` is also the backend-neutral result of a local
expression-backed refresh job.
"""
from __future__ import annotations
@@ -460,7 +462,11 @@ class FunctionBinding(_RemoteValue):
class RefreshColumnResult(_RemoteValue):
"""Terminal result of a remote Function-column refresh Job."""
"""Terminal result of an expression-backed or Function-backed refresh Job.
Local jobs produce this value in process. LanceDB Cloud and Enterprise
decode the same value from the durable server-job terminal payload.
"""
rows_assigned: _UInt64
rows_failed: _UInt64
+28 -30
View File
@@ -5,12 +5,11 @@
import asyncio
from datetime import timedelta
from typing import Any, Generic, Optional, TypeVar, cast
from typing import Any, Callable, Generic, Optional, TypeVar, cast
from lancedb.background_loop import LOOP
from . import _lancedb
from .functions import FunctionVersion
T = TypeVar("T")
@@ -18,11 +17,18 @@ T = TypeVar("T")
class AsyncJob(Generic[T]):
"""A handle to an operation that may still be running.
The operation may already be complete when the handle is created.
The operation may already be complete when the handle is created. ``T``
is the endpoint's terminal result type; unit-result jobs resolve to
``None``.
"""
def __init__(self, inner: Optional[Any]):
def __init__(
self,
inner: Optional[Any],
result_decoder: Optional[Callable[[Any], T]] = None,
):
self._inner = inner
self._result_decoder = result_decoder
@property
def id(self) -> Optional[str]:
@@ -50,17 +56,21 @@ class AsyncJob(Generic[T]):
async def wait(self, timeout: Optional[timedelta] = None) -> T:
"""Wait until the operation reaches a terminal state.
Returns the endpoint's typed result, or ``None`` for a unit-result
job.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
if self._inner is None:
return cast(T, None)
if timeout is None:
return cast(T, await self._inner.wait())
return cast(
T,
await asyncio.wait_for(self._inner.wait(), timeout.total_seconds()),
)
result = await self._inner.wait()
else:
result = await asyncio.wait_for(self._inner.wait(), timeout.total_seconds())
if self._result_decoder is not None:
return self._result_decoder(result)
return cast(T, result)
async def cancel(self):
"""Request cancellation. Cancelling a finished operation is a no-op."""
@@ -70,7 +80,7 @@ class AsyncJob(Generic[T]):
class Job(Generic[T]):
"""Synchronous counterpart of `AsyncJob`."""
"""Synchronous counterpart of `AsyncJob` with the same result type."""
def __init__(self, inner: Optional[AsyncJob[T]]):
self._inner = inner
@@ -96,6 +106,9 @@ class Job(Generic[T]):
def wait(self, timeout: Optional[timedelta] = None) -> T:
"""Block until the operation reaches a terminal state.
Returns the endpoint's typed result, or ``None`` for a unit-result
job.
Raises `JobFailedError` if the operation failed, `JobCancelledError`
if it was cancelled, and `TimeoutError` if `timeout` elapses first.
"""
@@ -110,23 +123,8 @@ class Job(Generic[T]):
LOOP.run(self._inner.cancel())
class _FunctionJobAdapter:
def __init__(self, inner: "_lancedb.FunctionJob"):
self._inner = inner
@property
def id(self) -> Optional[str]:
return self._inner.id
async def status(self) -> str:
return await self._inner.status()
async def wait(self) -> FunctionVersion:
return FunctionVersion.from_json(await self._inner.wait())
async def cancel(self):
await self._inner.cancel()
def _function_job(inner: "_lancedb.FunctionJob") -> AsyncJob[FunctionVersion]:
return AsyncJob(_FunctionJobAdapter(inner))
def _typed_job(
inner: "_lancedb.Job", result_decoder: Callable[[str], T]
) -> AsyncJob[T]:
"""Bind an internal JSON-producing job to its public result model."""
return AsyncJob(inner, result_decoder)
+178
View File
@@ -0,0 +1,178 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Materialized views: tables defined by a query over a source table and
maintained by refresh. See ``DBConnection.create_materialized_view``."""
from __future__ import annotations
import json
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Dict, List, Optional, Sequence, Tuple, Union
from .background_loop import LOOP
if TYPE_CHECKING:
import pyarrow as pa
from ._lancedb import RefreshMaterializedViewResult
from .table import AsyncTable, LanceTable
DEFINITION_META_KEY = b"mv.definition"
SelectArg = Union[
str,
Sequence[Union[str, Tuple[str, str]]],
Dict[str, str],
None,
]
@dataclass
class MaterializedViewDefinition:
"""The query that defines a materialized view."""
source_table: str
"""Name of the source table, in the same database as the view."""
projections: List[Tuple[str, str]]
"""``(output column, SQL expression)`` pairs, in view schema order."""
filter: Optional[str] = None
"""SQL predicate selecting the source rows the view holds."""
limit: Optional[int] = None
"""Cap on the number of rows the view holds."""
inputs: List[str] = field(default_factory=list)
"""Source columns the projections and filter read."""
def _definition_from_schema(
schema: "pa.Schema", name: str
) -> MaterializedViewDefinition:
metadata = schema.metadata or {}
raw = metadata.get(DEFINITION_META_KEY)
if raw is None:
raise ValueError(f"Table '{name}' is not a materialized view")
value = json.loads(raw)
kind = value.get("kind")
if kind != "select":
raise NotImplementedError(
f"materialized view '{name}' is defined by '{kind}', which this "
"version of lancedb cannot refresh"
)
return MaterializedViewDefinition(
source_table=value["source_table"],
projections=[
(p["output"], p["expression"]) for p in value.get("projections", [])
],
filter=value.get("filter"),
limit=value.get("limit"),
inputs=value.get("inputs", []),
)
def _quote_identifier(name: str) -> str:
"""Quote a column name as a Lance SQL identifier (backticks)."""
escaped = name.replace("`", "``")
return f"`{escaped}`"
def normalize_select(select: SelectArg) -> Optional[List[Tuple[str, str]]]:
"""``select`` items may be a column name, an ``(alias, expression)`` pair,
or a dict of the same. A bare name projects itself and is quoted, so any
valid column name works; dict and pair entries are kept verbatim because
their right side is an expression.
A lone string is one column, not a sequence of its characters."""
if select is None:
return None
if isinstance(select, str):
select = [select]
if isinstance(select, dict):
return list(select.items())
normalized = []
for item in select:
if isinstance(item, str):
normalized.append((item, _quote_identifier(item)))
else:
alias, expression = item
normalized.append((alias, expression))
return normalized
class AsyncMaterializedView:
"""A handle on a materialized view: its table plus its definition.
Obtained from ``AsyncConnection.create_materialized_view`` or
``AsyncConnection.open_materialized_view``.
"""
def __init__(self, table: "AsyncTable"):
self._table = table
def __repr__(self) -> str:
return f"AsyncMaterializedView(name={self.name!r})"
@property
def name(self) -> str:
return self._table.name
@property
def table(self) -> "AsyncTable":
"""The view, as the table it is. Queries, indexes and search all
apply; writes are not blocked, but a rebuild replaces them."""
return self._table
async def definition(self) -> MaterializedViewDefinition:
"""The query that defines the view, read from its stored schema."""
return _definition_from_schema(await self._table.schema(), self.name)
async def refresh(
self, *, full: bool = False, source_version: Optional[int] = None
) -> "RefreshMaterializedViewResult":
"""Recompute the view from its source.
The refresh is incremental when the source's changes can be
reconciled into the view -- rows added, changed or removed since the
last one -- and otherwise rebuilds. ``full=True`` forces a rebuild;
``source_version`` refreshes to that source version instead of the
latest.
Concurrent refreshes of one view do not duplicate its rows. Two that
plan the same source rows conflict on commit, and the loser raises
rather than writing them a second time.
"""
return await self._table._inner.refresh_materialized_view(
full=full, source_version=source_version
)
class MaterializedView:
"""Synchronous variant of
[AsyncMaterializedView][lancedb.materialized_view.AsyncMaterializedView]."""
def __init__(self, table: "LanceTable"):
self._table = table
self._async = AsyncMaterializedView(table._table)
def __repr__(self) -> str:
return f"MaterializedView(name={self.name!r})"
@property
def name(self) -> str:
return self._table.name
@property
def table(self) -> "LanceTable":
"""The view, as the table it is."""
return self._table
@property
def definition(self) -> MaterializedViewDefinition:
"""The query that defines the view, read from its stored schema."""
return _definition_from_schema(self._table.schema, self.name)
def refresh(
self, *, full: bool = False, source_version: Optional[int] = None
) -> "RefreshMaterializedViewResult":
"""Recompute the view from its source. See
[AsyncMaterializedView.refresh][lancedb.materialized_view.AsyncMaterializedView.refresh]."""
return LOOP.run(self._async.refresh(full=full, source_version=source_version))
+68
View File
@@ -61,6 +61,11 @@ from lance_namespace import (
NamespaceExistsRequest,
TableExistsRequest,
)
from lancedb.materialized_view import (
AsyncMaterializedView,
MaterializedView,
SelectArg,
)
from lancedb.table import AsyncTable, LanceTable, Table
from lancedb.util import validate_table_name
from lancedb.common import DATA
@@ -619,6 +624,42 @@ class LanceNamespaceDBConnection(DBConnection):
tbl.checkout(version)
return tbl
@override
def create_materialized_view(
self,
name: str,
source: str,
*,
select: "SelectArg" = None,
where: Optional[str] = None,
limit: Optional[int] = None,
) -> "MaterializedView":
"""Define a materialized view over a table in the root namespace.
See
[DBConnection.create_materialized_view][lancedb.DBConnection.create_materialized_view].
"""
return MaterializedView(
self.open_table(
LOOP.run(
self._inner.create_materialized_view(
name, source, select=select, where=where, limit=limit
)
).name
)
)
@override
def open_materialized_view(self, name: str) -> "MaterializedView":
"""Open the materialized view named ``name``."""
view = MaterializedView(self.open_table(name))
view.definition
return view
@override
def list_materialized_views(self) -> List[str]:
"""The names of the materialized views in the root namespace."""
return LOOP.run(self._inner.list_materialized_views())
@override
def drop_table(self, name: str, namespace_path: Optional[List[str]] = None):
if namespace_path is None:
@@ -1141,6 +1182,33 @@ class AsyncLanceNamespaceDBConnection:
route_pushdown_to_rust=self._route_pushdown_to_rust,
)
async def create_materialized_view(
self,
name: str,
source: str,
*,
select: "SelectArg" = None,
where: Optional[str] = None,
limit: Optional[int] = None,
) -> "AsyncMaterializedView":
"""Define a materialized view over a table in the root namespace."""
view = await self._inner.create_materialized_view(
name, source, select=select, where=where, limit=limit
)
# Reopen through the namespace so the view's table carries the
# namespace client and pushdown configuration a bare inner table lacks.
return AsyncMaterializedView(await self.open_table(view.name))
async def open_materialized_view(self, name: str) -> "AsyncMaterializedView":
"""Open the materialized view named ``name``."""
view = AsyncMaterializedView(await self.open_table(name))
await view.definition()
return view
async def list_materialized_views(self) -> List[str]:
"""The names of the materialized views in the root namespace."""
return await self._inner.list_materialized_views()
async def drop_table(self, name: str, namespace_path: Optional[List[str]] = None):
"""Drop a table from the namespace."""
if namespace_path is None:
+27
View File
@@ -25,6 +25,7 @@ from ..common import DATA
from ..db import DBConnection, LOOP
from ..functions import FunctionVersion, UdfDefinition
from ..job import AsyncJob, Job
from ..materialized_view import MaterializedView, SelectArg
if TYPE_CHECKING:
from .._lancedb import JobDescription, JobInfo
@@ -648,6 +649,32 @@ class RemoteDBConnection(DBConnection):
namespace_path=namespace_path,
)
@override
def create_materialized_view(
self,
name: str,
source: str,
*,
select: SelectArg = None,
where: Optional[str] = None,
limit: Optional[int] = None,
) -> MaterializedView:
raise NotImplementedError(
"materialized views are supported only on local databases"
)
@override
def open_materialized_view(self, name: str) -> MaterializedView:
raise NotImplementedError(
"materialized views are supported only on local databases"
)
@override
def list_materialized_views(self) -> List[str]:
raise NotImplementedError(
"materialized views are supported only on local databases"
)
@override
def drop_table(self, name: str, namespace_path: Optional[List[str]] = None):
"""Drop a table from the database.
+2 -2
View File
@@ -49,7 +49,7 @@ from lancedb.index import (
LabelList,
)
from lancedb.job import Job
from lancedb.functions import FunctionApplication
from lancedb.functions import FunctionApplication, RefreshColumnResult
from lancedb.remote.db import LOOP
from lancedb.table import IndexConfigType, KNOWN_METRICS
import pyarrow as pa
@@ -972,7 +972,7 @@ class RemoteTable(Table):
def refresh_column(self, column: str):
return LOOP.run(self._table.refresh_column(column))
def refresh_column_async(self, column: str) -> Job:
def refresh_column_async(self, column: str) -> Job[RefreshColumnResult]:
return Job(LOOP.run(self._table.refresh_column_async(column)))
def alter_columns(
+484 -70
View File
@@ -24,11 +24,16 @@ import os
import random
import threading
import time
import warnings
from collections import deque
from concurrent.futures import ThreadPoolExecutor
from copy import deepcopy
from multiprocessing import RawArray
from typing import Any, Callable, Iterator, Optional, Union
from typing import Any, Callable, cast, Iterator, Literal, Optional, Union
import pyarrow as pa
import pyarrow.compute as pc
import torch
from torch.utils.data import IterableDataset, get_worker_info
from .permutation import (
@@ -61,7 +66,7 @@ class StreamingDataset(IterableDataset):
Internally ``__iter__`` runs a two-stage pipeline:
- **Stage 1 (I/O)**: one thread pool with ``num_splits * prefetch_batches``
- **Stage 1 (I/O)**: one thread pool with ``num_splits * io_queue_depth``
workers fetches raw ``RecordBatch`` objects from LanceDB in parallel
across all splits and places them in a per-split raw-batch queue.
- **Stage 2 (transform)**: a second thread pool with
@@ -104,11 +109,11 @@ class StreamingDataset(IterableDataset):
call. Larger values amortise per-request overhead (critical on object
storage) at the cost of higher memory usage per split buffer. Defaults
to ``DEFAULT_READ_BATCH_SIZE`` (64).
prefetch_batches:
io_queue_depth:
Number of I/O batches to keep in flight per split. Higher values
overlap storage latency with transform and training compute at the cost
of more memory and threads. Defaults to ``DEFAULT_PREFETCH_BATCHES``
(4).
of more memory and threads. Must be greater than zero. Defaults to
``DEFAULT_PREFETCH_BATCHES`` (4).
columns:
Optional list of column names to read. When set, only those columns
are fetched from storage; all others are omitted. ``None`` (the
@@ -132,6 +137,39 @@ class StreamingDataset(IterableDataset):
Maximum number of transforms to run concurrently. Must be greater
than zero. When ``None`` (the default), uses ``os.cpu_count()`` or 1
when the CPU count is unavailable.
pack_sequences:
Sequence-packing mode: token lists from consecutive documents are
joined with ``eos_id`` and sliced into blocks of this many tokens.
Each item is then a dict of two ``(pack_sequences,)`` LongTensors —
``input_ids`` and ``doc_ids`` (per-position document index within
the block, for block-diagonal masks or position-id resets).
* Packing happens independently per owned split and preserves per-split
resume state.
* When a split cannot fill a real block for a cycle but an owned sibling
still can, or when only a short tail remains at epoch end, the short buffer
is padded to ``pack_sequences`` with ``pad_id`` so every local cycle emits
one block per owned split.
* ``eos_id``, ``pad_id``, and ``columns`` naming a single integer-list
column are required; incompatible with ``transform``.
eos_id:
Separator token id between packed documents. Required with
``pack_sequences``, ignored otherwise.
pad_id:
Padding token id used to complete blocks when a split runs out of
real tokens mid-cycle or at epoch end. Required with
``pack_sequences``, ignored otherwise. It must be reserved for padding:
padding positions retain the preceding document's ``doc_id`` (or zero
in an all-padding block), so callers must mask them separately using
``input_ids == pad_id``.
blocks_per_epoch:
Total number of packed blocks emitted globally per epoch. Required with
``pack_sequences``. An integer must be divisible by ``num_splits``.
Every logical split emits exactly ``blocks_per_epoch / num_splits``
blocks: exhausted splits emit padding, while tokens beyond the budget
are left out of the epoch. This fixed per-split budget keeps packed
iteration and checkpoints independent of rank topology.
Pass ``"auto"`` to estimate a corpus-level budget from a bounded sample
of token lists. The estimate may be inaccurate.
on_transform_error:
What to do when the transform raises an exception:
@@ -175,6 +213,16 @@ class StreamingDataset(IterableDataset):
Prefer the ``filter`` parameter when bad rows can be expressed as a
SQL predicate (e.g. ``"col IS NOT NULL"``) — filtering happens before
splits are built, so every guarantee is fully preserved.
transform_queue_depth:
Number of transform-result batches to buffer per split in the
post-transform queue before backpressure is applied to the transform
stage. When the combined count of in-flight transform futures and
already-buffered rows for a split reaches
``transform_queue_depth * read_batch_size``, no new transforms are
submitted for that split until the consumer catches up. Useful for
capping peak memory when the consumer (e.g. a GPU training step) is
slower than the transform stage. Must be greater than zero.
``None`` (the default) imposes no limit.
worker_info_override:
If set, used in place of ``torch.utils.data.get_worker_info()`` to
determine the DataLoader worker assignment. Intended for unit tests
@@ -194,17 +242,30 @@ class StreamingDataset(IterableDataset):
rank: int = 0,
world_size: int = 1,
read_batch_size: int = DEFAULT_READ_BATCH_SIZE,
prefetch_batches: int = DEFAULT_PREFETCH_BATCHES,
io_queue_depth: int = DEFAULT_PREFETCH_BATCHES,
columns: Optional[list[str]] = None,
shuffle_clump_size: Optional[int] = None,
filter: Optional[str] = None,
transform: Optional[Callable] = None,
transform_parallelism: Optional[int] = None,
pack_sequences: Optional[int] = None,
eos_id: Optional[int] = None,
pad_id: Optional[int] = None,
blocks_per_epoch: Optional[Union[int, Literal["auto"]]] = None,
on_transform_error: Union[str, Callable[[Exception], bool]] = "raise",
transform_queue_depth: Optional[int] = None,
connection_factory: Optional[Callable[[str], Any]] = None,
worker_info_override=None,
# Deprecated; use io_queue_depth instead.
prefetch_batches: Optional[int] = None,
):
super().__init__()
if prefetch_batches is not None:
logger.warning(
"prefetch_batches is deprecated and will be removed in a future "
"version; use io_queue_depth instead"
)
io_queue_depth = prefetch_batches
if num_splits is None:
num_splits = world_size
if shuffle_seed is None:
@@ -214,8 +275,59 @@ class StreamingDataset(IterableDataset):
f"num_splits ({num_splits}) must be divisible by "
f"world_size ({world_size})"
)
if io_queue_depth <= 0:
raise ValueError("io_queue_depth must be greater than 0")
if transform_parallelism is not None and transform_parallelism <= 0:
raise ValueError("transform_parallelism must be greater than 0")
if pack_sequences is not None:
if pack_sequences <= 0:
raise ValueError("pack_sequences must be greater than 0")
if eos_id is None:
raise ValueError("eos_id is required when pack_sequences is set")
if pad_id is None:
raise ValueError("pad_id is required when pack_sequences is set")
if blocks_per_epoch is None:
raise ValueError(
"blocks_per_epoch is required when pack_sequences is set"
)
if blocks_per_epoch != "auto":
if not isinstance(blocks_per_epoch, int) or isinstance(
blocks_per_epoch, bool
):
raise ValueError(
"blocks_per_epoch must be a positive integer or 'auto'"
)
if blocks_per_epoch <= 0:
raise ValueError("blocks_per_epoch must be greater than 0")
if blocks_per_epoch % num_splits != 0:
raise ValueError(
f"blocks_per_epoch ({blocks_per_epoch}) must be divisible by "
f"num_splits ({num_splits})"
)
if transform is not None:
raise ValueError("transform cannot be combined with pack_sequences")
if columns is None or len(columns) != 1:
raise ValueError(
"pack_sequences requires columns to name exactly one "
"list-typed column of token ids"
)
field = table.schema.field(columns[0])
if not (
pa.types.is_list(field.type)
or pa.types.is_large_list(field.type)
or pa.types.is_fixed_size_list(field.type)
):
raise ValueError(
f"pack_sequences requires a list-typed token column; "
f"{columns[0]} has type {field.type}"
)
if not pa.types.is_integer(field.type.value_type):
raise ValueError(
"pack_sequences requires a token column with integer values; "
f"{columns[0]} has value type {field.type.value_type}"
)
elif blocks_per_epoch is not None:
raise ValueError("blocks_per_epoch requires pack_sequences")
if on_transform_error not in ("raise", "skip", "warn") and not callable(
on_transform_error
):
@@ -223,6 +335,8 @@ class StreamingDataset(IterableDataset):
"on_transform_error must be 'raise', 'skip', 'warn', or a "
f"callable, got {on_transform_error!r}"
)
if transform_queue_depth is not None and transform_queue_depth <= 0:
raise ValueError("transform_queue_depth must be greater than 0")
self._table = table
self._num_splits = num_splits
@@ -232,16 +346,26 @@ class StreamingDataset(IterableDataset):
self._rank = rank
self._world_size = world_size
self._read_batch_size = read_batch_size
self._prefetch_batches = prefetch_batches
self._io_queue_depth = io_queue_depth
self._columns = columns
self._shuffle_clump_size = shuffle_clump_size
self._filter = filter
self._transform = transform
self._transform_parallelism = transform_parallelism
self._pack_sequences = pack_sequences
self._eos_id = eos_id
self._pad_id = pad_id
self._blocks_per_epoch = blocks_per_epoch
self._on_transform_error = on_transform_error
self._transform_queue_depth = transform_queue_depth
self._connection_factory = connection_factory
self._worker_info_override = worker_info_override
# Packing resume state: permutation positions and partial-block buffers.
self._pack_consumed: list[int] = [0] * num_splits
self._pack_buffers: dict[int, dict[str, list[int]]] = {}
self._pack_blocks_emitted: list[int] = [0] * num_splits
# Live references to pipeline state, set only while __iter__ is running
# in the same process. Used by the observability properties when the
# DataLoader runs with num_workers=0.
@@ -291,6 +415,9 @@ class StreamingDataset(IterableDataset):
else:
self._perm_table = builder.split_sequential(fixed=num_splits).execute()
if self._blocks_per_epoch == "auto":
self._blocks_per_epoch = self._estimate_blocks_per_epoch()
# Contiguous block of global split indices assigned to this rank.
splits_per_rank = num_splits // world_size
rank_start = rank * splits_per_rank
@@ -298,6 +425,71 @@ class StreamingDataset(IterableDataset):
range(rank_start, rank_start + splits_per_rank)
)
def _estimate_blocks_per_epoch(self) -> int:
"""Estimate a fixed packed-block budget from a bounded token sample."""
# TODO: Replace this fallback with Lance's dedicated exact token-count
# estimation API once it is available.
if self._pack_sequences is None or not self._columns:
raise RuntimeError(
"packing must be configured before estimating its budget"
)
pack_len = self._pack_sequences
token_column = self._columns[0]
sample_cap_per_split = max(1, 100_000 // self._num_splits)
sampled_tokens = 0
total_sampled = 0
total_rows = 0
rng = random.Random(self._shuffle_seed)
warnings.warn(
"blocks_per_epoch='auto' uses an approximate token-count sample; "
"pass an explicit value for exact epoch sizing",
)
for split in range(self._num_splits):
permutation = Permutation.from_tables(
self._table, self._perm_table, split=split
)
permutation = permutation.select_columns([token_column])
permutation = permutation.with_transform(Transforms.arrow2arrow)
split_rows = permutation.num_rows
if split_rows == 0:
raise ValueError(
"blocks_per_epoch='auto' cannot estimate an empty dataset"
)
# Sample roughly 1% from each logical split, with at least one row
# per split and a global target cap of 100,000 rows.
sample_rows = min(
split_rows,
max(1, min((split_rows + 99) // 100, sample_cap_per_split)),
)
sample_offsets = sorted(rng.sample(range(split_rows), sample_rows))
sample_batch_size = max(1, self._read_batch_size)
for start in range(0, sample_rows, sample_batch_size):
batch = permutation.__getitems__(
sample_offsets[start : start + sample_batch_size]
)
lengths = pc.list_value_length(batch.column(0))
if lengths.null_count:
raise ValueError("pack_sequences does not support null token lists")
sampled_tokens += int(pc.sum(lengths).as_py())
total_sampled += sample_rows
total_rows += split_rows
# Pool the samples into one global average. Each document contributes
# one EOS token.
estimated_tokens = (
(sampled_tokens + total_sampled) * total_rows // total_sampled
)
blocks = estimated_tokens // pack_len
return max(
self._num_splits,
blocks - blocks % self._num_splits,
)
def _resolve_my_splits(self) -> list[int]:
"""Return the split indices this instance should read in __iter__."""
torch_worker_info = get_worker_info()
@@ -348,8 +540,14 @@ class StreamingDataset(IterableDataset):
)
if self._columns is not None:
perm = perm.select_columns(self._columns)
perm = perm.with_transform(lambda batch: batch)
start_pos = self._resume_positions.get(split_idx, self._resume_offset)
perm = perm.with_transform(Transforms.arrow2arrow)
# 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)
)
if start_pos > 0:
perm = perm.with_skip(start_pos)
initial_positions.append(start_pos)
@@ -365,14 +563,36 @@ class StreamingDataset(IterableDataset):
pos_consumed = list(initial_positions)
batch_size = self._read_batch_size
max_prefetch = self._prefetch_batches
io_queue_depth = self._io_queue_depth
transform_workers = (
self._transform_parallelism
if self._transform_parallelism is not None
else (os.cpu_count() or 1)
)
final_transform = (
self._transform if self._transform is not None else Transforms.arrow2python
final_transform: Callable[[pa.RecordBatch], Any]
if self._pack_sequences is not None:
# Packing consumes raw token lists, one per document.
def arrow_tokens(batch: pa.RecordBatch) -> list[list[int]]:
token_column = batch.column(0)
if token_column.null_count or token_column.flatten().null_count:
raise ValueError(
"pack_sequences does not support null token lists or values"
)
return cast(list[list[int]], token_column.to_pylist())
final_transform = arrow_tokens
else:
final_transform = (
self._transform
if self._transform is not None
else Transforms.arrow2python
)
# None means no limit; otherwise cap rows per split to
# transform_queue_depth batches worth (including in-flight transforms).
max_cooked_rows = (
self._transform_queue_depth * batch_size
if self._transform_queue_depth is not None
else None
)
# Per-split pipeline state. Batches are paired with the absolute
@@ -409,7 +629,9 @@ class StreamingDataset(IterableDataset):
io_pending[i].append((abs_start, io_pool.submit(_io_call, perm_i, indices)))
def _fill_io(i: int) -> None:
while len(io_pending[i]) < max_prefetch and fetch_head[i] < split_sizes[i]:
while (
len(io_pending[i]) < io_queue_depth and fetch_head[i] < split_sizes[i]
):
_submit_io(i)
def _drain_io(i: int) -> None:
@@ -487,7 +709,19 @@ class StreamingDataset(IterableDataset):
def _try_submit_tx(i: int) -> None:
"""Submit transforms for raw_batches[i] up to available capacity."""
while raw_batches[i] and tx_semaphore.acquire(blocking=False):
while raw_batches[i]:
# Backpressure: only submit a new transform when there is room
# for a full batch in the post-transform queue. Checking for
# a full batch prevents submitting a transform that would
# overflow the limit mid-batch (e.g. 990 rows queued with a
# capacity of 1000 and a batch_size of 128 must wait until
# 128 rows have been consumed, not just 1).
if max_cooked_rows is not None:
in_pipeline = len(cooked[i]) + len(tx_pending[i]) * batch_size
if in_pipeline + batch_size > max_cooked_rows:
break
if not tx_semaphore.acquire(blocking=False):
break
abs_start, batch = raw_batches[i].popleft()
tx_pending[i].append(tx_pool.submit(_tx_call_guarded, abs_start, batch))
@@ -529,19 +763,127 @@ class StreamingDataset(IterableDataset):
else:
break # split exhausted
def _update_stats(*, idle: bool = False) -> None:
"""Refresh pipeline statistics visible to the parent process."""
ws = self._worker_stats
ws[0] = sum(split_sizes[j] - fetch_head[j] for j in range(n))
ws[1] = (
0
if idle
else sum(batch.num_rows for q in raw_batches for _, batch in q)
)
ws[2] = 0 if idle else sum(len(q) for q in cooked)
ws[3] = sum(local_consumed)
ws[4] = self._bytes_loaded
ws[5] = int(self._fetch_time * 1_000_000)
ws[6] = int(self._transform_time * 1_000_000)
ws[7] = self._rows_skipped
# Sequence-packing helpers
pack_len = cast(int, self._pack_sequences)
eos_id = cast(int, self._eos_id)
pad_id = cast(int, self._pad_id)
blocks_per_split = (
cast(int, self._blocks_per_epoch) // self._num_splits
if self._pack_sequences is not None
else 0
)
pack_consumed = list(self._pack_consumed)
pack_buffers = deepcopy(self._pack_buffers)
pack_blocks_emitted = list(self._pack_blocks_emitted)
def _pack_buffer(i: int) -> dict[str, list[int]]:
return pack_buffers.setdefault(my_splits[i], {"tokens": [], "starts": []})
def _fill_block(i: int) -> None:
"""Fill split i's buffer to one block or exhaust the split."""
buf = _pack_buffer(i)
while len(buf["tokens"]) < pack_len:
_ensure_cooked(i)
if not cooked[i]:
return
buf["starts"].append(len(buf["tokens"]))
pos, tokens = cooked[i].popleft()
buf["tokens"].extend(tokens)
buf["tokens"].append(eos_id)
pack_consumed[my_splits[i]] = pos + 1
local_consumed[i] += 1
_advance(i)
def _emit_block(i: int) -> dict[str, Any]:
buf = _pack_buffer(i)
tokens, starts = buf["tokens"], buf["starts"]
# doc_ids label document segments within the block; 0 also covers
# the continuation of a document begun in a prior block.
doc_ids = torch.zeros(pack_len, dtype=torch.int64)
doc_starts = [s for s in starts if 0 < s < pack_len]
doc_ids[doc_starts] = 1
doc_ids.cumsum_(dim=0) # cumulative sum marks document boundaries
block = {
"input_ids": torch.tensor(tokens[:pack_len], dtype=torch.int64),
"doc_ids": doc_ids,
}
del tokens[:pack_len]
# Shift start boundaries for the next call.
buf["starts"] = [s - pack_len for s in starts if s >= pack_len]
return block
def _commit_pack_state() -> None:
self._pack_consumed = list(pack_consumed)
self._pack_buffers = {
split: {
"tokens": list(buffer["tokens"]),
"starts": list(buffer["starts"]),
}
for split, buffer in pack_buffers.items()
}
self._pack_blocks_emitted = list(pack_blocks_emitted)
# ── Main loop ─────────────────────────────────────────────────────────
with ThreadPoolExecutor(max_workers=n * max_prefetch) as io_pool:
with ThreadPoolExecutor(max_workers=n * io_queue_depth) as io_pool:
with ThreadPoolExecutor(max_workers=transform_workers) as tx_pool:
self._raw_batches_ref = raw_batches
self._cooked_ref = cooked
self._fetch_head_ref = fetch_head
self._split_sizes_ref = split_sizes
self._local_consumed_ref = local_consumed
try:
for i in range(n):
_fill_io(i)
if self._pack_sequences is not None:
first_count = pack_blocks_emitted[my_splits[0]]
if any(
pack_blocks_emitted[split] != first_count
for split in my_splits[1:]
):
raise ValueError(
"Packed checkpoint is not aligned across the splits "
"owned by this iterator; merge every rank "
"state with merge_state_dicts before resuming on a "
"different topology"
)
while pack_blocks_emitted[my_splits[0]] < blocks_per_split:
# Each logical split gets one block per cycle. Exhausted
# splits are padded through the fixed global budget.
for i in range(n):
_fill_block(i)
for i in range(n):
tokens = _pack_buffer(i)["tokens"]
if len(tokens) < pack_len:
tokens.extend([pad_id] * (pack_len - len(tokens)))
block = _emit_block(i)
pack_blocks_emitted[my_splits[i]] += 1
if i == n - 1:
_commit_pack_state()
_update_stats()
yield block
return
while True:
# A cycle only runs if every split can still produce a
# row. Without skips all splits exhaust simultaneously
@@ -575,21 +917,7 @@ class StreamingDataset(IterableDataset):
self._resume_offset = initial_offset + local_consumed[i]
for j, split_idx in enumerate(my_splits):
self._resume_positions[split_idx] = pos_consumed[j]
ws = self._worker_stats
ws[0] = sum(
split_sizes[j] - fetch_head[j] for j in range(n)
)
ws[1] = sum(
batch.num_rows
for q in raw_batches
for _, batch in q
)
ws[2] = sum(len(q) for q in cooked)
ws[3] = sum(local_consumed)
ws[4] = self._bytes_loaded
ws[5] = int(self._fetch_time * 1_000_000)
ws[6] = int(self._transform_time * 1_000_000)
ws[7] = self._rows_skipped
_update_stats()
yield row
finally:
@@ -597,15 +925,7 @@ class StreamingDataset(IterableDataset):
# when iteration ends mid-cycle (e.g. a split whose rows
# were all skipped before completing a single cycle), so
# counters like rows_skipped would otherwise be stale.
ws = self._worker_stats
ws[0] = sum(split_sizes[j] - fetch_head[j] for j in range(n))
ws[1] = 0 # queue-depth properties document 0 when idle
ws[2] = 0
ws[3] = sum(local_consumed)
ws[4] = self._bytes_loaded
ws[5] = int(self._fetch_time * 1_000_000)
ws[6] = int(self._transform_time * 1_000_000)
ws[7] = self._rows_skipped
_update_stats(idle=True)
self._raw_batches_ref = None
self._cooked_ref = None
self._fetch_head_ref = None
@@ -763,21 +1083,31 @@ class StreamingDataset(IterableDataset):
def state_dict(self) -> dict:
"""Snapshot the dataset's consumption state.
The returned dict is topology-independent: at global step boundaries
every split has been consumed the same number of times (by the
round-robin design), so the per-split count is a single uniform value
that is identical across all ranks and DataLoader workers.
``positions_consumed_per_split`` records how far into each split's
permutation iteration has advanced. It only differs from
``samples_consumed_per_split`` when ``on_transform_error`` skipped
rows, in which case entries are exact for the splits this instance
iterated and a lower bound (the sample count) for splits owned by
other ranks or workers. Combine the state dicts from all ranks with
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
every rank with
[merge_state_dicts][lancedb.streaming.StreamingDataset.merge_state_dicts]
to recover the exact value for every split before resuming on a
different topology.
before resuming on a different topology.
Packed state includes partial token buffers and emitted block counts
for every logical split. When packing is sharded, merge every rank
state with ``merge_state_dicts`` before loading it.
"""
if self._pack_sequences is not None:
return {
"shuffle_seed": self._shuffle_seed,
"num_splits": self._num_splits,
"epoch": self._epoch,
"pack_sequences": self._pack_sequences,
"eos_id": self._eos_id,
"pad_id": self._pad_id,
"blocks_per_epoch": self._blocks_per_epoch,
"samples_consumed_per_split": list(self._pack_consumed),
"blocks_emitted_per_split": list(self._pack_blocks_emitted),
"pack_buffers": deepcopy(self._pack_buffers),
}
positions = [
self._resume_positions.get(split, self._resume_offset)
for split in range(self._num_splits)
@@ -795,7 +1125,9 @@ class StreamingDataset(IterableDataset):
Raises ``ValueError`` if ``num_splits`` or ``shuffle_seed`` differ
from the checkpoint, since a different split structure or shuffle order
makes mid-epoch resumption meaningless.
makes mid-epoch resumption meaningless. Packed checkpoints
pin ``pack_sequences``, ``eos_id``, ``pad_id``,
``blocks_per_epoch``, and ``epoch``.
"""
if state["num_splits"] != self._num_splits:
raise ValueError(
@@ -807,6 +1139,31 @@ class StreamingDataset(IterableDataset):
f"shuffle_seed mismatch: checkpoint has {state['shuffle_seed']}, "
f"current dataset has {self._shuffle_seed}"
)
if "pack_buffers" in state or self._pack_sequences is not None:
for key in (
"pack_sequences",
"eos_id",
"pad_id",
"blocks_per_epoch",
"epoch",
):
ours = getattr(self, f"_{key}")
if state.get(key) != ours:
raise ValueError(
f"{key} mismatch: checkpoint has {state.get(key)}, "
f"current dataset has {ours}"
)
self._pack_consumed = [int(c) for c in state["samples_consumed_per_split"]]
self._pack_blocks_emitted = [
int(c) for c in state["blocks_emitted_per_split"]
]
self._pack_buffers = {
int(g): {"tokens": list(b["tokens"]), "starts": list(b["starts"])}
for g, b in state["pack_buffers"].items()
}
return
consumed = state["samples_consumed_per_split"]
# All entries are equal at step boundaries; use the first.
if isinstance(consumed, list):
@@ -828,25 +1185,22 @@ class StreamingDataset(IterableDataset):
def merge_state_dicts(states: list[dict]) -> dict:
"""Merge state dicts saved by different ranks into one exact state.
Only needed when ``on_transform_error`` skips rows in multi-rank
training: each rank then knows the exact permutation position only for
its own splits, and records a lower bound for the rest. Because
exactly one rank owns each split, the elementwise maximum across all
ranks' ``positions_consumed_per_split`` recovers the exact position of
every split. Without skipped rows every rank's state is already
identical and merging is a no-op.
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
states must cover every rank at the same global step.
Raises ``ValueError`` if the states are empty or were not produced by
the same run (mismatched seed, split count, epoch, or sample counts).
Raises ``ValueError`` if the states are empty, were not produced by
the same run, or do not represent the same global step.
The merge is always all-to-all and topology-agnostic: collect the
``state_dict()`` from every rank of the *previous* run into one list,
merge that whole list, and hand the identical merged result to every
rank of the *next* run — regardless of whether the rank count grew,
shrank, or stayed the same. There is no pairwise or subset merging
step, because each split's exact position is only known to whichever
rank owned that split, and the elementwise maximum needs every rank's
contribution to be correct.
``state_dict()`` from every rank of the *previous* run into
one list, merge that whole list, and hand the identical merged result
to every rank of the *next* run — regardless of whether the
topology grew, shrank, or stayed the same. There is no pairwise or
subset merging step, because each split's exact state is only known to
whichever iterator owned that split.
For example, checkpointing 8 ranks and resuming on 4 (the same
pattern applies when growing, e.g. 4 ranks resuming on 8)::
@@ -879,13 +1233,73 @@ class StreamingDataset(IterableDataset):
if not states:
raise ValueError("merge_state_dicts requires at least one state dict")
first = states[0]
packed = "pack_buffers" in first
config_keys = ["shuffle_seed", "num_splits", "epoch"]
if packed:
config_keys.extend(
["pack_sequences", "eos_id", "pad_id", "blocks_per_epoch"]
)
for state in states[1:]:
for key in ("shuffle_seed", "num_splits", "epoch"):
if ("pack_buffers" in state) != packed:
raise ValueError("cannot merge packed and unpacked state dicts")
for key in config_keys:
if state[key] != first[key]:
raise ValueError(
f"{key} mismatch across state dicts: "
f"{state[key]} != {first[key]}"
)
if packed:
num_splits = first["num_splits"]
for state in states:
for key in (
"samples_consumed_per_split",
"blocks_emitted_per_split",
):
if len(state[key]) != num_splits:
raise ValueError(
f"{key} must contain one entry per logical split"
)
merged_consumed = []
merged_emitted = []
merged_buffers = {}
for split in range(num_splits):
owner = states[0]
owner_progress = (
owner["blocks_emitted_per_split"][split],
owner["samples_consumed_per_split"][split],
)
for state in states[1:]:
progress = (
state["blocks_emitted_per_split"][split],
state["samples_consumed_per_split"][split],
)
if progress > owner_progress:
owner = state
owner_progress = progress
merged_consumed.append(owner["samples_consumed_per_split"][split])
merged_emitted.append(owner["blocks_emitted_per_split"][split])
buffer = owner["pack_buffers"].get(
split, owner["pack_buffers"].get(str(split))
)
if buffer is not None:
merged_buffers[split] = deepcopy(buffer)
if len(set(merged_emitted)) > 1:
raise ValueError(
"packed state dicts were not captured at the same global "
"step or do not cover every rank"
)
merged = dict(first)
merged["samples_consumed_per_split"] = merged_consumed
merged["blocks_emitted_per_split"] = merged_emitted
merged["pack_buffers"] = merged_buffers
return merged
for state in states[1:]:
if (
state["samples_consumed_per_split"]
!= first["samples_consumed_per_split"]
+43 -18
View File
@@ -40,7 +40,7 @@ from ._blob import (
from .types import BlobMode
from lancedb.arrow import peek_reader
from lancedb.background_loop import LOOP, embedding_executor
from lancedb.job import AsyncJob, Job
from lancedb.job import AsyncJob, Job, _typed_job
from .dependencies import (
_check_for_hugging_face,
_check_for_lance,
@@ -72,7 +72,10 @@ from .index import (
FTS,
)
from .expr import Expr
from .functions import FunctionApplication
from .functions import (
FunctionApplication,
RefreshColumnResult as RefreshColumnJobResult,
)
from .merge import LanceMergeInsertBuilder
from .pydantic import LanceModel, model_to_dict
from .query import (
@@ -2039,7 +2042,7 @@ class Table(ABC):
"""
@abstractmethod
def refresh_column_async(self, column: str) -> Job:
def refresh_column_async(self, column: str) -> Job[RefreshColumnJobResult]:
"""
Like :meth:`refresh_column`, but returns a handle to the refresh job
instead of blocking until it completes.
@@ -2050,6 +2053,12 @@ class Table(ABC):
than failing the job. On local tables the job runs in-process; on
LanceDB Cloud and Enterprise it is the server's backfill job.
Returns
-------
Job[RefreshColumnResult]
A job whose successful ``wait`` returns row counts plus the source
and published table versions.
Examples
--------
>>> import lancedb
@@ -2058,7 +2067,9 @@ class Table(ABC):
>>> table.add_columns(computed={"doubled": "x * 2"})
AddColumnsResult(version=2)
>>> job = table.refresh_column_async("doubled")
>>> job.wait()
>>> result = job.wait()
>>> result.rows_assigned
2
>>> job.status()
'finished'
"""
@@ -4082,7 +4093,7 @@ class LanceTable(Table):
[`AsyncTable.refresh_column`][lancedb.AsyncTable.refresh_column]."""
return LOOP.run(self._table.refresh_column(column))
def refresh_column_async(self, column: str) -> Job:
def refresh_column_async(self, column: str) -> Job[RefreshColumnJobResult]:
"""Fill a computed column's unfilled rows, returning a handle to the
refresh job. See
[`Table.refresh_column_async`][lancedb.table.Table.refresh_column_async].
@@ -6122,7 +6133,9 @@ class AsyncTable:
"""
return await self._inner.refresh_column(column)
async def refresh_column_async(self, column: str) -> AsyncJob:
async def refresh_column_async(
self, column: str
) -> AsyncJob[RefreshColumnJobResult]:
"""
Like :meth:`refresh_column`, but returns a handle to the refresh job
instead of blocking until it completes.
@@ -6134,6 +6147,12 @@ class AsyncTable:
in-process; on LanceDB Cloud and Enterprise it is the server's
backfill job.
Returns
-------
AsyncJob[RefreshColumnResult]
A job whose successful ``wait`` returns row counts plus the source
and published table versions.
Examples
--------
>>> import asyncio
@@ -6143,12 +6162,16 @@ class AsyncTable:
... table = await db.create_table("computed_job_async_demo", [{"x": 1}])
... await table.add_columns(computed={"doubled": "x * 2"})
... job = await table.refresh_column_async("doubled")
... await job.wait()
... result = await job.wait()
... assert result.rows_assigned == 1
... return await job.status()
>>> asyncio.run(refresh_in_background())
'finished'
"""
return AsyncJob(await self._inner.refresh_column_async(column))
return _typed_job(
await self._inner.refresh_column_async(column),
RefreshColumnJobResult.from_json,
)
async def alter_columns(
self, *alterations: Iterable[dict[str, Any]]
@@ -6804,21 +6827,21 @@ class Branches:
"""Diff a branch against main."""
return LOOP.run(self._table.branches.diff(from_branch))
def merge(self, from_branch: str, dry_run: bool = False) -> Dict[str, Any]:
"""Merge a branch into main, or dry-run.
def cherry_pick(self, from_branch: str, dry_run: bool = False) -> Dict[str, Any]:
"""Cherry-pick a branch onto main, or dry-run.
Parameters
----------
from_branch: str
Branch to merge from.
Branch to cherry-pick from.
dry_run: bool, default False
When True, only preview. When False, attempt the merge.
When True, only preview. When False, attempt the cherry-pick.
Notes
-----
A rejected merge returns ``status="rejected"`` instead of raising.
A failed cherry-pick returns ``status="failed"`` instead of raising.
"""
return LOOP.run(self._table.branches.merge(from_branch, dry_run))
return LOOP.run(self._table.branches.cherry_pick(from_branch, dry_run))
def _wrap(
self, async_table: "AsyncTable", version: Optional[int] = None
@@ -6954,9 +6977,11 @@ class AsyncBranches:
"""Diff a branch against main."""
return await self._table.branches.diff(from_branch)
async def merge(self, from_branch: str, dry_run: bool = False) -> Dict[str, Any]:
"""Merge a branch into main, or dry-run.
async def cherry_pick(
self, from_branch: str, dry_run: bool = False
) -> Dict[str, Any]:
"""Cherry-pick a branch onto main, or dry-run.
A rejected merge returns ``status="rejected"`` instead of raising.
A failed cherry-pick returns ``status="failed"`` instead of raising.
"""
return await self._table.branches.merge(from_branch, dry_run)
return await self._table.branches.cherry_pick(from_branch, dry_run)
+2 -2
View File
@@ -774,7 +774,7 @@ def test_drop_table_async(tmp_db: lancedb.DBConnection):
job = tmp_db.drop_table_async("test")
assert job.id is None
assert job.status() == "finished"
job.wait()
assert job.wait() is None
assert tmp_db.table_names() == []
tmp_db.create_table("test", data=data)
@@ -790,7 +790,7 @@ async def test_drop_table_async_connection(tmp_db_async: lancedb.AsyncConnection
job = await tmp_db_async.drop_table_async("test")
assert job.id is None
assert await job.status() == "finished"
await job.wait()
assert await job.wait() is None
assert await tmp_db_async.table_names() == []
+392 -2
View File
@@ -1374,6 +1374,188 @@ def test_transform_parallelism_must_be_positive(lance_table, transform_paralleli
)
# ---------------------------------------------------------------------------
# Backpressure / transform_queue_depth tests
# ---------------------------------------------------------------------------
@pytest.mark.parametrize("transform_queue_depth", [0, -1])
def test_transform_queue_depth_must_be_positive(lance_table, transform_queue_depth):
"""transform_queue_depth=0 or negative must raise ValueError."""
with pytest.raises(
ValueError, match="transform_queue_depth must be greater than 0"
):
StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
transform_queue_depth=transform_queue_depth,
)
@pytest.mark.parametrize("transform_queue_depth", [1, 2, 4])
def test_transform_queue_depth_correctness(lance_table, transform_queue_depth):
"""With backpressure enabled, every row is still yielded exactly once."""
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle_seed=SHUFFLE_SEED,
transform_queue_depth=transform_queue_depth,
read_batch_size=8,
)
items = list(ds)
assert sorted(item["id"] for item in items) == list(range(NUM_ROWS))
def test_transform_queue_depth_matches_no_backpressure(lance_table):
"""With backpressure enabled the same samples are produced as without it."""
ds_unlimited = StreamingDataset(
lance_table, num_splits=NUM_SPLITS, shuffle_seed=SHUFFLE_SEED
)
ds_limited = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle_seed=SHUFFLE_SEED,
transform_queue_depth=1,
)
assert [item["id"] for item in ds_unlimited] == [
item["id"] for item in ds_limited
], "transform_queue_depth must not affect the sample ordering or set"
def test_transform_queue_depth_bounds_cooked_rows(lance_table):
"""prefetch_queue_depth stays within transform_queue_depth * read_batch_size
per split when observed from the main thread during iteration."""
n_splits = 4
batch_size = 8
cooked_depth = 2
# max cooked rows across all 4 splits: 4 * 2 * 8 = 64
max_allowed = n_splits * cooked_depth * batch_size
ds = StreamingDataset(
lance_table,
num_splits=n_splits,
shuffle_seed=SHUFFLE_SEED,
transform_queue_depth=cooked_depth,
read_batch_size=batch_size,
transform_parallelism=1,
world_size=1,
)
peak = 0
for _ in ds:
depth = ds.prefetch_queue_depth
if depth > peak:
peak = depth
# The main thread observes depth *after* popping a row, so the peak is at
# most max_allowed (one row already popped from the split just served).
assert peak <= max_allowed, (
f"prefetch_queue_depth peaked at {peak}, expected <= {max_allowed}"
)
def test_transform_queue_depth_does_not_admit_at_capacity_minus_one(tmp_path):
"""Admission requires a full read_batch_size of free space, not just one slot.
The test intercepts ThreadPoolExecutor.submit to make I/O calls execute
synchronously on the main thread. This ensures all raw batches land in
raw_batches (via _drain_io) before _try_submit_tx evaluates the admission
predicate for the first time. Without this, the I/O future for batch N+1
might still be in io_pending at the capacity-minus-one transition, leaving
raw_batches empty and causing _try_submit_tx to skip the admission check
entirely so both the correct and the broken predicate produce depth=0
observations and the test cannot distinguish them.
With all raw batches pre-loaded in raw_batches the 43 cooked transition
(consuming one row from a full cooked queue) always triggers _try_submit_tx
against a non-empty raw_batches.
With transform_queue_depth=1 and batch_size=4, max_cooked_rows=4.
A transform may only be submitted when in_pipeline + batch_size <= 4, i.e.
when in_pipeline == 0 (cooked is completely empty). Under the old broken
predicate (in_pipeline >= max_cooked_rows) the second transform would be
admitted with cooked containing batch_size-1 rows still unconsumed.
"""
import concurrent.futures as cf
from concurrent.futures import ThreadPoolExecutor
from unittest.mock import patch
db = lancedb.connect(tmp_path)
batch_size = 4
# Four full batches → four transform submissions to observe.
table = db.create_table("t", pa.table({"id": list(range(batch_size * 4))}))
cooked_at_submit: list[int] = []
original_submit = ThreadPoolExecutor.submit
def tracking_submit(self, fn, *args, **kwargs):
name = getattr(fn, "__name__", "")
if name == "_io_call":
# Run I/O synchronously on the calling (main) thread and return an
# already-completed Future. _drain_io checks fut.done(), so a
# completed Future is moved to raw_batches immediately on the next
# _advance call — making raw-batch readiness deterministic at the
# capacity-minus-one transition instead of depending on I/O thread
# scheduling.
fut = cf.Future()
try:
fut.set_result(fn(*args, **kwargs))
except Exception as exc:
fut.set_exception(exc)
return fut
if name == "_tx_call_guarded":
# Capture cooked depth synchronously on the main thread before the
# transform worker can drain the queue.
ref = ds._cooked_ref
cooked_at_submit.append(len(ref[0]) if ref is not None else -1)
return original_submit(self, fn, *args, **kwargs)
with patch.object(ThreadPoolExecutor, "submit", tracking_submit):
ds = StreamingDataset(
table,
num_splits=1,
shuffle_seed=42,
read_batch_size=batch_size,
transform_queue_depth=1,
transform_parallelism=1,
)
list(ds)
assert len(cooked_at_submit) == 4, (
f"Expected 4 transform submissions (one per batch), got {len(cooked_at_submit)}"
)
# With full-batch backpressure each transform is only admitted when the
# cooked queue is completely empty (depth == 0). The old broken predicate
# would admit at depth == batch_size - 1 == 3.
assert all(depth == 0 for depth in cooked_at_submit), (
"Transform admitted with non-empty cooked queue; full-batch backpressure "
"requires in_pipeline + batch_size <= max_cooked_rows before admission. "
f"Cooked depths at each submission: {cooked_at_submit}"
)
# ---------------------------------------------------------------------------
# Deprecated parameter name tests
# ---------------------------------------------------------------------------
def test_prefetch_batches_deprecated_warns(lance_table, caplog):
"""prefetch_batches logs a deprecation warning and behaves like io_queue_depth."""
with caplog.at_level(logging.WARNING, logger="lancedb.streaming"):
ds = StreamingDataset(
lance_table,
num_splits=NUM_SPLITS,
shuffle_seed=SHUFFLE_SEED,
prefetch_batches=2,
)
messages = [r.message for r in caplog.records if r.levelno >= logging.WARNING]
assert any("deprecated" in m.lower() and "io_queue_depth" in m for m in messages), (
f"Expected deprecation warning mentioning io_queue_depth; got: {messages}"
)
assert sorted(item["id"] for item in ds) == list(range(NUM_ROWS))
def test_filter_limits_rows(tmp_path):
"""A filter expression is applied to the permutation so only matching rows
are yielded. IDs 0..59 pass ``id < 60``; the other 60 are excluded."""
@@ -1931,6 +2113,214 @@ def test_shuffle_seed_none_generates_stable_seed(lance_table):
assert first == second, "Same resolved seed must produce the same ordering"
# Sequence packing tests
def _create_token_table(tmp_path, documents):
db = lancedb.connect(tmp_path)
tokens = pa.array(documents, type=pa.list_(pa.int64()))
return db.create_table("tokens", pa.table({"tokens": tokens}))
def _packed_dataset(table, pack_sequences, *, blocks_per_epoch, pad_id=0, **kwargs):
return StreamingDataset(
table,
shuffle=False,
columns=["tokens"],
pack_sequences=pack_sequences,
eos_id=9,
pad_id=pad_id,
blocks_per_epoch=blocks_per_epoch,
**kwargs,
)
def test_pack_sequences_emits_blocks_and_pads_final_tail(tmp_path):
table = _create_token_table(tmp_path, [[1, 2], [3, 4], [5]])
dataset = _packed_dataset(table, 6, blocks_per_epoch=2)
blocks = list(dataset)
assert len(blocks) == 2
assert blocks[0]["input_ids"].tolist() == [1, 2, 9, 3, 4, 9]
assert blocks[0]["doc_ids"].tolist() == [0, 0, 0, 1, 1, 1]
assert blocks[1]["input_ids"].tolist() == [5, 9, 0, 0, 0, 0]
assert blocks[1]["doc_ids"].tolist() == [0, 0, 0, 0, 0, 0]
assert blocks[0]["input_ids"].dtype == torch.int64
assert blocks[0]["doc_ids"].dtype == torch.int64
def test_pack_sequences_pads_lagging_splits(tmp_path):
table = _create_token_table(
tmp_path,
[[1], [2], [10, 11, 12, 13, 14, 15, 16, 17], [20]],
)
dataset = _packed_dataset(table, 5, blocks_per_epoch=6, num_splits=2)
input_ids = [block["input_ids"].tolist() for block in dataset]
# Split 0 has four real tokens including EOS markers, while split 1 has
# eleven. Packing must emit three complete two-split cycles.
assert input_ids == [
[1, 9, 2, 9, 0],
[10, 11, 12, 13, 14],
[0, 0, 0, 0, 0],
[15, 16, 17, 9, 20],
[0, 0, 0, 0, 0],
[9, 0, 0, 0, 0],
]
per_rank = []
for rank in range(2):
rank_dataset = _packed_dataset(
table,
5,
blocks_per_epoch=6,
num_splits=2,
world_size=2,
rank=rank,
)
per_rank.append([block["input_ids"].tolist() for block in rank_dataset])
assert [len(blocks) for blocks in per_rank] == [3, 3]
sharded = [block for cycle in zip(*per_rank) for block in cycle]
assert sharded == input_ids
def test_pack_sequences_auto_estimates_filtered_token_column(tmp_path):
db = lancedb.connect(tmp_path)
table = db.create_table(
"tokens",
pa.table(
{
"tokens": pa.array([[1] * 4, [2] * 9], type=pa.list_(pa.int64())),
"keep": [True, False],
}
),
)
table.add(
pa.table(
{
"tokens": pa.array([[3] * 4, [4] * 9], type=pa.list_(pa.int64())),
"keep": [True, False],
}
)
)
with pytest.warns(UserWarning, match="approximate token-count sample"):
dataset = _packed_dataset(
table,
5,
blocks_per_epoch="auto",
num_splits=2,
filter="keep",
)
# Two kept documents contain 8 tokens plus 2 EOS tokens: two blocks.
assert dataset.state_dict()["blocks_per_epoch"] == 2
def test_pack_sequences_checkpoint_resumes_on_new_topology(tmp_path):
table = _create_token_table(
tmp_path,
[[1], [2], [10, 11, 12, 13, 14, 15, 16, 17], [20]],
)
kwargs = dict(pack_sequences=5, blocks_per_epoch=6, num_splits=2)
reference = list(_packed_dataset(table, **kwargs))
datasets = [
_packed_dataset(table, world_size=2, rank=rank, **kwargs) for rank in range(2)
]
iterators = [iter(dataset) for dataset in datasets]
first_cycle = [next(iterator) for iterator in iterators]
checkpoint = StreamingDataset.merge_state_dicts(
[dataset.state_dict() for dataset in datasets]
)
for iterator in iterators:
iterator.close()
resumed = _packed_dataset(table, **kwargs)
resumed.load_state_dict(checkpoint)
actual_remaining = list(resumed)
assert [block["input_ids"].tolist() for block in first_cycle] == [
[1, 9, 2, 9, 0],
[10, 11, 12, 13, 14],
]
assert checkpoint["blocks_emitted_per_split"] == [1, 1]
assert [block["input_ids"].tolist() for block in actual_remaining] == [
block["input_ids"].tolist() for block in reference[2:]
]
assert [block["doc_ids"].tolist() for block in actual_remaining] == [
block["doc_ids"].tolist() for block in reference[2:]
]
def test_pack_sequences_validates_configuration_and_tokens(tmp_path):
table = _create_token_table(tmp_path, [[1, 2]])
with pytest.raises(ValueError, match="pad_id is required"):
StreamingDataset(
table,
shuffle=False,
columns=["tokens"],
pack_sequences=4,
eos_id=9,
)
with pytest.raises(ValueError, match="blocks_per_epoch is required"):
StreamingDataset(
table,
shuffle=False,
columns=["tokens"],
pack_sequences=4,
eos_id=9,
pad_id=0,
)
with pytest.raises(ValueError, match="must be divisible"):
_packed_dataset(table, 4, blocks_per_epoch=3, num_splits=2)
with pytest.raises(ValueError, match="positive integer or 'auto'"):
_packed_dataset(table, 4, blocks_per_epoch="estimate")
checkpoint = _packed_dataset(table, 4, blocks_per_epoch=1).state_dict()
resumed = _packed_dataset(table, 4, blocks_per_epoch=1, pad_id=8)
with pytest.raises(ValueError, match="pad_id mismatch"):
resumed.load_state_dict(checkpoint)
float_db = lancedb.connect(tmp_path / "float")
float_table = float_db.create_table(
"tokens",
pa.table({"tokens": pa.array([[1.5, 2.5]], type=pa.list_(pa.float64()))}),
)
with pytest.raises(ValueError, match="token column with integer values"):
_packed_dataset(float_table, 4, blocks_per_epoch=1)
null_db = lancedb.connect(tmp_path / "null")
null_table = null_db.create_table(
"tokens",
pa.table({"tokens": pa.array([None], type=pa.list_(pa.int64()))}),
)
with pytest.raises(ValueError, match="does not support null token lists"):
list(_packed_dataset(null_table, 4, blocks_per_epoch=1))
null_value_db = lancedb.connect(tmp_path / "null_value")
null_value_table = null_value_db.create_table(
"tokens",
pa.table(
{"tokens": pa.array([[1], [2, None], [3]], type=pa.list_(pa.int64()))}
),
)
blocks = list(
_packed_dataset(
null_value_table,
2,
blocks_per_epoch=2,
on_transform_error="skip",
)
)
assert [block["input_ids"].tolist() for block in blocks] == [[1, 9], [3, 9]]
# ---------------------------------------------------------------------------
# Doc examples — each test mirrors the code snippet in index.mdx so that
# broken doc examples are caught before they ship.
@@ -1954,7 +2344,7 @@ def test_doc_example_basic(tmp_path):
def test_doc_example_prefetch_params(tmp_path):
"""doc: Prefetching — read_batch_size and prefetch_batches still cover all rows."""
"""doc: Prefetching — read_batch_size and io_queue_depth still cover all rows."""
db = lancedb.connect(tmp_path)
table = db.create_table("t", pa.table({"id": list(range(NUM_ROWS))}))
@@ -1963,7 +2353,7 @@ def test_doc_example_prefetch_params(tmp_path):
num_splits=NUM_SPLITS,
shuffle_seed=SHUFFLE_SEED,
read_batch_size=8,
prefetch_batches=2,
io_queue_depth=2,
)
assert sorted(s["id"] for s in ds) == list(range(NUM_ROWS))
@@ -187,7 +187,7 @@ def _mock_remote_function_catalog():
"job_state": "DONE",
"result": state["version"],
}
elif self.path == "/v1/functions/get":
elif self.path == "/v1/functions/describe":
assert body == {
"name": "normalize_score",
"version": "fv_exact",
+1 -1
View File
@@ -88,7 +88,7 @@ async def binary_table(db_async):
async def test_create_index_async_returns_done_job(some_table: AsyncTable):
job = await some_table.create_index_async("id", config=BTree())
assert job.id is None
await job.wait()
assert await job.wait() is None
assert len(await some_table.list_indices()) == 1
await job.cancel()
@@ -0,0 +1,268 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import lancedb
import pytest
from lancedb.materialized_view import MaterializedViewDefinition
STABLE_ROW_IDS = {"new_table_enable_stable_row_ids": "true"}
def make_db(tmp_path):
db = lancedb.connect(tmp_path, storage_options=STABLE_ROW_IDS)
db.create_table(
"people",
[
{"name": "ada", "age": 36},
{"name": "kid", "age": 7},
{"name": "grace", "age": 85},
],
)
return db
def test_create_refresh_and_query(tmp_path):
db = make_db(tmp_path)
view = db.create_materialized_view(
"adults",
"people",
select=["name", ("shout", "upper(name)")],
where="age >= 18",
)
assert view.name == "adults"
assert view.table.count_rows() == 0
result = view.refresh()
assert result.mode == "rebuild"
assert result.rows_written == 2
rows = view.table.search().to_list()
assert sorted(row["shout"] for row in rows) == ["ADA", "GRACE"]
def test_definition_round_trips(tmp_path):
db = make_db(tmp_path)
db.create_materialized_view("adults", "people", where="age >= 18")
view = db.open_materialized_view("adults")
assert view.definition == MaterializedViewDefinition(
source_table="people",
projections=[("name", "`name`"), ("age", "`age`")],
filter="age >= 18",
inputs=["age", "name"],
)
def test_incremental_refresh_after_append(tmp_path):
db = make_db(tmp_path)
view = db.create_materialized_view("copy", "people")
view.refresh()
db.open_table("people").add([{"name": "alan", "age": 41}])
result = view.refresh()
assert result.mode == "incremental"
assert result.rows_written == 1
assert view.table.count_rows() == 4
assert view.refresh().mode == "no_op"
def test_incremental_refresh_after_update(tmp_path):
db = make_db(tmp_path)
view = db.create_materialized_view("copy", "people")
view.refresh()
db.open_table("people").update(where="name = 'kid'", values={"age": 8})
result = view.refresh()
assert result.mode == "incremental"
assert result.rows_written == 1
rows = view.table.search().to_list()
assert sorted(row["age"] for row in rows) == [8, 36, 85]
def test_legacy_storage_source_update_rebuilds(tmp_path):
db = lancedb.connect(
tmp_path,
storage_options={**STABLE_ROW_IDS, "new_table_data_storage_version": "legacy"},
)
db.create_table("people", [{"name": "ada", "age": 36}, {"name": "kid", "age": 7}])
view = db.create_materialized_view("copy", "people")
view.refresh()
db.open_table("people").update(where="name = 'kid'", values={"age": 8})
result = view.refresh()
assert result.mode == "rebuild"
rows = view.table.search().to_list()
assert sorted(row["age"] for row in rows) == [8, 36]
def test_list_and_not_a_view(tmp_path):
db = make_db(tmp_path)
db.create_materialized_view("adults", "people", where="age >= 18")
assert db.list_materialized_views() == ["adults"]
with pytest.raises(ValueError, match="not a materialized view"):
db.open_materialized_view("people")
def test_invalid_expression_fails_at_create(tmp_path):
db = make_db(tmp_path)
with pytest.raises(Exception, match="missing"):
db.create_materialized_view("bad", "people", select=[("x", "missing + 1")])
assert "bad" not in db.list_tables().tables
@pytest.mark.asyncio
async def test_async_create_refresh_and_open(tmp_path):
db = await lancedb.connect_async(tmp_path, storage_options=STABLE_ROW_IDS)
await db.create_table("people", [{"name": "ada", "age": 36}])
view = await db.create_materialized_view(
"shouts", "people", select=[("shout", "upper(name)")]
)
result = await view.refresh()
assert result.mode == "rebuild"
assert result.rows_written == 1
reopened = await db.open_materialized_view("shouts")
definition = await reopened.definition()
assert definition.projections == [("shout", "upper(name)")]
assert await db.list_materialized_views() == ["shouts"]
@pytest.mark.asyncio
async def test_async_incremental(tmp_path):
db = await lancedb.connect_async(tmp_path, storage_options=STABLE_ROW_IDS)
await db.create_table("people", [{"name": "ada", "age": 36}])
view = await db.create_materialized_view("copy", "people")
await view.refresh()
table = await db.open_table("people")
await table.add([{"name": "alan", "age": 41}])
result = await view.refresh()
assert result.mode == "incremental"
assert result.rows_written == 1
def test_source_requires_stable_row_ids(tmp_path):
db = lancedb.connect(tmp_path)
db.create_table("plain", [{"x": 1}])
with pytest.raises(Exception, match="stable row ids"):
db.create_materialized_view("v", "plain")
def test_bare_select_names_are_quoted(tmp_path):
db = lancedb.connect(tmp_path, storage_options=STABLE_ROW_IDS)
db.create_table("odd_names", [{"order item": "widget", "select": 2}])
view = db.create_materialized_view(
"quoted", "odd_names", select=["order item", "select"]
)
result = view.refresh()
assert result.rows_written == 1
rows = view.table.search().to_list()
assert rows[0]["order item"] == "widget"
assert rows[0]["select"] == 2
@pytest.mark.asyncio
async def test_async_remote_is_refused_without_network():
db = await lancedb.connect_async(
"db://nowhere", api_key="sk_test", region="us-east-1"
)
with pytest.raises(NotImplementedError, match="local"):
await db.create_materialized_view("v", "src")
with pytest.raises(NotImplementedError, match="local"):
await db.open_materialized_view("v")
with pytest.raises(NotImplementedError, match="local"):
await db.list_materialized_views()
def test_scalar_select_is_one_column(tmp_path):
db = make_db(tmp_path)
view = db.create_materialized_view("just_name", "people", select="name")
view.refresh()
rows = view.table.search().to_list()
assert set(rows[0]) - {"__source_row_id"} == {"name"}
assert sorted(row["name"] for row in rows) == ["ada", "grace", "kid"]
@pytest.mark.asyncio
async def test_async_scalar_select_is_one_column(tmp_path):
db = await lancedb.connect_async(tmp_path, storage_options=STABLE_ROW_IDS)
await db.create_table("people", [{"name": "ada", "age": 36}])
view = await db.create_materialized_view("just_name", "people", select="name")
await view.refresh()
rows = await view.table.query().to_list()
assert set(rows[0]) - {"__source_row_id"} == {"name"}
def test_limit_above_i64_max_is_refused(tmp_path):
db = make_db(tmp_path)
with pytest.raises(ValueError, match="exceeds the maximum"):
db.create_materialized_view("too_big", "people", limit=2**63)
# The boundary is fine, and zero still means an empty view.
db.create_materialized_view("at_max", "people", limit=2**63 - 1)
empty = db.create_materialized_view("none", "people", limit=0)
empty.refresh()
assert empty.table.count_rows() == 0
def _namespace_db(tmp_path):
return lancedb.connect_namespace(
"dir",
{"root": str(tmp_path)},
storage_options=STABLE_ROW_IDS,
)
def test_namespace_connection_materialized_views(tmp_path):
db = _namespace_db(tmp_path)
db.create_table(
"people",
[{"name": "ada", "age": 36}, {"name": "kid", "age": 7}],
storage_options=STABLE_ROW_IDS,
)
view = db.create_materialized_view("adults", "people", where="age >= 18")
view.refresh()
assert view.table.count_rows() == 1
assert db.list_materialized_views() == ["adults"]
reopened = db.open_materialized_view("adults")
assert reopened.definition.source_table == "people"
with pytest.raises(ValueError, match="not a materialized view"):
db.open_materialized_view("people")
@pytest.mark.asyncio
async def test_async_namespace_connection_materialized_views(tmp_path):
db = lancedb.connect_namespace_async(
"dir",
{"root": str(tmp_path)},
storage_options=STABLE_ROW_IDS,
)
await db.create_table(
"people",
[{"name": "ada", "age": 36}, {"name": "kid", "age": 7}],
storage_options=STABLE_ROW_IDS,
)
view = await db.create_materialized_view("adults", "people", where="age >= 18")
await view.refresh()
assert await view.table.count_rows() == 1
assert await db.list_materialized_views() == ["adults"]
reopened = await db.open_materialized_view("adults")
assert (await reopened.definition()).source_table == "people"
# The view's table came through the namespace, not straight from the
# inner connection: a bare inner table carries no namespace context, so
# its pushdown routing differs from a table the namespace opened.
through_namespace = await db.open_table("adults")
for handle in (view.table, reopened.table):
assert (
handle._route_pushdown_to_rust == through_namespace._route_pushdown_to_rust
)
assert handle._namespace_path == through_namespace._namespace_path
+83 -10
View File
@@ -242,8 +242,8 @@ def test_remote_table_branches_sync():
table.branches.delete("exp")
def test_remote_table_branch_merge_defaults_to_execute():
merge_bodies = []
def test_remote_table_cherry_pick_defaults_to_execute():
cherry_pick_bodies = []
diff = {
"fromBranch": "exp",
"parentVersion": 1,
@@ -265,8 +265,7 @@ def test_remote_table_branch_merge_defaults_to_execute():
"changedColumns": [],
"addedIndexes": [],
"removedIndexes": [],
"mergeable": True,
"mergeBlockers": [],
"errors": [],
}
def handler(request):
@@ -276,11 +275,11 @@ def test_remote_table_branch_merge_defaults_to_execute():
else:
content_len = int(request.headers.get("Content-Length"))
request_body = json.loads(request.rfile.read(content_len))
merge_bodies.append(request_body)
cherry_pick_bodies.append(request_body)
dry_run = request_body["dry_run"]
status = 200 if dry_run else 409
body = {
"status": "ready" if dry_run else "rejected",
"status": "ready" if dry_run else "failed",
"diff": diff,
"preview": {"promotedColumns": []},
}
@@ -292,10 +291,10 @@ def test_remote_table_branch_merge_defaults_to_execute():
with mock_lancedb_connection(handler) as db:
branches = db.open_table("test").branches
assert branches.merge("exp")["status"] == "rejected"
assert branches.merge("exp", dry_run=True)["status"] == "ready"
assert branches.cherry_pick("exp")["status"] == "failed"
assert branches.cherry_pick("exp", dry_run=True)["status"] == "ready"
assert merge_bodies == [
assert cherry_pick_bodies == [
{"from_branch": "exp", "dry_run": False},
{"from_branch": "exp", "dry_run": True},
]
@@ -901,11 +900,85 @@ def test_remote_create_index_async_returns_job():
table = db.create_table("test", [{"id": 1}])
job = table.create_index_async("id", config=BTree())
assert job.id == "job-1"
job.wait(timeout=timedelta(seconds=30))
assert job.wait(timeout=timedelta(seconds=30)) is None
assert len(describe_calls) == 2
job.cancel()
def test_remote_refresh_async_returns_typed_terminal_result():
terminal_result = {
"rows_assigned": 12,
"rows_failed": 0,
"rows_remaining": 0,
"source_version": 7,
"published_version": 8,
}
def handler(request):
content_len = int(request.headers.get("Content-Length", 0))
body = request.rfile.read(content_len) if content_len > 0 else b""
if request.path == "/v1/table/test/backfill_column":
assert json.loads(body)["column"] == "derived"
request.send_response(202)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b'{"job_id": "refresh-1"}')
elif request.path == "/v1/jobs/describe":
assert json.loads(body)["job_id"] == "refresh-1"
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": "refresh-1",
"job_type": "function_refresh",
"job_state": "DONE",
"result": terminal_result,
}
).encode()
)
elif request.path == "/v1/table/test/create/?mode=create":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(b"{}")
elif request.path == "/v1/table/test/describe/":
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"version": 1,
"schema": {
"fields": [
{
"name": "id",
"type": {"type": "int64"},
"nullable": False,
}
]
},
}
).encode()
)
else:
request.send_response(404)
request.end_headers()
with mock_lancedb_connection(handler) as db:
table = db.create_table("test", [{"id": 1}])
job = table.refresh_column_async("derived")
assert job.id == "refresh-1"
result = job.wait(timeout=timedelta(seconds=30))
assert isinstance(result, lancedb.RefreshColumnResult)
assert result.model_dump() == terminal_result
assert result.rows_filled == 12
assert result.version == 8
def test_remote_job_wait_raises_on_failure():
from lancedb.exceptions import JobFailedError
from lancedb.index import BTree
+18 -3
View File
@@ -1467,7 +1467,7 @@ def test_create_index_async_returns_done_job(mem_db: DBConnection):
table = mem_db.create_table("job_test", [{"id": i} for i in range(10)])
job = table.create_index_async("id", config=BTree())
assert job.id is None
job.wait()
assert job.wait() is None
assert len(table.list_indices()) == 1
job.cancel()
@@ -3947,10 +3947,21 @@ def test_refresh_column_async_returns_job(tmp_path):
job = table.refresh_column_async("doubled")
assert job.id is None # in-process jobs have no server id
assert job.wait() is None
result = job.wait()
assert isinstance(result, lancedb.RefreshColumnResult)
assert result.rows_assigned == 2
assert result.rows_failed == 0
assert result.rows_remaining == 0
assert result.source_version == 2
assert result.published_version == 3
assert job.status() == "finished"
assert sorted(table.to_arrow()["doubled"].to_pylist()) == [2, 4]
no_op = table.refresh_column_async("doubled").wait()
assert no_op.rows_assigned == 0
assert no_op.source_version == 3
assert no_op.published_version is None
# Bad input raises at the call, not through the job.
with pytest.raises(Exception, match="not a computed column"):
table.refresh_column_async("x")
@@ -3963,6 +3974,10 @@ async def test_refresh_column_async_job_async_table(tmp_path):
await table.add_columns(computed={"tripled": "x * 3"})
job = await table.refresh_column_async("tripled")
assert await job.wait() is None
result = await job.wait()
assert isinstance(result, lancedb.RefreshColumnResult)
assert result.rows_assigned == 1
assert result.source_version == 2
assert result.published_version == 3
assert await job.status() == "finished"
assert (await table.to_arrow())["tripled"].to_pylist() == [9]
+35 -1
View File
@@ -333,6 +333,40 @@ impl Connection {
})
}
#[pyo3(signature = (name, source, projections=None, filter=None, limit=None))]
pub fn create_materialized_view(
self_: PyRef<'_, Self>,
name: String,
source: String,
projections: Option<Vec<(String, String)>>,
filter: Option<String>,
limit: Option<u64>,
) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.get_inner()?.clone();
future_into_py(self_.py(), async move {
let mut builder = inner.create_materialized_view(name, source);
if let Some(projections) = projections {
builder = builder.select(projections);
}
if let Some(filter) = filter {
builder = builder.only_if(filter);
}
if let Some(limit) = limit {
builder = builder.limit(limit);
}
let view = builder.execute().await.infer_error()?;
Ok(Table::new(view.table().clone()))
})
}
pub fn list_materialized_views(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.get_inner()?.clone();
future_into_py(self_.py(), async move {
let views = inner.list_materialized_views().await.infer_error()?;
Ok(views.into_iter().map(|view| view.name).collect::<Vec<_>>())
})
}
#[pyo3(signature = (name, namespace_path=None))]
pub fn drop_table(
self_: PyRef<'_, Self>,
@@ -575,7 +609,7 @@ impl Connection {
.create_function_async(request)
.await
.infer_error()
.map(crate::job::FunctionJob::new)
.map(crate::job::Job::new_typed)
})
}
+18 -55
View File
@@ -5,72 +5,33 @@ use std::sync::Arc;
use crate::runtime::future_into_py;
use pyo3::{Bound, PyAny, PyRef, PyResult, pyclass, pymethods};
use serde::Serialize;
use crate::error::PythonErrorExt;
#[pyclass]
pub struct Job {
inner: Arc<lancedb::Job>,
}
/// Python bridge for a typed remote Function registration job.
///
/// The public Python layer decodes the canonical JSON returned by `wait`
/// into its immutable `FunctionVersion` model.
#[pyclass]
pub struct FunctionJob {
inner: Arc<lancedb::Job<lancedb::function::FunctionVersion>>,
}
impl FunctionJob {
pub(crate) fn new(inner: lancedb::Job<lancedb::function::FunctionVersion>) -> Self {
Self {
inner: Arc::new(inner),
}
}
inner: Arc<lancedb::Job<std::result::Result<Option<String>, String>>>,
}
impl Job {
pub(crate) fn new(inner: lancedb::Job) -> Self {
Self {
inner: Arc::new(inner),
inner: Arc::new(inner.map(|()| Ok(None))),
}
}
}
#[pymethods]
impl FunctionJob {
#[getter]
pub fn id(&self) -> Option<String> {
self.inner.id().map(str::to_string)
}
pub fn status(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner.clone();
future_into_py(
self_.py(),
async move { inner.status().await.infer_error() },
)
}
pub fn wait(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner.clone();
future_into_py(self_.py(), async move {
inner
.wait()
.await
.infer_error()?
.to_canonical_json()
.infer_error()
})
}
pub fn cancel(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner.clone();
future_into_py(self_.py(), async move {
inner.cancel().await.infer_error()?;
Ok(())
})
pub(crate) fn new_typed<T>(inner: lancedb::Job<T>) -> Self
where
T: Clone + Serialize + Send + Sync + 'static,
{
Self {
inner: Arc::new(inner.map(|result| {
serde_json::to_string(&result)
.map(Some)
.map_err(|error| format!("failed to serialize typed job result: {error}"))
})),
}
}
}
@@ -92,8 +53,10 @@ impl Job {
pub fn wait(self_: PyRef<'_, Self>) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner.clone();
future_into_py(self_.py(), async move {
inner.wait().await.infer_error()?;
Ok(None::<()>)
let result = inner.wait().await.infer_error()?;
result
.map_err(|message| lancedb::Error::Runtime { message })
.infer_error()
})
}
+3 -3
View File
@@ -16,8 +16,8 @@ use query::{FTSQuery, HybridQuery, Query, VectorQuery};
use session::Session;
use table::{
AddColumnsResult, AddResult, AlterColumnsResult, DeleteResult, DropColumnsResult, FtsToken,
LsmWriteSpec, MergeResult, PyBlobFile, RefreshColumnResult, Table, UpdateFieldMetadataResult,
UpdateResult,
LsmWriteSpec, MergeResult, PyBlobFile, RefreshColumnResult, RefreshMaterializedViewResult,
Table, UpdateFieldMetadataResult, UpdateResult,
};
pub mod arrow;
@@ -47,7 +47,6 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<Session>()?;
m.add_class::<Table>()?;
m.add_class::<crate::job::Job>()?;
m.add_class::<crate::job::FunctionJob>()?;
m.add_class::<crate::job::JobInfo>()?;
m.add_class::<crate::job::JobDescription>()?;
m.add_class::<crate::job::JobFailureInfo>()?;
@@ -60,6 +59,7 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<RecordBatchStream>()?;
m.add_class::<AddColumnsResult>()?;
m.add_class::<RefreshColumnResult>()?;
m.add_class::<RefreshMaterializedViewResult>()?;
m.add_class::<AlterColumnsResult>()?;
m.add_class::<UpdateFieldMetadataResult>()?;
m.add_class::<AddResult>()?;
+58 -3
View File
@@ -441,6 +441,41 @@ impl From<lancedb::table::RefreshColumnResult> for RefreshColumnResult {
}
}
#[pyclass(get_all, from_py_object)]
#[derive(Clone, Debug)]
pub struct RefreshMaterializedViewResult {
pub mode: String,
pub rows_written: u64,
pub source_version: u64,
pub version: u64,
}
#[pymethods]
impl RefreshMaterializedViewResult {
pub fn __repr__(&self) -> String {
format!(
"RefreshMaterializedViewResult(mode={}, rows_written={}, source_version={}, version={})",
self.mode, self.rows_written, self.source_version, self.version
)
}
}
impl From<lancedb::RefreshMaterializedViewResult> for RefreshMaterializedViewResult {
fn from(result: lancedb::RefreshMaterializedViewResult) -> Self {
let mode = match result.mode {
lancedb::RefreshMode::Rebuild => "rebuild",
lancedb::RefreshMode::Incremental => "incremental",
lancedb::RefreshMode::NoOp => "no_op",
};
Self {
mode: mode.to_string(),
rows_written: result.rows_written,
source_version: result.source_version,
version: result.version,
}
}
}
#[pymethods]
impl AddColumnsResult {
pub fn __repr__(&self) -> String {
@@ -1584,7 +1619,27 @@ impl Table {
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
let job = inner.refresh_column_async(column).await.infer_error()?;
Ok(crate::job::Job::new(job))
Ok(crate::job::Job::new_typed(job))
})
}
#[pyo3(signature = (full=false, source_version=None))]
pub fn refresh_materialized_view(
self_: PyRef<'_, Self>,
full: bool,
source_version: Option<u64>,
) -> PyResult<Bound<'_, PyAny>> {
let inner = self_.inner_ref()?.clone();
future_into_py(self_.py(), async move {
let view = lancedb::MaterializedView::from_table(inner)
.await
.infer_error()?;
let mut builder = view.refresh().full(full);
if let Some(version) = source_version {
builder = builder.source_version(version);
}
let result = builder.execute().await.infer_error()?;
Ok(RefreshMaterializedViewResult::from(result))
})
}
@@ -1885,7 +1940,7 @@ impl Branches {
}
#[pyo3(signature = (from_branch, dry_run=false))]
pub fn merge(
pub fn cherry_pick(
self_: PyRef<'_, Self>,
from_branch: String,
dry_run: bool,
@@ -1893,7 +1948,7 @@ impl Branches {
let inner = self_.inner.clone();
future_into_py(self_.py(), async move {
let result = inner
.merge_branch(&from_branch, dry_run)
.cherry_pick(&from_branch, dry_run)
.await
.infer_error()?;
Python::attach(|py| struct_to_wire_py(py, &result))