mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-01 19:18:38 +00:00
Merge remote-tracking branch 'origin/main' into gatekeeper/fix-2057-1
# Conflicts: # rust/lancedb/src/remote/table.rs
This commit is contained in:
+5
-6
@@ -1,6 +1,6 @@
|
||||
[package]
|
||||
name = "lancedb-python"
|
||||
version = "0.38.0-beta.3"
|
||||
version = "0.38.0-beta.10"
|
||||
publish = false
|
||||
edition.workspace = true
|
||||
description = "Python bindings for LanceDB"
|
||||
@@ -26,7 +26,9 @@ lance-namespace-impls.workspace = true
|
||||
lance-io.workspace = true
|
||||
env_logger.workspace = true
|
||||
log.workspace = true
|
||||
pyo3 = { version = "0.28", features = ["extension-module", "abi3-py310", "chrono"] }
|
||||
# Maturin enables extension-module mode for Python builds. Keeping it out of
|
||||
# Cargo features lets Rust unit tests link against libpython.
|
||||
pyo3 = { version = "0.28", features = ["abi3-py310", "chrono"] }
|
||||
chrono.workspace = true
|
||||
pyo3-async-runtimes = { version = "0.28", features = [
|
||||
"attributes",
|
||||
@@ -41,10 +43,7 @@ tokio.workspace = true
|
||||
libc = "0.2"
|
||||
|
||||
[build-dependencies]
|
||||
pyo3-build-config = { version = "0.28", features = [
|
||||
"extension-module",
|
||||
"abi3-py310",
|
||||
] }
|
||||
pyo3-build-config = { version = "0.28", features = ["abi3-py310"] }
|
||||
|
||||
[features]
|
||||
default = ["remote", "lancedb/aws", "lancedb/gcs", "lancedb/azure", "lancedb/dynamodb", "lancedb/oss", "lancedb/huggingface", "lancedb/cos", "lancedb/goosefs", "lancedb/metrics-otel"]
|
||||
|
||||
@@ -38,6 +38,25 @@ Stable releases are created about every 2 weeks. For the latest features and bug
|
||||
pip install --pre --extra-index-url https://pypi.fury.io/lancedb/ lancedb
|
||||
```
|
||||
|
||||
### Threading in CPU-limited containers
|
||||
|
||||
LanceDB uses separate pools for compute work and storage I/O. On a container with
|
||||
two visible CPUs, current releases intentionally use one compute worker by default;
|
||||
no manual configuration is needed. If every query logs an I/O core reservation
|
||||
warning on a two-CPU container, upgrade from LanceDB 0.21.1 or earlier.
|
||||
|
||||
The two commonly tuned environment variables control different resources:
|
||||
|
||||
- `LANCE_CPU_THREADS` overrides the number of compute workers. One worker is the
|
||||
appropriate setting for a two-CPU container when an explicit override is needed.
|
||||
- `LANCE_IO_THREADS` controls concurrent storage operations, not reserved CPU
|
||||
cores. Its default can be greater than the number of CPUs because I/O workers
|
||||
spend much of their time waiting for storage.
|
||||
|
||||
Keep the defaults unless measurements show that the workload benefits from an
|
||||
override. See the [Lance threading model](https://lance.org/guide/performance/#threading-model)
|
||||
for the current defaults and tuning guidance.
|
||||
|
||||
## Usage
|
||||
|
||||
### Basic Example
|
||||
|
||||
@@ -101,9 +101,12 @@ azure = ["adlfs>=2024.2.0"]
|
||||
[tool.maturin]
|
||||
python-source = "python"
|
||||
module-name = "lancedb._lancedb"
|
||||
# uv installs the project as an editable package before `uv run`, so keep that
|
||||
# bootstrap build consistent with `maturin develop`.
|
||||
editable-profile = "dev"
|
||||
|
||||
[build-system]
|
||||
requires = ["maturin>=1.4"]
|
||||
requires = ["maturin>=1.10"]
|
||||
build-backend = "maturin"
|
||||
|
||||
[tool.ruff.lint]
|
||||
|
||||
@@ -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
|
||||
@@ -173,6 +179,18 @@ def connect(
|
||||
... },
|
||||
... )
|
||||
|
||||
For Azure Blob Storage, credentials can be passed directly without setting
|
||||
environment variables:
|
||||
|
||||
>>> azure_storage_options = {
|
||||
... "account_name": "some-account",
|
||||
... "account_key": "some-key",
|
||||
... }
|
||||
>>> db = lancedb.connect( # doctest: +SKIP
|
||||
... "az://my-container/my-database",
|
||||
... storage_options=azure_storage_options,
|
||||
... )
|
||||
|
||||
For tests and temporary data, use an in-memory database:
|
||||
|
||||
>>> db = lancedb.connect("memory://")
|
||||
@@ -459,6 +477,10 @@ async def connect_async(
|
||||
--------
|
||||
|
||||
>>> import lancedb
|
||||
>>> azure_storage_options = {
|
||||
... "account_name": "some-account",
|
||||
... "account_key": "some-key",
|
||||
... }
|
||||
>>> async def doctest_example():
|
||||
... # For a local directory, provide a path to the database
|
||||
... db = await lancedb.connect_async("~/.lancedb")
|
||||
@@ -466,6 +488,11 @@ async def connect_async(
|
||||
... db = await lancedb.connect_async("s3://my-bucket/lancedb",
|
||||
... storage_options={
|
||||
... "aws_access_key_id": "***"})
|
||||
... # Azure credentials can also be passed directly
|
||||
... db = await lancedb.connect_async(
|
||||
... "az://my-container/my-database",
|
||||
... storage_options=azure_storage_options,
|
||||
... )
|
||||
... # For tests and temporary data, use an in-memory database
|
||||
... db = await lancedb.connect_async("memory://")
|
||||
... # Connect to LanceDB cloud
|
||||
@@ -506,6 +533,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:
|
||||
@@ -281,6 +283,7 @@ class Table:
|
||||
mode: Literal["append", "overwrite"],
|
||||
progress: Optional[Any] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
allow_external_blob_outside_bases: bool = False,
|
||||
) -> AddResult: ...
|
||||
async def update(
|
||||
self, updates: Dict[str, str], where: Optional[str]
|
||||
@@ -355,6 +358,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 +426,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 +710,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."""
|
||||
|
||||
@@ -85,8 +85,9 @@ class Expr:
|
||||
# for dict keys / set membership.
|
||||
__hash__ = None # type: ignore[assignment]
|
||||
|
||||
def __init__(self, inner: PyExpr) -> None:
|
||||
def __init__(self, inner: PyExpr, *, column_path: str | None = None) -> None:
|
||||
self._inner = inner
|
||||
self._column_path = column_path
|
||||
|
||||
# ── comparisons ──────────────────────────────────────────────────────────
|
||||
|
||||
@@ -273,7 +274,7 @@ def col(name: str) -> Expr:
|
||||
>>> col("age") > lit(18)
|
||||
Expr((age > 18))
|
||||
"""
|
||||
return Expr(expr_col(name))
|
||||
return Expr(expr_col(name), column_path=name)
|
||||
|
||||
|
||||
def lit(value: Union[bool, int, float, str, bytes, date, datetime, Decimal]) -> Expr:
|
||||
|
||||
@@ -1,19 +1,24 @@
|
||||
# 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.
|
||||
environment bake, and execution are owned by Sophon.
|
||||
``RefreshColumnResult`` is also the backend-neutral result of a local
|
||||
expression-backed refresh job.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import builtins
|
||||
import base64
|
||||
import functools
|
||||
import hashlib
|
||||
import importlib
|
||||
import inspect
|
||||
import symtable
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
@@ -223,7 +228,7 @@ class PythonEnvironmentSpec(_RemoteValue):
|
||||
|
||||
|
||||
class PythonRuntimeSpec(_RemoteValue):
|
||||
"""Remote runtime definition with non-secret environment values.
|
||||
"""Remote runtime definition with environment values.
|
||||
|
||||
V1 supports ``kind="python"``. Newer runtime kinds remain readable, while
|
||||
their unknown payload fields are intentionally not retained by the client.
|
||||
@@ -262,22 +267,73 @@ class FunctionVersion(_RemoteValue):
|
||||
runtime: PythonRuntimeSpec
|
||||
runtime_digest: str
|
||||
environment_digest: str
|
||||
required_secrets: tuple[str, ...] = ()
|
||||
created_at: str
|
||||
|
||||
def __call__(self, **inputs: Any) -> FunctionApplication:
|
||||
"""Bind this exact version to named table columns.
|
||||
|
||||
Every input must be a direct [lancedb.col][lancedb.expr.col]
|
||||
reference. The returned application is immutable and retains a
|
||||
named-struct output as one binding, so every row's sibling values
|
||||
come from one logical Function evaluation. Map result fields to table
|
||||
columns with
|
||||
[FunctionApplication.rename][lancedb.functions.FunctionApplication.rename],
|
||||
then pass the application to
|
||||
[Table.add_columns][lancedb.table.Table.add_columns].
|
||||
|
||||
Examples
|
||||
--------
|
||||
>>> from lancedb import col
|
||||
>>> application = function( # doctest: +SKIP
|
||||
... title=col("title"),
|
||||
... body=col("body"),
|
||||
... ).rename(columns={
|
||||
... "normalized_text": "search_text",
|
||||
... "token_count": "search_token_count",
|
||||
... })
|
||||
>>> table.add_columns(application) # doctest: +SKIP
|
||||
"""
|
||||
from lancedb.expr import Expr
|
||||
|
||||
parameters = tuple(parameter.name for parameter in self.signature.inputs)
|
||||
missing = [parameter for parameter in parameters if parameter not in inputs]
|
||||
unknown = sorted(set(inputs) - set(parameters))
|
||||
if missing or unknown:
|
||||
details = []
|
||||
if missing:
|
||||
details.append(f"missing inputs: {missing!r}")
|
||||
if unknown:
|
||||
details.append(f"unknown inputs: {unknown!r}")
|
||||
raise TypeError("invalid Function inputs (" + "; ".join(details) + ")")
|
||||
|
||||
bindings = []
|
||||
for parameter in parameters:
|
||||
value = inputs[parameter]
|
||||
if not isinstance(value, Expr) or value._column_path is None:
|
||||
raise TypeError(
|
||||
f"Function input {parameter!r} must be a direct col(...) reference"
|
||||
)
|
||||
bindings.append(
|
||||
ApplicationInput(
|
||||
parameter=parameter,
|
||||
kind="column",
|
||||
value={"path": value._column_path},
|
||||
)
|
||||
)
|
||||
return FunctionApplication(
|
||||
function=FunctionVersionRef(name=self.name, version=self.version),
|
||||
inputs=tuple(bindings),
|
||||
output=self.signature.output,
|
||||
)
|
||||
|
||||
|
||||
class FunctionRegistrationRequest(_RemoteValue):
|
||||
"""Stable remote registration envelope produced by :func:`udf`.
|
||||
|
||||
Only secret names are represented. Secret values are resolved inside the
|
||||
remote service and have no client request field.
|
||||
"""
|
||||
"""Stable remote registration envelope produced by :func:`udf`."""
|
||||
|
||||
name: str
|
||||
artifact: FunctionArtifactRequest
|
||||
signature: FunctionSignature
|
||||
runtime: PythonRuntimeSpec
|
||||
required_secrets: tuple[str, ...] = ()
|
||||
|
||||
|
||||
class FunctionVersionRef(_OpenRemoteValue):
|
||||
@@ -304,12 +360,18 @@ class ApplicationInput(_OpenRemoteValue):
|
||||
|
||||
|
||||
class FunctionApplication(_OpenRemoteValue):
|
||||
"""Immutable pre-declaration application of an exact Function version."""
|
||||
"""Immutable pre-declaration application of an exact Function version.
|
||||
|
||||
A named-struct output remains one application through table
|
||||
declaration and execution.
|
||||
[FunctionApplication.rename][lancedb.functions.FunctionApplication.rename]
|
||||
records the result-field to table-column mapping without splitting sibling
|
||||
outputs into separate UDF calls.
|
||||
"""
|
||||
|
||||
function: FunctionVersionRef
|
||||
inputs: tuple[ApplicationInput, ...]
|
||||
output: FunctionOutput
|
||||
group_id: str
|
||||
columns: Mapping[str, str] = Field(default_factory=dict)
|
||||
|
||||
def _known_dict(self) -> dict[str, Any]:
|
||||
@@ -381,12 +443,10 @@ class OutputMapping(_RemoteValue):
|
||||
|
||||
|
||||
class FunctionBinding(_RemoteValue):
|
||||
"""Immutable grouped binding persisted by the Enterprise table service."""
|
||||
"""Immutable Function binding persisted by the Enterprise table service."""
|
||||
|
||||
binding_id: str
|
||||
revision: _UInt64
|
||||
function: FunctionVersionRef
|
||||
group_id: str
|
||||
inputs: tuple[InputBinding, ...]
|
||||
outputs: tuple[OutputMapping, ...]
|
||||
input_schema: Optional[Mapping[str, Any]] = None
|
||||
@@ -394,7 +454,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
|
||||
@@ -414,62 +478,60 @@ class RefreshColumnResult(_RemoteValue):
|
||||
|
||||
|
||||
_FUNCTION_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_.-]*$")
|
||||
_SECRET_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||
|
||||
|
||||
_GRAMMAR_PRIMITIVES = (
|
||||
(pa.bool_(), "bool"),
|
||||
(pa.int8(), "int8"),
|
||||
(pa.int16(), "int16"),
|
||||
(pa.int32(), "int32"),
|
||||
(pa.int64(), "int64"),
|
||||
(pa.uint8(), "uint8"),
|
||||
(pa.uint16(), "uint16"),
|
||||
(pa.uint32(), "uint32"),
|
||||
(pa.uint64(), "uint64"),
|
||||
(pa.float16(), "float16"),
|
||||
(pa.float32(), "float32"),
|
||||
(pa.float64(), "float64"),
|
||||
(pa.string(), "utf8"),
|
||||
(pa.binary(), "binary"),
|
||||
(pa.date32(), "date32"),
|
||||
(pa.date64(), "date64"),
|
||||
)
|
||||
|
||||
|
||||
def _canonical_arrow_type(data_type: pa.DataType) -> str:
|
||||
primitive_types = (
|
||||
(pa.bool_(), "bool"),
|
||||
(pa.int8(), "int8"),
|
||||
(pa.int16(), "int16"),
|
||||
(pa.int32(), "int32"),
|
||||
(pa.int64(), "int64"),
|
||||
(pa.uint8(), "uint8"),
|
||||
(pa.uint16(), "uint16"),
|
||||
(pa.uint32(), "uint32"),
|
||||
(pa.uint64(), "uint64"),
|
||||
(pa.float16(), "float16"),
|
||||
(pa.float32(), "float32"),
|
||||
(pa.float64(), "float64"),
|
||||
(pa.string(), "utf8"),
|
||||
(pa.large_utf8(), "large_utf8"),
|
||||
(pa.binary(), "binary"),
|
||||
(pa.large_binary(), "large_binary"),
|
||||
(pa.date32(), "date32"),
|
||||
(pa.date64(), "date64"),
|
||||
)
|
||||
for candidate, name in primitive_types:
|
||||
"""The server's V1 Function type grammar. Anything outside it is rejected
|
||||
here rather than at registration."""
|
||||
for candidate, name in _GRAMMAR_PRIMITIVES:
|
||||
if data_type == candidate:
|
||||
return name
|
||||
if pa.types.is_fixed_size_binary(data_type):
|
||||
return f"fixed_size_binary[{data_type.byte_width}]"
|
||||
if pa.types.is_list(data_type):
|
||||
return f"list<{_canonical_arrow_type(data_type.value_type)}>"
|
||||
if pa.types.is_large_list(data_type):
|
||||
return f"large_list<{_canonical_arrow_type(data_type.value_type)}>"
|
||||
if pa.types.is_fixed_size_list(data_type):
|
||||
if pa.types.is_list(data_type) or pa.types.is_large_list(data_type):
|
||||
prefix = "list" if pa.types.is_list(data_type) else "large_list"
|
||||
return f"{prefix}<{_canonical_list_item(data_type)}>"
|
||||
if pa.types.is_fixed_size_list(data_type) and data_type.list_size > 0:
|
||||
return (
|
||||
f"fixed_size_list<{_canonical_arrow_type(data_type.value_type)}>"
|
||||
f"[{data_type.list_size}]"
|
||||
f"fixed_size_list<{_canonical_list_item(data_type)}, {data_type.list_size}>"
|
||||
)
|
||||
if pa.types.is_struct(data_type):
|
||||
fields = ",".join(
|
||||
f"{field.name}:{_canonical_arrow_type(field.type)}" for field in data_type
|
||||
)
|
||||
return f"struct<{fields}>"
|
||||
if pa.types.is_timestamp(data_type):
|
||||
timezone = f",tz={data_type.tz}" if data_type.tz is not None else ""
|
||||
return f"timestamp[{data_type.unit}{timezone}]"
|
||||
if pa.types.is_time32(data_type) or pa.types.is_time64(data_type):
|
||||
return f"time[{data_type.unit}]"
|
||||
if pa.types.is_duration(data_type):
|
||||
return f"duration[{data_type.unit}]"
|
||||
if pa.types.is_decimal(data_type):
|
||||
bit_width = data_type.bit_width
|
||||
return f"decimal{bit_width}({data_type.precision},{data_type.scale})"
|
||||
raise TypeError(f"unsupported Arrow type for Function signature: {data_type}")
|
||||
|
||||
|
||||
def _canonical_list_item(data_type: pa.DataType) -> str:
|
||||
"""The grammar names only the item type; it always means a non-nullable
|
||||
child called `item`, so any other child metadata cannot be represented."""
|
||||
child = data_type.value_field
|
||||
if child.name != "item" or child.nullable or child.metadata:
|
||||
raise TypeError(
|
||||
"unsupported Arrow type for Function signature: list items must be a "
|
||||
f"non-nullable field named 'item', got {child}"
|
||||
)
|
||||
return _canonical_arrow_type(child.type)
|
||||
|
||||
|
||||
def _list_of(item: pa.DataType) -> pa.DataType:
|
||||
return pa.list_(pa.field("item", item, nullable=False))
|
||||
|
||||
|
||||
def _annotation_type(annotation: Any) -> tuple[pa.DataType, bool]:
|
||||
nullable = False
|
||||
origin = get_origin(annotation)
|
||||
@@ -517,7 +579,7 @@ def _annotation_type(annotation: Any) -> tuple[pa.DataType, bool]:
|
||||
value_type, value_nullable = _annotation_type(arguments[0])
|
||||
if value_nullable:
|
||||
raise TypeError("nullable Function list elements are not supported")
|
||||
return pa.list_(value_type), nullable
|
||||
return _list_of(value_type), nullable
|
||||
raise TypeError(f"unsupported Function annotation: {annotation!r}")
|
||||
|
||||
|
||||
@@ -664,6 +726,104 @@ def _literal_source(value: Any) -> str:
|
||||
)
|
||||
|
||||
|
||||
_DYNAMIC_NAMESPACE_ACCESS = frozenset(
|
||||
{"globals", "locals", "vars", "eval", "exec", "compile", "__import__"}
|
||||
)
|
||||
# Modules that hand out namespaces (`sys.modules`, `builtins`, importers,
|
||||
# introspection). The artifact's module namespace holds only the names it was
|
||||
# packaged with, so reaching around it cannot be represented.
|
||||
_NAMESPACE_MODULES = frozenset(
|
||||
{"sys", "builtins", "importlib", "inspect", "gc", "ctypes", "types"}
|
||||
)
|
||||
|
||||
|
||||
def _namespace_acquisition(
|
||||
definition: ast.FunctionDef, references: set[str]
|
||||
) -> list[str]:
|
||||
found = set(references & _DYNAMIC_NAMESPACE_ACCESS)
|
||||
for node in ast.walk(definition):
|
||||
if isinstance(node, ast.Import):
|
||||
found.update(
|
||||
alias.name
|
||||
for alias in node.names
|
||||
if alias.name.split(".")[0] in _NAMESPACE_MODULES
|
||||
)
|
||||
elif isinstance(node, ast.ImportFrom) and node.module:
|
||||
if node.module.split(".")[0] in _NAMESPACE_MODULES:
|
||||
found.add(node.module)
|
||||
return sorted(found)
|
||||
|
||||
|
||||
def _module_references(module_source: str) -> set[str]:
|
||||
"""Names any scope in `module_source` binds or loads at module scope.
|
||||
Python's own scope analysis on the exact text that ships: free variables
|
||||
belong to an enclosing scope inside the function, and postponed
|
||||
annotations are not runtime loads."""
|
||||
|
||||
def visit(table: symtable.SymbolTable, found: set[str]) -> None:
|
||||
for symbol in table.get_symbols():
|
||||
if symbol.is_global() and (
|
||||
symbol.is_referenced() or symbol.is_declared_global()
|
||||
):
|
||||
found.add(symbol.get_name())
|
||||
for child in table.get_children():
|
||||
visit(child, found)
|
||||
|
||||
found: set[str] = set()
|
||||
for table in symtable.symtable(module_source, "<udf>", "exec").get_children():
|
||||
visit(table, found)
|
||||
return found
|
||||
|
||||
|
||||
def _global_source(name: str, value: Any) -> str:
|
||||
"""One module-level line that rebinds `name` to `value` in the artifact:
|
||||
an import for modules and importable classes/functions, a literal otherwise."""
|
||||
if isinstance(value, types.ModuleType):
|
||||
if value.__name__.split(".")[0] in _NAMESPACE_MODULES:
|
||||
raise ValueError(
|
||||
f"@udf cannot package dynamic namespace access: {value.__name__!r}"
|
||||
)
|
||||
try:
|
||||
imported = importlib.import_module(value.__name__)
|
||||
except ImportError:
|
||||
imported = None
|
||||
if imported is not value:
|
||||
raise TypeError(
|
||||
f"Function source references module {name!r} that does not import "
|
||||
f"as {value.__name__!r}"
|
||||
)
|
||||
return f"import {value.__name__} as {name}"
|
||||
module_name = getattr(value, "__module__", None)
|
||||
qualname = getattr(value, "__qualname__", None)
|
||||
if (
|
||||
isinstance(module_name, str)
|
||||
and isinstance(qualname, str)
|
||||
and module_name != "__main__"
|
||||
and "." not in qualname
|
||||
and "<" not in qualname
|
||||
):
|
||||
try:
|
||||
imported = getattr(importlib.import_module(module_name), qualname)
|
||||
except (ImportError, AttributeError):
|
||||
imported = None
|
||||
if imported is value:
|
||||
return f"from {module_name} import {qualname} as {name}"
|
||||
return f"{name} = {_literal_source(value)}"
|
||||
|
||||
|
||||
def _is_recursive_reference(function: Callable[..., Any], name: str) -> bool:
|
||||
"""`name` inside the body means the function itself unless the module has
|
||||
since bound it to something else."""
|
||||
if name != function.__name__:
|
||||
return False
|
||||
bound = function.__globals__.get(name, function)
|
||||
if bound is function:
|
||||
return True
|
||||
# The decorator's own result is the one wrapper known to call `function`
|
||||
# unchanged; any other binding may behave differently from a self-call.
|
||||
return type(bound) is UdfDefinition and bound._function is function
|
||||
|
||||
|
||||
def _package_source(function: Callable[..., Any]) -> bytes:
|
||||
if not inspect.isfunction(function) or inspect.iscoroutinefunction(function):
|
||||
raise TypeError("@udf requires a synchronous Python function")
|
||||
@@ -688,23 +848,46 @@ def _package_source(function: Callable[..., Any]) -> bytes:
|
||||
closure = inspect.getclosurevars(function)
|
||||
if closure.nonlocals:
|
||||
raise ValueError("@udf cannot package functions that capture closure values")
|
||||
if closure.unbound:
|
||||
raise ValueError(
|
||||
f"@udf source contains unresolved global names: {sorted(closure.unbound)!r}"
|
||||
)
|
||||
globals_source = []
|
||||
for name, value in sorted(closure.globals.items()):
|
||||
if isinstance(value, types.ModuleType):
|
||||
globals_source.append(f"import {value.__name__} as {name}")
|
||||
else:
|
||||
globals_source.append(f"{name} = {_literal_source(value)}")
|
||||
|
||||
function_source = ast.unparse(definition)
|
||||
parts = ["from __future__ import annotations"]
|
||||
module_header = "from __future__ import annotations"
|
||||
references = _module_references(f"{module_header}\n\n{function_source}\n")
|
||||
dynamic = _namespace_acquisition(definition, references)
|
||||
if dynamic:
|
||||
raise ValueError(f"@udf cannot package dynamic namespace access: {dynamic!r}")
|
||||
# Resolve every module-scope reference the way the interpreter would: the
|
||||
# function's own globals first (a module global may shadow a builtin, and
|
||||
# nested scopes are not visible to getclosurevars), then its builtins.
|
||||
# The artifact runs under the standard builtins; only the exact mapping is
|
||||
# provably equivalent (a subclass or copy can change lookups and hooks).
|
||||
if function.__builtins__ is not vars(builtins):
|
||||
raise ValueError("@udf cannot package a non-standard builtins environment")
|
||||
globals_source = []
|
||||
unresolved = []
|
||||
for name in sorted(references):
|
||||
if name == function.__name__:
|
||||
if not _is_recursive_reference(function, name):
|
||||
raise ValueError(
|
||||
f"@udf cannot package {name!r}: the module binds that name to "
|
||||
"another value, which the artifact's own definition would shadow"
|
||||
)
|
||||
continue
|
||||
if name in function.__globals__:
|
||||
globals_source.append(_global_source(name, function.__globals__[name]))
|
||||
elif hasattr(builtins, name):
|
||||
pass
|
||||
else:
|
||||
unresolved.append(name)
|
||||
if unresolved:
|
||||
raise ValueError(
|
||||
f"@udf source contains unresolved global names: {unresolved!r}"
|
||||
)
|
||||
|
||||
parts = [module_header]
|
||||
if globals_source:
|
||||
parts.extend(["", *globals_source])
|
||||
parts.extend(["", function_source, ""])
|
||||
return "\n".join(parts).encode("utf-8")
|
||||
packaged = "\n".join(parts)
|
||||
return packaged.encode("utf-8")
|
||||
|
||||
|
||||
class UdfDefinition:
|
||||
@@ -725,7 +908,6 @@ class UdfDefinition:
|
||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema],
|
||||
pip: tuple[str, ...],
|
||||
env: Mapping[str, str],
|
||||
secrets: tuple[str, ...],
|
||||
python_version: Optional[str],
|
||||
):
|
||||
function_name = name or function.__name__
|
||||
@@ -740,18 +922,6 @@ class UdfDefinition:
|
||||
for key, value in environment.items()
|
||||
):
|
||||
raise TypeError("Function env keys and values must be strings")
|
||||
required_secrets = tuple(sorted(set(secrets)))
|
||||
invalid_secrets = [
|
||||
secret for secret in required_secrets if not _SECRET_NAME.fullmatch(secret)
|
||||
]
|
||||
if invalid_secrets:
|
||||
raise ValueError(f"invalid Function secret names: {invalid_secrets!r}")
|
||||
overlap = set(environment) & set(required_secrets)
|
||||
if overlap:
|
||||
raise ValueError(
|
||||
f"Function env and secret names must be disjoint: {sorted(overlap)!r}"
|
||||
)
|
||||
|
||||
signature = _infer_signature(function, input_schema, output_schema)
|
||||
source = _package_source(function)
|
||||
digest = f"sha256:{hashlib.sha256(source).hexdigest()}"
|
||||
@@ -780,7 +950,6 @@ class UdfDefinition:
|
||||
),
|
||||
signature=signature,
|
||||
runtime=runtime,
|
||||
required_secrets=required_secrets,
|
||||
)
|
||||
functools.update_wrapper(self, function)
|
||||
|
||||
@@ -806,7 +975,6 @@ def udf(
|
||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
||||
pip: tuple[str, ...] | list[str] = (),
|
||||
env: Optional[Mapping[str, str]] = None,
|
||||
secrets: tuple[str, ...] | list[str] = (),
|
||||
python_version: Optional[str] = None,
|
||||
) -> Callable[[Callable[..., Any]], UdfDefinition]: ...
|
||||
|
||||
@@ -819,7 +987,6 @@ def udf(
|
||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
||||
pip: tuple[str, ...] | list[str] = (),
|
||||
env: Optional[Mapping[str, str]] = None,
|
||||
secrets: tuple[str, ...] | list[str] = (),
|
||||
python_version: Optional[str] = None,
|
||||
):
|
||||
"""Prepare a scalar Python callable for remote Function registration.
|
||||
@@ -844,13 +1011,17 @@ def udf(
|
||||
pip : sequence of str, optional
|
||||
Pip requirements for the remote environment.
|
||||
env : mapping of str to str, optional
|
||||
Non-secret environment variables. Use ``secrets`` for credentials.
|
||||
secrets : sequence of str, optional
|
||||
Names of secrets resolved by the remote service. Secret values are not
|
||||
accepted by this API or included in the registration request.
|
||||
Environment variables included in the Function definition.
|
||||
python_version : str, optional
|
||||
Remote Python major/minor version. Defaults to the client version.
|
||||
|
||||
The packaged artifact is a snapshot: the function source plus exactly
|
||||
the module-level names it references (modules as imports, importable
|
||||
classes and functions as imports, literals inline). Code that reaches the
|
||||
module namespace another way -- ``globals()``/``eval``, ``sys.modules``,
|
||||
``builtins`` -- is rejected where it can be seen and otherwise
|
||||
unsupported; closures and a non-standard ``__builtins__`` are rejected.
|
||||
|
||||
Returns
|
||||
-------
|
||||
UdfDefinition
|
||||
@@ -862,7 +1033,7 @@ def udf(
|
||||
Examples
|
||||
--------
|
||||
>>> from lancedb import udf
|
||||
>>> @udf(pip=["numpy==2.2.0"], secrets=["MODEL_TOKEN"])
|
||||
>>> @udf(pip=["numpy==2.2.0"])
|
||||
... def score(value: float) -> float:
|
||||
... return value * 2
|
||||
>>> score(1.5)
|
||||
@@ -877,7 +1048,6 @@ def udf(
|
||||
output_schema=output_schema,
|
||||
pip=tuple(pip),
|
||||
env={} if env is None else env,
|
||||
secrets=tuple(secrets),
|
||||
python_version=python_version,
|
||||
)
|
||||
|
||||
|
||||
@@ -163,6 +163,15 @@ class FTS:
|
||||
The number of documents per compressed posting block. Supported values
|
||||
are 128 and 256. A value of 256 uses the experimental FTS V3 format
|
||||
and may introduce breaking changes.
|
||||
memory_limit : int, optional
|
||||
The total memory limit in MiB for the local FTS build stage. The limit
|
||||
is divided evenly among indexing workers. This build-only setting is
|
||||
not persisted with the index and does not apply to remote tables.
|
||||
num_workers : int, optional
|
||||
The number of workers for a local FTS build. By default Lance uses
|
||||
roughly half of the available CPU cores. The effective value is
|
||||
limited by the available compute capacity. This build-only setting is
|
||||
not persisted with the index and does not apply to remote tables.
|
||||
|
||||
Notes
|
||||
-----
|
||||
@@ -185,6 +194,8 @@ class FTS:
|
||||
prefix_only: bool = False
|
||||
block_size: int = 128
|
||||
custom_stop_words: Optional[List[str]] = None
|
||||
memory_limit: Optional[int] = None
|
||||
num_workers: Optional[int] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -41,21 +41,15 @@ class PermutationBuilder:
|
||||
The permutation is stored in memory and will be lost when the program exits.
|
||||
"""
|
||||
|
||||
def __init__(self, table: LanceTable):
|
||||
def __init__(self, table: Table):
|
||||
"""
|
||||
Creates a new permutation builder for the given table.
|
||||
|
||||
By default, the permutation builder will create a single split that contains all
|
||||
rows in the same order as the base table.
|
||||
|
||||
Tables with an LSM write spec are rejected: unflushed rows have no row id.
|
||||
"""
|
||||
if not hasattr(table, "_inner"):
|
||||
raise TypeError(
|
||||
f"PermutationBuilder requires a local LanceTable, "
|
||||
f"got {type(table).__name__}. "
|
||||
"The permutation API is not supported on remote tables. "
|
||||
"Remote tables connect to LanceDB Cloud or Enterprise and do not have "
|
||||
"direct access to the underlying Lance dataset needed for permutations."
|
||||
)
|
||||
self._async = async_permutation_builder(table)
|
||||
|
||||
def split_random(
|
||||
@@ -231,7 +225,7 @@ class PermutationBuilder:
|
||||
return LOOP.run(do_execute())
|
||||
|
||||
|
||||
def permutation_builder(table: LanceTable) -> PermutationBuilder:
|
||||
def permutation_builder(table: Table) -> PermutationBuilder:
|
||||
return PermutationBuilder(table)
|
||||
|
||||
|
||||
@@ -248,7 +242,7 @@ class Permutations:
|
||||
|
||||
Attributes
|
||||
----------
|
||||
base_table: LanceTable
|
||||
base_table: Table
|
||||
The base table that the permutations are based on.
|
||||
permutation_table: LanceTable
|
||||
The permutation table that defines the splits.
|
||||
@@ -282,7 +276,7 @@ class Permutations:
|
||||
{'train': 0, 'test': 1}
|
||||
"""
|
||||
|
||||
def __init__(self, base_table: LanceTable, permutation_table: LanceTable):
|
||||
def __init__(self, base_table: Table, permutation_table: LanceTable):
|
||||
self.base_table = base_table
|
||||
self.permutation_table = permutation_table
|
||||
|
||||
@@ -397,6 +391,15 @@ def _table_to_pickle_state(table: Table) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _drop_base_version(permutation_data: pa.Table) -> pa.Table:
|
||||
"""Strip the recorded base version so the reader leaves the base table unpinned."""
|
||||
metadata = dict(permutation_data.schema.metadata or {})
|
||||
if metadata.pop(b"base_version", None) is None:
|
||||
return permutation_data
|
||||
metadata.pop(b"base_branch", None)
|
||||
return permutation_data.replace_schema_metadata(metadata)
|
||||
|
||||
|
||||
def _table_from_pickle_state(state: dict[str, Any]) -> Table:
|
||||
from . import connect
|
||||
|
||||
@@ -685,11 +688,15 @@ class Permutation:
|
||||
from . import connect
|
||||
|
||||
connection_factory = state["connection_factory"]
|
||||
rebuilt_base = False
|
||||
if connection_factory is not None:
|
||||
base_table = connection_factory(state["base_table_name"])
|
||||
elif "base_table_state" in state:
|
||||
base_table = _table_from_pickle_state(state["base_table_state"])
|
||||
base_state = state["base_table_state"]
|
||||
rebuilt_base = base_state["kind"] == "memory"
|
||||
base_table = _table_from_pickle_state(base_state)
|
||||
elif "base_table_data" in state:
|
||||
rebuilt_base = True
|
||||
# In-memory base table inlined into the pickle; rebuild the same
|
||||
# way we rebuild the in-memory permutation table.
|
||||
mem_db = connect("memory://")
|
||||
@@ -707,11 +714,14 @@ class Permutation:
|
||||
)
|
||||
|
||||
permutation_table: Optional[Table] = None
|
||||
if state["permutation_data"] is not None:
|
||||
permutation_data = state["permutation_data"]
|
||||
if permutation_data is not None:
|
||||
if rebuilt_base:
|
||||
# The base table was materialized from Arrow, so it is a fresh
|
||||
# single-version dataset and the recorded pin cannot resolve on it.
|
||||
permutation_data = _drop_base_version(permutation_data)
|
||||
mem_db = connect("memory://")
|
||||
permutation_table = mem_db.create_table(
|
||||
"permutation", state["permutation_data"]
|
||||
)
|
||||
permutation_table = mem_db.create_table("permutation", permutation_data)
|
||||
|
||||
self.base_table = base_table
|
||||
self.permutation_table = permutation_table
|
||||
|
||||
@@ -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
|
||||
@@ -610,6 +610,7 @@ class RemoteTable(Table):
|
||||
fill_value: float = 0.0,
|
||||
progress: Optional[Union[bool, Callable, Any]] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
allow_external_blob_outside_bases: bool = False,
|
||||
) -> AddResult:
|
||||
"""Add more data to the [Table][lancedb.table.Table].
|
||||
|
||||
@@ -642,6 +643,8 @@ class RemoteTable(Table):
|
||||
data in flight. Defaults to an estimate based on the data size,
|
||||
capped at the number of CPU cores. Lower this if bulk ingestion is
|
||||
using too much memory.
|
||||
allow_external_blob_outside_bases: bool, default False
|
||||
Not supported on LanceDB Cloud. Setting this raises.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -658,6 +661,7 @@ class RemoteTable(Table):
|
||||
fill_value=fill_value,
|
||||
progress=progress,
|
||||
write_parallelism=write_parallelism,
|
||||
allow_external_blob_outside_bases=allow_external_blob_outside_bases,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
@@ -972,7 +976,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(
|
||||
|
||||
+1037
-91
File diff suppressed because it is too large
Load Diff
@@ -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 (
|
||||
@@ -1266,6 +1269,7 @@ class Table(ABC):
|
||||
fill_value: float = 0.0,
|
||||
progress: Optional[Union[bool, Callable, Any]] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
allow_external_blob_outside_bases: bool = False,
|
||||
) -> AddResult:
|
||||
"""Add more data to the [Table][lancedb.table.Table].
|
||||
|
||||
@@ -1317,6 +1321,10 @@ class Table(ABC):
|
||||
data in flight. Defaults to an estimate based on the data size,
|
||||
capped at the number of CPU cores. Lower this if bulk ingestion is
|
||||
using too much memory.
|
||||
allow_external_blob_outside_bases: bool, default False
|
||||
Store blob URIs that sit outside registered blob bases. The row
|
||||
keeps a reference, so the object has to stay readable. Local
|
||||
tables only.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -1969,7 +1977,7 @@ class Table(ABC):
|
||||
A mapping with one ``FunctionApplication`` value keeps its scalar
|
||||
or named-struct result in the named table column. A bare
|
||||
named-struct application expands its ordered result fields as one
|
||||
atomic sibling group; aliases come from ``rename(columns=...)``.
|
||||
atomic binding; aliases come from ``rename(columns=...)``.
|
||||
Function columns are supported only on LanceDB Cloud and
|
||||
Enterprise.
|
||||
computed: Dict[str, str], optional
|
||||
@@ -2039,7 +2047,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 +2058,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 +2072,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'
|
||||
"""
|
||||
@@ -3398,6 +3414,7 @@ class LanceTable(Table):
|
||||
fill_value: float = 0.0,
|
||||
progress: Optional[Union[bool, Callable, Any]] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
allow_external_blob_outside_bases: bool = False,
|
||||
) -> AddResult:
|
||||
"""Add data to the table.
|
||||
If vector columns are missing and the table
|
||||
@@ -3425,6 +3442,9 @@ class LanceTable(Table):
|
||||
data in flight. Defaults to an estimate based on the data size,
|
||||
capped at the number of CPU cores. Lower this if bulk ingestion is
|
||||
using too much memory.
|
||||
allow_external_blob_outside_bases: bool, default False
|
||||
Allow blob URIs outside registered bases. See :meth:`Table.add`.
|
||||
Local tables only.
|
||||
|
||||
Returns
|
||||
-------
|
||||
@@ -3441,6 +3461,7 @@ class LanceTable(Table):
|
||||
fill_value=fill_value,
|
||||
progress=progress,
|
||||
write_parallelism=write_parallelism,
|
||||
allow_external_blob_outside_bases=allow_external_blob_outside_bases,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
@@ -4082,7 +4103,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].
|
||||
@@ -5354,6 +5375,7 @@ class AsyncTable:
|
||||
fill_value: Optional[float] = None,
|
||||
progress: Optional[Union[bool, Callable, Any]] = None,
|
||||
write_parallelism: Optional[int] = None,
|
||||
allow_external_blob_outside_bases: bool = False,
|
||||
) -> AddResult:
|
||||
"""Add more data to the [AsyncTable][lancedb.table.AsyncTable].
|
||||
|
||||
@@ -5384,6 +5406,9 @@ class AsyncTable:
|
||||
data in flight. Defaults to an estimate based on the data size,
|
||||
capped at the number of CPU cores. Lower this if bulk ingestion is
|
||||
using too much memory.
|
||||
allow_external_blob_outside_bases: bool, default False
|
||||
Allow blob URIs outside registered bases. See :meth:`Table.add`.
|
||||
Local tables only.
|
||||
|
||||
"""
|
||||
schema = await self.schema()
|
||||
@@ -5420,6 +5445,7 @@ class AsyncTable:
|
||||
mode or "append",
|
||||
progress=progress,
|
||||
write_parallelism=write_parallelism,
|
||||
allow_external_blob_outside_bases=allow_external_blob_outside_bases,
|
||||
)
|
||||
except RuntimeError as e:
|
||||
if "Cast error" in str(e):
|
||||
@@ -6027,7 +6053,7 @@ class AsyncTable:
|
||||
A mapping with one ``FunctionApplication`` value keeps its scalar
|
||||
or named-struct result in the named table column. A bare
|
||||
named-struct application expands its ordered result fields as one
|
||||
atomic sibling group; aliases come from ``rename(columns=...)``.
|
||||
atomic binding; aliases come from ``rename(columns=...)``.
|
||||
Function columns are supported only on LanceDB Cloud and
|
||||
Enterprise.
|
||||
computed: Dict[str, str], optional
|
||||
@@ -6064,7 +6090,7 @@ class AsyncTable:
|
||||
isinstance(value, FunctionApplication) for value in transforms.values()
|
||||
):
|
||||
raise ValueError(
|
||||
"one add_columns call declares exactly one Function sibling group"
|
||||
"one add_columns call declares exactly one Function binding"
|
||||
)
|
||||
function_output_name, function_application = next(iter(transforms.items()))
|
||||
|
||||
@@ -6122,7 +6148,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 +6162,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 +6177,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]]
|
||||
@@ -6801,21 +6839,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
|
||||
@@ -6951,9 +6989,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)
|
||||
|
||||
@@ -617,3 +617,71 @@ def test_fetch_blobs_nested_path_survives_sort_after_query():
|
||||
def _identifiable_payload(size: int) -> bytes:
|
||||
block = 256
|
||||
return b"".join(bytes([i % 256]) * block for i in range(size // block))
|
||||
|
||||
|
||||
def _external_uri_blob_array(uris):
|
||||
blob_type = lancedb.blob("image").type
|
||||
storage_type = blob_type.storage_type
|
||||
child_names = [field.name for field in storage_type]
|
||||
assert "uri" in child_names, "blob layout no longer has a uri child"
|
||||
children = [
|
||||
pa.array(uris if field.name == "uri" else [None] * len(uris), type=field.type)
|
||||
for field in storage_type
|
||||
]
|
||||
storage = pa.StructArray.from_arrays(children, fields=list(storage_type))
|
||||
return pa.ExtensionArray.from_storage(blob_type, storage)
|
||||
|
||||
|
||||
def _external_uri_table_and_rows(name, uris):
|
||||
db = lancedb.connect("memory:///")
|
||||
schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")])
|
||||
table = db.create_table(name, schema=schema)
|
||||
rows = pa.Table.from_arrays(
|
||||
[
|
||||
pa.array(range(len(uris)), type=pa.int64()),
|
||||
_external_uri_blob_array(uris),
|
||||
],
|
||||
schema=schema,
|
||||
)
|
||||
return table, rows
|
||||
|
||||
|
||||
def test_add_external_uri_struct_round_trips_with_flag(tmp_path):
|
||||
payload = b"external-uri-bytes"
|
||||
blob_path = tmp_path / "payload.bin"
|
||||
blob_path.write_bytes(payload)
|
||||
|
||||
table, rows = _external_uri_table_and_rows("external_struct", [blob_path.as_uri()])
|
||||
table.add(rows, allow_external_blob_outside_bases=True)
|
||||
|
||||
hits = table.search().to_arrow()
|
||||
blobs = table.fetch_blobs("image", hits)
|
||||
assert blobs[0].as_py() == payload
|
||||
|
||||
|
||||
def test_add_external_uri_without_flag_raises(tmp_path):
|
||||
blob_path = tmp_path / "payload.bin"
|
||||
blob_path.write_bytes(b"unreachable")
|
||||
|
||||
table, rows = _external_uri_table_and_rows("external_no_flag", [blob_path.as_uri()])
|
||||
with pytest.raises(ValueError, match="allow_external_blob_outside_bases"):
|
||||
table.add(rows)
|
||||
assert table.count_rows() == 0
|
||||
|
||||
|
||||
def test_add_external_uri_string_round_trips_with_flag(tmp_path):
|
||||
payload = b"external-uri-bytes"
|
||||
blob_path = tmp_path / "payload.bin"
|
||||
blob_path.write_bytes(payload)
|
||||
|
||||
db = lancedb.connect("memory:///")
|
||||
schema = pa.schema([pa.field("id", pa.int64()), lancedb.blob("image")])
|
||||
table = db.create_table("external_string", schema=schema)
|
||||
table.add(
|
||||
[{"id": 1, "image": blob_path.as_uri()}],
|
||||
allow_external_blob_outside_bases=True,
|
||||
)
|
||||
|
||||
hits = table.search().to_arrow()
|
||||
blobs = table.fetch_blobs("image", hits)
|
||||
assert blobs[0].as_py() == payload
|
||||
|
||||
@@ -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() == []
|
||||
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -6,6 +6,7 @@ from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from lancedb import col
|
||||
import lancedb.functions as functions
|
||||
from lancedb.functions import (
|
||||
FunctionApplication,
|
||||
@@ -36,21 +37,6 @@ def job_result(name: str) -> dict:
|
||||
return json.loads(fixture(name))["result"]
|
||||
|
||||
|
||||
def assert_no_secret_values(value):
|
||||
if isinstance(value, dict):
|
||||
for key, child in value.items():
|
||||
assert key not in {
|
||||
"secret_value",
|
||||
"secret_values",
|
||||
"resolved_secret",
|
||||
"resolved_secrets",
|
||||
}
|
||||
assert_no_secret_values(child)
|
||||
elif isinstance(value, list):
|
||||
for child in value:
|
||||
assert_no_secret_values(child)
|
||||
|
||||
|
||||
def test_public_function_values_are_in_api_reference():
|
||||
docs = Path(__file__).parents[3] / "docs" / "src" / "python" / "python.md"
|
||||
rendered = docs.read_text()
|
||||
@@ -108,7 +94,6 @@ def test_function_version_identity_is_immutable_and_exact():
|
||||
version = FunctionVersion.from_json(json.dumps(value))
|
||||
assert version.name == "embed"
|
||||
assert version.version == "fv_01K3EXACT"
|
||||
assert version.required_secrets == ("HF_TOKEN",)
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
version.version = "fv_changed"
|
||||
@@ -120,6 +105,80 @@ def test_function_version_identity_is_immutable_and_exact():
|
||||
assert FunctionVersion(**changed) != version
|
||||
|
||||
|
||||
def test_function_version_binds_named_columns_as_one_immutable_application():
|
||||
version = FunctionVersion.from_json(
|
||||
json.dumps(job_result("remote_function_job.json"))
|
||||
)
|
||||
|
||||
application = version(text=col("documents.body"))
|
||||
|
||||
assert application.function.name == version.name
|
||||
assert application.function.version == version.version
|
||||
assert application.output is version.signature.output
|
||||
assert [
|
||||
(value.parameter, value.kind, value.value["path"])
|
||||
for value in application.inputs
|
||||
] == [("text", "column", "documents.body")]
|
||||
|
||||
|
||||
def test_function_version_binding_validates_names_and_direct_columns():
|
||||
version = FunctionVersion.from_json(
|
||||
json.dumps(job_result("remote_function_job.json"))
|
||||
)
|
||||
|
||||
with pytest.raises(TypeError, match=r"missing inputs: \['text'\]"):
|
||||
version()
|
||||
with pytest.raises(TypeError, match=r"unknown inputs: \['body'\]"):
|
||||
version(text=col("text"), body=col("body"))
|
||||
with pytest.raises(TypeError, match="direct col"):
|
||||
version(text=col("text").lower())
|
||||
|
||||
|
||||
def test_function_version_keeps_named_struct_outputs_in_one_application():
|
||||
value = job_result("remote_function_job.json")
|
||||
value["name"] = "text_features"
|
||||
value["version"] = "fv_multi_output"
|
||||
value["signature"] = {
|
||||
"inputs": [
|
||||
{"name": "title", "arrow_type": "utf8", "nullable": True},
|
||||
{"name": "body", "arrow_type": "utf8", "nullable": True},
|
||||
],
|
||||
"output": {
|
||||
"kind": "named_struct",
|
||||
"fields": [
|
||||
{
|
||||
"name": "normalized_text",
|
||||
"arrow_type": "utf8",
|
||||
"nullable": False,
|
||||
},
|
||||
{
|
||||
"name": "token_count",
|
||||
"arrow_type": "int64",
|
||||
"nullable": False,
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
version = FunctionVersion(**value)
|
||||
|
||||
application = version(body=col("body"), title=col("title")).rename(
|
||||
columns={
|
||||
"normalized_text": "search_text",
|
||||
"token_count": "search_token_count",
|
||||
}
|
||||
)
|
||||
|
||||
assert [value.parameter for value in application.inputs] == ["title", "body"]
|
||||
assert [field.name for field in application.output.fields] == [
|
||||
"normalized_text",
|
||||
"token_count",
|
||||
]
|
||||
assert dict(application.columns) == {
|
||||
"normalized_text": "search_text",
|
||||
"token_count": "search_token_count",
|
||||
}
|
||||
|
||||
|
||||
def test_unknown_fields_and_discriminators_are_forward_decodable():
|
||||
value = job_result("remote_function_job.json")
|
||||
value["future_version_metadata"] = {"retention_class": "catalog"}
|
||||
@@ -143,7 +202,6 @@ def test_function_application_uses_rename_columns_only():
|
||||
assert application.columns["normalized_text"] == "search_text"
|
||||
assert renamed.columns["normalized_text"] == "body_normalized"
|
||||
assert renamed.function == application.function
|
||||
assert renamed.group_id == application.group_id
|
||||
assert not hasattr(application, "rename_outputs")
|
||||
with pytest.raises(TypeError, match="immutable"):
|
||||
renamed.columns["normalized_text"] = "changed"
|
||||
@@ -164,7 +222,6 @@ def test_function_application_uses_rename_columns_only():
|
||||
|
||||
def test_binding_and_refresh_result_keep_stable_remote_fields():
|
||||
binding = FunctionBinding.from_json(fixture("remote_function_binding.json"))
|
||||
assert binding.revision == 3
|
||||
assert binding.function.version == "fv_01K3TEXT"
|
||||
assert [output.output_ordinal for output in binding.outputs] == [0, 1]
|
||||
assert binding.input_schema is not None
|
||||
@@ -219,15 +276,6 @@ def test_refresh_result_rejects_non_u64_values(field):
|
||||
RefreshColumnResult.from_json(json.dumps(value))
|
||||
|
||||
|
||||
def test_canonical_client_values_contain_secret_names_only():
|
||||
version = FunctionVersion.from_json(
|
||||
json.dumps(job_result("remote_function_job.json"))
|
||||
)
|
||||
canonical = json.loads(version.to_canonical_json())
|
||||
assert canonical["required_secrets"] == ["HF_TOKEN"]
|
||||
assert_no_secret_values(canonical)
|
||||
|
||||
|
||||
class _FunctionDeclarationInner:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
@@ -244,7 +292,7 @@ def known_application() -> FunctionApplication:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_columns_routes_struct_as_one_and_grouped_expansion_atomically():
|
||||
async def test_add_columns_routes_struct_as_one_and_multi_output_binding_atomically():
|
||||
inner = _FunctionDeclarationInner()
|
||||
table = AsyncTable(inner)
|
||||
application = known_application()
|
||||
@@ -265,12 +313,12 @@ async def test_add_columns_routes_struct_as_one_and_grouped_expansion_atomically
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_columns_rejects_mixed_groups_and_unknown_newer_application():
|
||||
async def test_add_columns_rejects_multiple_bindings_and_unknown_newer_application():
|
||||
inner = _FunctionDeclarationInner()
|
||||
table = AsyncTable(inner)
|
||||
application = known_application()
|
||||
|
||||
with pytest.raises(ValueError, match="exactly one Function sibling group"):
|
||||
with pytest.raises(ValueError, match="exactly one Function binding"):
|
||||
await table.add_columns({"a": application, "b": application})
|
||||
|
||||
future = json.loads(fixture("remote_function_application.json"))
|
||||
@@ -298,7 +346,6 @@ def test_rename_requires_named_struct_and_keeps_partial_mapping_immutable():
|
||||
"arrow_type": "list<float32>",
|
||||
"nullable": False,
|
||||
},
|
||||
"group_id": "fg_scalar",
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@@ -3,7 +3,12 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import contextlib
|
||||
import functools
|
||||
import importlib.util
|
||||
import types
|
||||
from datetime import date
|
||||
import http.server
|
||||
import json
|
||||
from pathlib import Path
|
||||
@@ -16,6 +21,9 @@ import pytest
|
||||
import lancedb
|
||||
from lancedb.functions import UdfDefinition, udf
|
||||
|
||||
THRESHOLD = 20
|
||||
_CACHE = None
|
||||
|
||||
|
||||
FIXTURES = (
|
||||
Path(__file__).parents[3]
|
||||
@@ -31,28 +39,12 @@ FIXTURES = (
|
||||
@udf(
|
||||
pip=["numpy>=2"],
|
||||
env={"MODE": "test"},
|
||||
secrets=["API_TOKEN"],
|
||||
python_version="3.12",
|
||||
)
|
||||
def normalize_score(value: float) -> float:
|
||||
return value / 100.0
|
||||
|
||||
|
||||
def _assert_no_secret_values(value):
|
||||
if isinstance(value, dict):
|
||||
for key, child in value.items():
|
||||
assert key not in {
|
||||
"secret_value",
|
||||
"secret_values",
|
||||
"resolved_secret",
|
||||
"resolved_secrets",
|
||||
}
|
||||
_assert_no_secret_values(child)
|
||||
elif isinstance(value, list):
|
||||
for child in value:
|
||||
_assert_no_secret_values(child)
|
||||
|
||||
|
||||
def test_scalar_udf_matches_shared_registration_golden_and_remains_callable():
|
||||
assert isinstance(normalize_score, UdfDefinition)
|
||||
assert normalize_score(25.0) == 0.25
|
||||
@@ -67,13 +59,397 @@ def test_scalar_udf_matches_shared_registration_golden_and_remains_callable():
|
||||
"kind": "scalar_to_arrow_batch",
|
||||
"version": 1,
|
||||
}
|
||||
assert request["required_secrets"] == ["API_TOKEN"]
|
||||
_assert_no_secret_values(request)
|
||||
|
||||
|
||||
def _run_packaged(definition, *args):
|
||||
"""Execute the shipped artifact in a fresh namespace, as a worker would."""
|
||||
source = base64.b64decode(definition.registration_request.artifact.content.data)
|
||||
namespace: dict = {}
|
||||
exec(compile(source, "<udf>", "exec"), namespace)
|
||||
return namespace[definition.registration_request.artifact.entrypoint](*args)
|
||||
|
||||
|
||||
def test_udf_packages_attribute_access_and_body_imports():
|
||||
@udf
|
||||
def word_norm(body: str) -> float:
|
||||
import numpy as np
|
||||
|
||||
try:
|
||||
words = body.split()
|
||||
except AttributeError as error:
|
||||
raise ValueError(str(error)) from error
|
||||
return float(np.linalg.norm([len(w) for w in words]))
|
||||
|
||||
assert _run_packaged(word_norm, "aa bb") == pytest.approx(8**0.5)
|
||||
|
||||
|
||||
def test_udf_packages_module_globals_and_global_caches():
|
||||
@udf
|
||||
def label(value: int) -> str:
|
||||
return "big" if value >= THRESHOLD else "small"
|
||||
|
||||
assert _run_packaged(label, 21) == "big"
|
||||
|
||||
@udf
|
||||
def cached(value: int) -> int:
|
||||
global _CACHE
|
||||
if _CACHE is None:
|
||||
_CACHE = 40
|
||||
return _CACHE + value
|
||||
|
||||
assert _run_packaged(cached, 2) == 42
|
||||
|
||||
|
||||
def test_udf_annotations_are_not_runtime_names():
|
||||
@udf
|
||||
def identity(value: date) -> date:
|
||||
return value
|
||||
|
||||
assert _run_packaged(identity, date(2026, 8, 25)) == date(2026, 8, 25)
|
||||
|
||||
|
||||
def test_udf_nested_scopes_resolve_lexically():
|
||||
@udf
|
||||
def score(value: int) -> int:
|
||||
offset = 2
|
||||
|
||||
def add_offset() -> int:
|
||||
return value + offset
|
||||
|
||||
return add_offset() + sum(v for v in [0])
|
||||
|
||||
assert _run_packaged(score, 3) == 5
|
||||
|
||||
|
||||
def test_udf_resolves_module_globals_before_builtins(tmp_path):
|
||||
module_path = tmp_path / "shadowing_udfs.py"
|
||||
module_path.write_text(
|
||||
"max = 7\n"
|
||||
"len = lambda _: 99\n"
|
||||
"\n"
|
||||
"def uses_literal_shadow(value: int) -> int:\n"
|
||||
" def nested() -> int:\n"
|
||||
" return max\n"
|
||||
" return nested() + value\n"
|
||||
"\n"
|
||||
"def uses_callable_shadow(value: int) -> int:\n"
|
||||
" def nested() -> int:\n"
|
||||
" return len([1])\n"
|
||||
" return nested() + value\n"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location("shadowing_udfs", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
|
||||
# The module's `max = 7` is what the interpreter would use, so it ships.
|
||||
assert _run_packaged(udf(module.uses_literal_shadow), 1) == 8
|
||||
# A callable global cannot ship; it must not be silently swapped for the builtin.
|
||||
with pytest.raises(TypeError, match="unsupported global value of type function"):
|
||||
udf(module.uses_callable_shadow)
|
||||
|
||||
|
||||
def test_canonical_arrow_type_is_exactly_the_grammar():
|
||||
from lancedb.functions import _GRAMMAR_PRIMITIVES, _canonical_arrow_type
|
||||
|
||||
golden = json.loads(
|
||||
(
|
||||
Path(__file__).parents[3]
|
||||
/ "rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json"
|
||||
).read_text()
|
||||
)
|
||||
primitives = [
|
||||
case["arrow_type"] for case in golden["valid"] if "<" not in case["arrow_type"]
|
||||
]
|
||||
assert [name for _, name in _GRAMMAR_PRIMITIVES] == primitives
|
||||
for outside in [
|
||||
pa.timestamp("us"),
|
||||
pa.decimal128(10, 2),
|
||||
pa.large_string(),
|
||||
pa.large_binary(),
|
||||
pa.binary(4),
|
||||
pa.duration("s"),
|
||||
pa.struct([pa.field("a", pa.int32())]),
|
||||
pa.list_(pa.float32(), 0),
|
||||
pa.list_(pa.timestamp("us")),
|
||||
]:
|
||||
with pytest.raises(TypeError, match="unsupported Arrow type"):
|
||||
_canonical_arrow_type(outside)
|
||||
|
||||
|
||||
def test_udf_nested_annotations_are_postponed_in_the_artifact():
|
||||
@udf
|
||||
def score(value: int) -> int:
|
||||
def identity(item: date) -> date:
|
||||
return item
|
||||
|
||||
identity(date(2026, 8, 25))
|
||||
return value
|
||||
|
||||
assert _run_packaged(score, 3) == 3
|
||||
|
||||
|
||||
def test_udf_ships_globals_the_body_deletes():
|
||||
@udf
|
||||
def clear(value: int) -> int:
|
||||
global _CACHE
|
||||
del _CACHE
|
||||
return value
|
||||
|
||||
assert _run_packaged(clear, 3) == 3
|
||||
|
||||
|
||||
def test_udf_rejects_a_module_global_that_does_not_import_as_itself(tmp_path):
|
||||
module_path = tmp_path / "fake_module_udfs.py"
|
||||
module_path.write_text(
|
||||
"import types\n"
|
||||
"np = types.ModuleType('numpy')\n"
|
||||
"np.sqrt = lambda x: 0\n"
|
||||
"\n"
|
||||
"def score(value: int) -> int:\n"
|
||||
" return int(np.sqrt(value))\n"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location("fake_module_udfs", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
with pytest.raises(TypeError, match="does not import as 'numpy'"):
|
||||
udf(module.score)
|
||||
|
||||
|
||||
def test_udf_rejects_a_module_level_namespace_alias(tmp_path):
|
||||
module_path = tmp_path / "aliasing_udfs.py"
|
||||
module_path.write_text(
|
||||
"import builtins as b\n"
|
||||
"THRESHOLD = 5\n"
|
||||
"\n"
|
||||
"def score(value: int) -> int:\n"
|
||||
" return value + b.vars(b.__import__('aliasing_udfs'))['THRESHOLD']\n"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location("aliasing_udfs", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
with pytest.raises(ValueError, match="dynamic namespace access"):
|
||||
udf(module.score)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"access",
|
||||
[
|
||||
"globals()['THRESHOLD']",
|
||||
"eval('THRESHOLD')",
|
||||
"(lambda g: g()['THRESHOLD'])(globals)",
|
||||
"__import__('sys').modules[__name__].THRESHOLD",
|
||||
"sys.modules[__name__].THRESHOLD",
|
||||
],
|
||||
)
|
||||
def test_udf_rejects_dynamic_namespace_access(access):
|
||||
namespace: dict = {}
|
||||
exec(
|
||||
f"def score(value: int) -> int:\n return value + {access}\n",
|
||||
{"THRESHOLD": 5},
|
||||
namespace,
|
||||
)
|
||||
with pytest.raises(ValueError, match="dynamic namespace access"):
|
||||
_package_from_text(
|
||||
"def score(value: int) -> int:\n"
|
||||
" import sys\n"
|
||||
f" return value + {access}\n"
|
||||
)
|
||||
|
||||
|
||||
def _package_from_text(source: str, module_globals: dict | None = None):
|
||||
"""Load `source` as a real module file so the packager can inspect it."""
|
||||
import tempfile
|
||||
|
||||
directory = tempfile.mkdtemp()
|
||||
path = Path(directory) / "generated_udf_module.py"
|
||||
path.write_text(source)
|
||||
spec = importlib.util.spec_from_file_location(f"generated_udf_{id(source)}", path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
if module_globals:
|
||||
module.__dict__.update(module_globals)
|
||||
spec.loader.exec_module(module)
|
||||
functions = [
|
||||
value
|
||||
for value in vars(module).values()
|
||||
if callable(value) and getattr(value, "__module__", None) == module.__name__
|
||||
]
|
||||
return udf(functions[0])
|
||||
|
||||
|
||||
def test_udf_rejects_a_non_standard_builtins_environment():
|
||||
def score(value: int) -> int:
|
||||
return len([1]) + value
|
||||
|
||||
score.__globals__ # noqa: B018 -- real function, real globals
|
||||
import builtins
|
||||
|
||||
patched = types.FunctionType(
|
||||
score.__code__,
|
||||
{"__builtins__": {**vars(builtins), "len": lambda _: 99}},
|
||||
"score",
|
||||
)
|
||||
patched.__annotations__ = score.__annotations__
|
||||
assert patched(3) == 102
|
||||
with pytest.raises(ValueError, match="non-standard builtins environment"):
|
||||
udf(patched)
|
||||
|
||||
class ReportingDict(dict): # reports standard entries, resolves differently
|
||||
def __missing__(self, key):
|
||||
return vars(builtins)[key]
|
||||
|
||||
disguised = types.FunctionType(
|
||||
score.__code__, {"__builtins__": ReportingDict(len=lambda _: 99)}, "score"
|
||||
)
|
||||
disguised.__annotations__ = score.__annotations__
|
||||
assert disguised(3) == 102
|
||||
with pytest.raises(ValueError, match="non-standard builtins environment"):
|
||||
udf(disguised)
|
||||
|
||||
hooked = types.FunctionType(
|
||||
score.__code__,
|
||||
{"__builtins__": {**vars(builtins), "__import__": lambda *a, **k: None}},
|
||||
"score",
|
||||
)
|
||||
hooked.__annotations__ = score.__annotations__
|
||||
with pytest.raises(ValueError, match="non-standard builtins environment"):
|
||||
udf(hooked)
|
||||
|
||||
|
||||
def test_udf_recursion_versus_a_rebound_module_name(tmp_path):
|
||||
module_path = tmp_path / "rebound_udfs.py"
|
||||
module_path.write_text(
|
||||
"def fact(value: int) -> int:\n"
|
||||
" return 1 if value <= 1 else value * fact(value - 1)\n"
|
||||
"\n"
|
||||
"def score(value: int) -> int:\n"
|
||||
" return score + value\n"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location("rebound_udfs", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
assert _run_packaged(udf(module.fact), 5) == 120
|
||||
raw = module.score
|
||||
module.score = 10
|
||||
with pytest.raises(ValueError, match="binds that name to another value"):
|
||||
udf(raw)
|
||||
# A wrapper that merely exposes __wrapped__ is not the function.
|
||||
module.score = functools.wraps(raw)(lambda value: 41)
|
||||
with pytest.raises(ValueError, match="binds that name to another value"):
|
||||
udf(raw)
|
||||
# The decorator's own result is; a subclass of it is not.
|
||||
module.fact = udf(module.fact)
|
||||
assert _run_packaged(module.fact, 4) == 24
|
||||
|
||||
class Twisted(UdfDefinition):
|
||||
def __call__(self, *args, **kwargs):
|
||||
return 41
|
||||
|
||||
raw_fact = module.fact._function
|
||||
module.fact = Twisted(
|
||||
raw_fact,
|
||||
name=None,
|
||||
input_schema=None,
|
||||
output_schema=None,
|
||||
pip=(),
|
||||
env={},
|
||||
python_version=None,
|
||||
)
|
||||
with pytest.raises(ValueError, match="binds that name to another value"):
|
||||
udf(raw_fact)
|
||||
|
||||
|
||||
def test_canonical_arrow_type_rejects_unrepresentable_list_children():
|
||||
from lancedb.functions import _canonical_arrow_type
|
||||
|
||||
for outside in [
|
||||
pa.list_(pa.float32()), # pyarrow default: nullable child
|
||||
pa.list_(pa.field("custom", pa.float32(), nullable=False)),
|
||||
pa.list_(pa.field("item", pa.float32(), nullable=False, metadata={"k": "v"})),
|
||||
pa.list_(pa.field("item", pa.float32(), nullable=False), 0),
|
||||
]:
|
||||
with pytest.raises(TypeError, match="unsupported Arrow type"):
|
||||
_canonical_arrow_type(outside)
|
||||
assert (
|
||||
_canonical_arrow_type(
|
||||
pa.list_(pa.field("item", pa.float32(), nullable=False), 3)
|
||||
)
|
||||
== "fixed_size_list<float32, 3>"
|
||||
)
|
||||
|
||||
|
||||
def _calls_missing(value: int) -> int:
|
||||
return missing(value) # noqa: F821
|
||||
|
||||
|
||||
def _shadows_missing_in_a_comprehension(value: int) -> int:
|
||||
return missing(value) + sum(missing for missing in ()) # noqa: F821
|
||||
|
||||
|
||||
def _shadows_missing_in_a_lambda(value: int) -> int:
|
||||
return (lambda missing: missing)(value) + missing # noqa: F821
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"function",
|
||||
[_calls_missing, _shadows_missing_in_a_comprehension, _shadows_missing_in_a_lambda],
|
||||
)
|
||||
def test_udf_rejects_a_truly_unresolved_global(function):
|
||||
with pytest.raises(ValueError, match=r"unresolved global names: \['missing'\]"):
|
||||
udf(function)
|
||||
|
||||
|
||||
def _arrow_type_from_golden(spec: dict) -> pa.DataType:
|
||||
kind = spec["type"]
|
||||
if kind in ("list", "large_list", "fixed_size_list"):
|
||||
item = _arrow_type_from_golden(spec["fields"][0]["type"])
|
||||
field = pa.field("item", item, nullable=False)
|
||||
if kind == "list":
|
||||
return pa.list_(field)
|
||||
if kind == "large_list":
|
||||
return pa.large_list(field)
|
||||
return pa.list_(field, spec["length"])
|
||||
return {
|
||||
"null": pa.null(),
|
||||
"bool": pa.bool_(),
|
||||
"utf8": pa.string(),
|
||||
"binary": pa.binary(),
|
||||
"float16": pa.float16(),
|
||||
"float32": pa.float32(),
|
||||
"float64": pa.float64(),
|
||||
"date32": pa.date32(),
|
||||
"date64": pa.date64(),
|
||||
}.get(kind) or getattr(pa, kind)()
|
||||
|
||||
|
||||
def test_arrow_type_grammar_matches_the_shared_golden():
|
||||
golden = json.loads(
|
||||
(
|
||||
Path(__file__).parents[3]
|
||||
/ "rust/lancedb/tests/fixtures/first_class_functions/v1/arrow_types.json"
|
||||
).read_text()
|
||||
)
|
||||
from lancedb.functions import _canonical_arrow_type
|
||||
|
||||
emitted = {
|
||||
case["arrow_type"]: _canonical_arrow_type(_arrow_type_from_golden(case["json"]))
|
||||
for case in golden["valid"]
|
||||
}
|
||||
assert emitted == {
|
||||
case["arrow_type"]: case["arrow_type"] for case in golden["valid"]
|
||||
}
|
||||
assert not set(emitted) & set(golden["invalid"])
|
||||
for case in golden["server_only"]:
|
||||
with pytest.raises(TypeError, match="unsupported Arrow type"):
|
||||
_canonical_arrow_type(_arrow_type_from_golden(case["json"]))
|
||||
|
||||
|
||||
def test_explicit_arrow_schema_is_deterministic():
|
||||
input_schema = pa.schema([pa.field("value", pa.float32(), nullable=True)])
|
||||
output_schema = pa.field("embedding", pa.list_(pa.float32(), 3), nullable=False)
|
||||
output_schema = pa.field(
|
||||
"embedding",
|
||||
pa.list_(pa.field("item", pa.float32(), nullable=False), 3),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
@udf(input_schema=input_schema, output_schema=output_schema)
|
||||
def explicit(value):
|
||||
@@ -82,7 +458,7 @@ def test_explicit_arrow_schema_is_deterministic():
|
||||
signature = explicit.registration_request.signature
|
||||
assert signature.inputs[0].arrow_type == "float32"
|
||||
assert signature.inputs[0].nullable is True
|
||||
assert signature.output.arrow_type == "fixed_size_list<float32>[3]"
|
||||
assert signature.output.arrow_type == "fixed_size_list<float32, 3>"
|
||||
assert signature.output.nullable is False
|
||||
|
||||
|
||||
@@ -130,14 +506,6 @@ def test_annotation_and_explicit_schema_validation_fail_closed():
|
||||
return value
|
||||
|
||||
|
||||
def test_environment_rejects_secret_value_overlap():
|
||||
with pytest.raises(ValueError, match="must be disjoint"):
|
||||
|
||||
@udf(env={"TOKEN": "plaintext"}, secrets=["TOKEN"])
|
||||
def overlapping(value: int) -> int:
|
||||
return value
|
||||
|
||||
|
||||
def test_local_function_catalog_operations_are_not_supported(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
message = "Function catalog operations are not supported by this database"
|
||||
@@ -162,7 +530,7 @@ def _mock_remote_function_catalog():
|
||||
body = json.loads(self.rfile.read(length) or b"{}")
|
||||
state["requests"].append((self.path, body))
|
||||
status = 200
|
||||
if self.path == "/v1/function/create":
|
||||
if self.path == "/v1/functions/create":
|
||||
state["version"] = {
|
||||
"name": body["name"],
|
||||
"version": "fv_exact",
|
||||
@@ -174,7 +542,6 @@ def _mock_remote_function_catalog():
|
||||
"runtime": body["runtime"],
|
||||
"runtime_digest": "sha256:runtime",
|
||||
"environment_digest": "sha256:environment",
|
||||
"required_secrets": body.get("required_secrets", []),
|
||||
"created_at": "2026-08-21T00:00:00Z",
|
||||
}
|
||||
response = {"job_id": "job-register"}
|
||||
@@ -187,7 +554,7 @@ def _mock_remote_function_catalog():
|
||||
"job_state": "DONE",
|
||||
"result": state["version"],
|
||||
}
|
||||
elif self.path == "/v1/function/describe":
|
||||
elif self.path == "/v1/functions/describe":
|
||||
assert body == {
|
||||
"name": "normalize_score",
|
||||
"version": "fv_exact",
|
||||
@@ -233,7 +600,6 @@ def test_remote_registration_job_and_exact_version_reopen_round_trip():
|
||||
assert create_request == json.loads(
|
||||
normalize_score.registration_request.to_canonical_json()
|
||||
)
|
||||
_assert_no_secret_values(create_request)
|
||||
|
||||
|
||||
def test_blocking_remote_registration_returns_function_version():
|
||||
@@ -249,6 +615,6 @@ def test_blocking_remote_registration_returns_function_version():
|
||||
assert created.name == "normalize_score"
|
||||
assert created.version == "fv_exact"
|
||||
assert [path for path, _ in state["requests"]] == [
|
||||
"/v1/function/create",
|
||||
"/v1/functions/create",
|
||||
"/v1/jobs/describe",
|
||||
]
|
||||
|
||||
@@ -245,6 +245,14 @@ def test_create_inverted_index_rejects_invalid_block_size(table):
|
||||
table.create_index("text", config=FTS(block_size=129))
|
||||
|
||||
|
||||
def test_create_inverted_index_respects_build_memory_limit(table):
|
||||
with pytest.raises(ValueError, match="exceeds worker memory limit"):
|
||||
table.create_index(
|
||||
"text",
|
||||
config=FTS(memory_limit=0, num_workers=1),
|
||||
)
|
||||
|
||||
|
||||
def test_custom_stop_words_list(table):
|
||||
table.create_index(
|
||||
"text",
|
||||
|
||||
@@ -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
|
||||
@@ -8,6 +8,11 @@ import pytest
|
||||
from lancedb import DBConnection, Table, connect
|
||||
from lancedb.background_loop import LOOP
|
||||
from lancedb.permutation import Permutation, Permutations, permutation_builder
|
||||
from utils import (
|
||||
MockPermutationServer,
|
||||
assert_server_safe_row_id_requests,
|
||||
mock_remote_table,
|
||||
)
|
||||
|
||||
|
||||
def test_split_random_ratios(mem_db):
|
||||
@@ -51,6 +56,31 @@ def test_execute_does_not_reenter_background_loop(tmp_path, monkeypatch):
|
||||
assert permutation_tbl._conn.read_consistency_interval is None
|
||||
|
||||
|
||||
def test_pickled_permutation_reads_pinned_version(tmp_path):
|
||||
"""An unpickled copy must still read the pinned version, which also covers the
|
||||
version surviving the ``to_arrow()`` round trip in ``__getstate__``."""
|
||||
import pickle
|
||||
|
||||
db = connect(tmp_path)
|
||||
tbl = db.create_table("base", pa.table({"idx": range(20)}))
|
||||
permutation_tbl = permutation_builder(tbl).execute()
|
||||
perm = Permutation.from_tables(tbl, permutation_tbl)
|
||||
|
||||
payload = pickle.dumps(perm)
|
||||
|
||||
# Compact so the stored row addresses no longer describe these rows at latest.
|
||||
tbl.delete("true")
|
||||
tbl.optimize()
|
||||
assert tbl.count_rows() == 0
|
||||
|
||||
# Unpickle after the mutation: __setstate__ reopens at latest, so this only
|
||||
# passes if the recorded version is applied on reopen.
|
||||
restored = pickle.loads(payload)
|
||||
assert len(restored) == 20
|
||||
rows = restored.__getitems__(list(range(20)))
|
||||
assert sorted(row["idx"] for row in rows) == list(range(20))
|
||||
|
||||
|
||||
def test_split_random_counts(mem_db):
|
||||
"""Test random splitting with absolute counts."""
|
||||
tbl = mem_db.create_table(
|
||||
@@ -1214,3 +1244,57 @@ def test_remove_rowid_after_select(some_permutation: Permutation):
|
||||
perm_without_rowid = perm_with_rowid.remove_columns(["_rowid"])
|
||||
assert "_rowid" not in perm_without_rowid.column_names
|
||||
assert perm_without_rowid.column_names == ["id"]
|
||||
|
||||
|
||||
def test_permutation_is_stable_when_remote_scan_order_varies():
|
||||
"""Splits are assigned by scan position, and every rank builds its own
|
||||
permutation, so two ranks seeing different scan orders must still agree."""
|
||||
server = MockPermutationServer(num_rows=16, vary_scan_order=True)
|
||||
|
||||
def split_of_each_row(permutation_tbl):
|
||||
# Sequential splits are assigned by position, so a reversed scan would put
|
||||
# the last rows in split 0. Compare the mapping rather than the table order,
|
||||
# which the split-id sort does not pin down.
|
||||
rows = permutation_tbl.search(None).to_arrow().to_pydict()
|
||||
return dict(zip(rows["row_id"], rows["split_id"]))
|
||||
|
||||
with mock_remote_table(server) as table:
|
||||
first = split_of_each_row(
|
||||
permutation_builder(table).split_sequential(fixed=2).execute()
|
||||
)
|
||||
second = split_of_each_row(
|
||||
permutation_builder(table).split_sequential(fixed=2).execute()
|
||||
)
|
||||
|
||||
assert server.scan_calls == 2, "both builds must have scanned"
|
||||
assert first == second
|
||||
assert first[0] == 0 and first[server.num_rows - 1] == 1, first
|
||||
|
||||
|
||||
def test_permutation_over_remote_table():
|
||||
"""The permutation API accepts a remote table, addressing rows by `_rowid` just
|
||||
as `take_row_ids` does. Also pins the request shapes sent to the server.
|
||||
"""
|
||||
server = MockPermutationServer()
|
||||
|
||||
with mock_remote_table(server) as table:
|
||||
permutation_tbl = permutation_builder(table).split_sequential(fixed=2).execute()
|
||||
assert permutation_tbl.count_rows() == server.num_rows
|
||||
|
||||
permutation = Permutation.from_tables(table, permutation_tbl, 0)
|
||||
assert permutation.num_rows == server.num_rows // 2
|
||||
|
||||
# Compare against the permutation's own order; the split-id sort is not stable.
|
||||
rows = permutation_tbl.search(None).to_arrow().to_pydict()
|
||||
split0 = [
|
||||
row_id
|
||||
for row_id, split in zip(rows["row_id"], rows["split_id"])
|
||||
if not split
|
||||
]
|
||||
# The mock table's `id` equals its `_rowid`.
|
||||
assert permutation.take_offsets([2, 0]) == [
|
||||
{"id": split0[2]},
|
||||
{"id": split0[0]},
|
||||
]
|
||||
|
||||
assert_server_safe_row_id_requests(server)
|
||||
|
||||
@@ -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},
|
||||
]
|
||||
@@ -876,11 +875,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
|
||||
job.wait()
|
||||
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")
|
||||
await job.wait()
|
||||
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]
|
||||
|
||||
@@ -1,7 +1,17 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
import contextlib
|
||||
import http.server
|
||||
import json
|
||||
import re
|
||||
import threading
|
||||
|
||||
import lancedb
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
ARROW_FILE_CONTENT_TYPE = "application/vnd.apache.arrow.file"
|
||||
|
||||
|
||||
def exception_output(e_info: pytest.ExceptionInfo):
|
||||
import traceback
|
||||
@@ -9,3 +19,199 @@ def exception_output(e_info: pytest.ExceptionInfo):
|
||||
# skip traceback part, since it's not worth checking in tests
|
||||
lines = traceback.format_exception_only(e_info.type, e_info.value)
|
||||
return "".join(lines).strip()
|
||||
|
||||
|
||||
def parse_in_list(filter_sql: str) -> list[int]:
|
||||
"""Pull the integers out of a `<col> IN (a, b, c)` predicate.
|
||||
|
||||
Scoped to the parenthesised list so a cast in the SQL adds no phantom values.
|
||||
"""
|
||||
match = re.search(r"\bIN\s*\(([^)]*)\)", filter_sql, re.IGNORECASE)
|
||||
assert match is not None, f"expected an IN list, got: {filter_sql}"
|
||||
return [int(m) for m in re.findall(r"-?\d+", match.group(1))]
|
||||
|
||||
|
||||
def is_row_id_take(body) -> bool:
|
||||
"""True when a query body fetches specific rows by row id."""
|
||||
return "_rowid" in (body.get("filter") or "")
|
||||
|
||||
|
||||
def arrow_file_bytes(table: pa.Table) -> bytes:
|
||||
"""Serialize to the Arrow IPC *file* framing the /query/ route answers with."""
|
||||
sink = pa.BufferOutputStream()
|
||||
with pa.ipc.new_file(sink, table.schema) as writer:
|
||||
writer.write_table(table)
|
||||
return sink.getvalue().to_pybytes()
|
||||
|
||||
|
||||
class MockPermutationServer:
|
||||
"""A stand-in LanceDB server hosting one table whose ``id`` equals its ``_rowid``.
|
||||
|
||||
Records every ``/query/`` body so tests can assert on the request shapes sent to
|
||||
the server, which is the part that has to stay compatible.
|
||||
"""
|
||||
|
||||
def __init__(self, name="remote_data", num_rows=8, vary_scan_order=False):
|
||||
self.name = name
|
||||
self.num_rows = num_rows
|
||||
self.query_bodies = []
|
||||
# Stand in for a distributed scan that answers in no fixed order.
|
||||
self.vary_scan_order = vary_scan_order
|
||||
self.scan_calls = 0
|
||||
|
||||
def __call__(self, request):
|
||||
path = request.path
|
||||
if path == f"/v1/table/{self.name}/describe/":
|
||||
return self._json(
|
||||
request,
|
||||
{
|
||||
"version": 1,
|
||||
"schema": {
|
||||
"fields": [
|
||||
{"name": "id", "type": {"type": "int64"}, "nullable": False}
|
||||
]
|
||||
},
|
||||
},
|
||||
)
|
||||
if path == f"/v1/table/{self.name}/get_lsm_write_spec/":
|
||||
self._read_body(request)
|
||||
# Null spec: this table has no LSM write path.
|
||||
return self._json(request, {"lsm_write_spec": None})
|
||||
if path == f"/v1/table/{self.name}/count_rows/":
|
||||
self._read_body(request)
|
||||
return self._json(request, self.num_rows)
|
||||
if path == f"/v1/table/{self.name}/query/":
|
||||
return self._query(request, self._read_body(request))
|
||||
|
||||
# Drain first, so an unexpected route cannot desync a keep-alive connection.
|
||||
self._read_body(request)
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
@property
|
||||
def scans(self):
|
||||
"""Bodies of the permutation build scan: the row id column, nothing else."""
|
||||
return [b for b in self.query_bodies if b.get("columns") == ["_rowid"]]
|
||||
|
||||
@property
|
||||
def takes(self):
|
||||
"""Bodies of the row-id takes the loader fetches batches with.
|
||||
|
||||
Keyed on `_rowid`, not "has a filter": the schema probe also has a predicate.
|
||||
"""
|
||||
return [b for b in self.query_bodies if is_row_id_take(b)]
|
||||
|
||||
@staticmethod
|
||||
def _read_body(request):
|
||||
content_len = int(request.headers.get("Content-Length") or 0)
|
||||
return json.loads(request.rfile.read(content_len)) if content_len else {}
|
||||
|
||||
@staticmethod
|
||||
def _json(request, payload):
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(json.dumps(payload).encode())
|
||||
|
||||
@staticmethod
|
||||
def _arrow(request, table):
|
||||
body = arrow_file_bytes(table)
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", ARROW_FILE_CONTENT_TYPE)
|
||||
request.send_header("Content-Length", str(len(body)))
|
||||
request.end_headers()
|
||||
request.wfile.write(body)
|
||||
|
||||
def _query(self, request, body):
|
||||
self.query_bodies.append(body)
|
||||
|
||||
if is_row_id_take(body):
|
||||
# A row-id take. Answer ascending, so tests prove the client reorders.
|
||||
row_ids = sorted(parse_in_list(body["filter"]))
|
||||
return self._arrow(
|
||||
request,
|
||||
pa.table(
|
||||
{
|
||||
"id": pa.array(row_ids, pa.int64()),
|
||||
"_rowid": pa.array(row_ids, pa.uint64()),
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
if body.get("columns") == ["_rowid"]:
|
||||
# The permutation build scan: row ids and nothing else.
|
||||
row_ids = list(range(self.num_rows))
|
||||
if self.vary_scan_order and self.scan_calls % 2:
|
||||
row_ids.reverse()
|
||||
self.scan_calls += 1
|
||||
return self._arrow(
|
||||
request,
|
||||
pa.table({"_rowid": pa.array(row_ids, pa.uint64())}),
|
||||
)
|
||||
|
||||
# The schema probe: filtered to nothing, so it carries schema and no rows.
|
||||
return self._arrow(request, pa.table({"id": pa.array([], pa.int64())}))
|
||||
|
||||
|
||||
def _make_handler(serve):
|
||||
class MockLanceDBHandler(http.server.BaseHTTPRequestHandler):
|
||||
def do_GET(self):
|
||||
serve(self)
|
||||
|
||||
def do_POST(self):
|
||||
serve(self)
|
||||
|
||||
def log_message(self, *args):
|
||||
pass # keep pytest output readable
|
||||
|
||||
return MockLanceDBHandler
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def mock_remote_table(server):
|
||||
"""Run ``server`` on a local port and yield an open remote table against it.
|
||||
|
||||
Threading: the loader fans out fetch threads a single-threaded server would
|
||||
serialize, hiding the prefetch overlap under test.
|
||||
"""
|
||||
with http.server.ThreadingHTTPServer(
|
||||
("localhost", 0), _make_handler(server)
|
||||
) as srv:
|
||||
thread = threading.Thread(target=srv.serve_forever)
|
||||
thread.start()
|
||||
try:
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=f"http://localhost:{srv.server_address[1]}",
|
||||
client_config={"timeout_config": {"connect_timeout": 5}},
|
||||
)
|
||||
yield db.open_table(server.name)
|
||||
finally:
|
||||
srv.shutdown()
|
||||
thread.join()
|
||||
|
||||
|
||||
def assert_server_safe_row_id_requests(server):
|
||||
"""Assert the loader fetched rows by row id and bounded everything else.
|
||||
|
||||
`.get`, not `[...]`, so a dropped field reads as the assertion, not a KeyError.
|
||||
"""
|
||||
for body in server.takes:
|
||||
# The fetch needs the row id back to restore the requested order.
|
||||
assert body.get("with_row_id") is True, body
|
||||
assert "_rowid" in body["filter"], body
|
||||
|
||||
# Only the one-off permutation scan may scan the whole table; the schema probe is
|
||||
# built once per split per epoch. `k == 0` counts as unbounded: lance reads a zero
|
||||
# limit as "no limit".
|
||||
def is_unbounded(body):
|
||||
if is_row_id_take(body):
|
||||
return False
|
||||
k = body.get("k")
|
||||
return k is None or k == 0 or k > server.num_rows
|
||||
|
||||
unbounded = [b for b in server.query_bodies if is_unbounded(b)]
|
||||
assert unbounded == server.scans, (
|
||||
f"only the permutation scan may be unbounded, got {unbounded}"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
+57
-1
@@ -42,7 +42,7 @@ pub fn extract_index_params(source: &Option<Bound<'_, PyAny>>) -> PyResult<Lance
|
||||
"Fm" => Ok(LanceDbIndex::Fm(FmIndexBuilder::default())),
|
||||
"FTS" => {
|
||||
let params = source.extract::<FtsParams>()?;
|
||||
let inner_opts = FtsIndexBuilder::default()
|
||||
let mut inner_opts = FtsIndexBuilder::default()
|
||||
.base_tokenizer(params.base_tokenizer)
|
||||
.language(¶ms.language)
|
||||
.map_err(|_| {
|
||||
@@ -61,6 +61,12 @@ pub fn extract_index_params(source: &Option<Bound<'_, PyAny>>) -> PyResult<Lance
|
||||
.ngram_max_length(params.ngram_max_length)
|
||||
.ngram_prefix_only(params.prefix_only)
|
||||
.custom_stop_words(params.custom_stop_words);
|
||||
if let Some(memory_limit) = params.memory_limit {
|
||||
inner_opts = inner_opts.memory_limit_mb(memory_limit);
|
||||
}
|
||||
if let Some(num_workers) = params.num_workers {
|
||||
inner_opts = inner_opts.num_workers(num_workers);
|
||||
}
|
||||
let inner_opts = inner_opts
|
||||
.block_size(params.block_size)
|
||||
.map_err(|err| PyValueError::new_err(err.to_string()))?;
|
||||
@@ -213,6 +219,8 @@ struct FtsParams {
|
||||
ngram_max_length: u32,
|
||||
prefix_only: bool,
|
||||
block_size: usize,
|
||||
memory_limit: Option<u64>,
|
||||
num_workers: Option<usize>,
|
||||
}
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
@@ -444,3 +452,51 @@ impl IndexConfig {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use pyo3::types::{PyDict, PyDictMethods};
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn fts_build_controls_are_forwarded() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"class FTS:
|
||||
with_position = True
|
||||
base_tokenizer = 'simple'
|
||||
language = 'English'
|
||||
max_token_length = None
|
||||
lower_case = True
|
||||
stem = False
|
||||
remove_stop_words = False
|
||||
custom_stop_words = None
|
||||
ascii_folding = False
|
||||
ngram_min_length = 3
|
||||
ngram_max_length = 3
|
||||
prefix_only = False
|
||||
block_size = 128
|
||||
memory_limit = 2048
|
||||
num_workers = 7
|
||||
|
||||
config = FTS()",
|
||||
None,
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
let config = locals.get_item("config").unwrap().unwrap();
|
||||
let index = extract_index_params(&Some(config)).unwrap();
|
||||
let LanceDbIndex::FTS(params) = index else {
|
||||
panic!("expected FTS index parameters");
|
||||
};
|
||||
let training_json = params.to_training_json().unwrap();
|
||||
|
||||
assert_eq!(training_json.get("memory_limit"), Some(&json!(2048)));
|
||||
assert_eq!(training_json.get("num_workers"), Some(&json!(7)));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
+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(())
|
||||
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>()?;
|
||||
|
||||
@@ -268,7 +268,9 @@ impl PyPermutationReader {
|
||||
.await
|
||||
.infer_error()?
|
||||
} else {
|
||||
PermutationReader::identity(base_table).await
|
||||
PermutationReader::identity(base_table)
|
||||
.await
|
||||
.infer_error()?
|
||||
};
|
||||
Ok(Self::from_reader(reader))
|
||||
})
|
||||
|
||||
+64
-5
@@ -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 {
|
||||
@@ -745,15 +780,19 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
#[pyo3(signature = (data, mode, progress=None, write_parallelism=None))]
|
||||
#[pyo3(signature = (data, mode, progress=None, write_parallelism=None, allow_external_blob_outside_bases=false))]
|
||||
pub fn add<'a>(
|
||||
self_: PyRef<'a, Self>,
|
||||
data: PyScannable,
|
||||
mode: String,
|
||||
progress: Option<Py<PyAny>>,
|
||||
write_parallelism: Option<usize>,
|
||||
allow_external_blob_outside_bases: bool,
|
||||
) -> PyResult<Bound<'a, PyAny>> {
|
||||
let mut op = self_.inner_ref()?.add(data);
|
||||
let mut op = self_
|
||||
.inner_ref()?
|
||||
.add(data)
|
||||
.allow_external_blob_outside_bases(allow_external_blob_outside_bases);
|
||||
if mode == "append" {
|
||||
op = op.mode(AddDataMode::Append);
|
||||
} else if mode == "overwrite" {
|
||||
@@ -1584,7 +1623,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 +1944,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 +1952,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