mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-30 09:58:20 +00:00
Merge remote-tracking branch 'origin/main' into gatekeeper/fix-2820-1
This commit is contained in:
+1
-1
@@ -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"
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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))
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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() == []
|
||||
|
||||
|
||||
|
||||
@@ -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 4→3 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",
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user