mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-04 04:28:44 +00:00
feat: add_bases registers extra table storage prefixes
TableBase plus add_bases on native, memory, namespace, and Cloud clients.
This commit is contained in:
@@ -21,7 +21,7 @@ from .remote.db import RemoteDBConnection
|
||||
from .expr import Expr, col, lit, func
|
||||
from .schema import blob, vector, BlobType
|
||||
from .job import AsyncJob, Job
|
||||
from .table import AsyncTable, Table
|
||||
from .table import AsyncTable, Table, TableBase
|
||||
from .types import BaseTokenizerType
|
||||
from ._lancedb import Session
|
||||
from .namespace import (
|
||||
@@ -521,5 +521,6 @@ __all__ = [
|
||||
"RemoteDBConnection",
|
||||
"Session",
|
||||
"Table",
|
||||
"TableBase",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
@@ -377,6 +377,7 @@ class Table:
|
||||
def take_offsets(self, offsets: list[int]) -> TakeQuery: ...
|
||||
def take_row_ids(self, row_ids: list[int]) -> TakeQuery: ...
|
||||
async def blob_columns(self) -> list[str]: ...
|
||||
async def add_bases(self, bases: list[Any]) -> None: ...
|
||||
async def fetch_blobs(
|
||||
self, column: str, row_ids: list[int]
|
||||
) -> pa.LargeBinaryArray: ...
|
||||
|
||||
@@ -50,7 +50,7 @@ from lancedb.index import (
|
||||
)
|
||||
from lancedb.job import Job
|
||||
from lancedb.remote.db import LOOP
|
||||
from lancedb.table import IndexConfigType, KNOWN_METRICS
|
||||
from lancedb.table import IndexConfigType, KNOWN_METRICS, TableBase
|
||||
import pyarrow as pa
|
||||
|
||||
from lancedb.common import DATA, VEC, VECTOR_COLUMN_NAME
|
||||
@@ -1082,6 +1082,13 @@ class RemoteTable(Table):
|
||||
def blob_columns(self) -> list[str]:
|
||||
return LOOP.run(self._table.blob_columns())
|
||||
|
||||
def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
"""Register additional storage bases for this table."""
|
||||
LOOP.run(self._table.add_bases(bases))
|
||||
|
||||
def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> pa.LargeBinaryArray:
|
||||
|
||||
@@ -19,6 +19,7 @@ from typing import (
|
||||
Iterable,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
@@ -710,6 +711,21 @@ def _normalize_progress(progress):
|
||||
return progress, False
|
||||
|
||||
|
||||
@dataclass
|
||||
class TableBase:
|
||||
"""An extra storage prefix registered on a table.
|
||||
|
||||
``path`` is an object-store URI. ``name`` is an optional alias.
|
||||
``is_dataset_root`` is true when ``path`` points to a Lance dataset
|
||||
root. When false, ``path`` points directly to the directory containing
|
||||
the referenced files.
|
||||
"""
|
||||
|
||||
path: str
|
||||
name: Optional[str] = None
|
||||
is_dataset_root: bool = False
|
||||
|
||||
|
||||
class Table(ABC):
|
||||
"""
|
||||
A Table is a collection of Records in a LanceDB Database.
|
||||
@@ -1568,6 +1584,18 @@ class Table(ABC):
|
||||
def blob_columns(self) -> list[str]:
|
||||
"""Names of the blob v2 columns declared on this table."""
|
||||
|
||||
def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
"""Register additional storage bases for this table.
|
||||
|
||||
A URI string is a non-root base with no alias::
|
||||
|
||||
table.add_bases("s3://bucket/media/")
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
@@ -2414,6 +2442,12 @@ class LanceTable(Table):
|
||||
def blob_columns(self) -> list[str]:
|
||||
return LOOP.run(self._table.blob_columns())
|
||||
|
||||
def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
LOOP.run(self._table.add_bases(bases))
|
||||
|
||||
def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> pa.LargeBinaryArray:
|
||||
@@ -6266,6 +6300,18 @@ class AsyncTable:
|
||||
async def blob_columns(self) -> list[str]:
|
||||
return await self._inner.blob_columns()
|
||||
|
||||
async def add_bases(
|
||||
self,
|
||||
bases: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> None:
|
||||
"""Register additional storage bases for this table.
|
||||
|
||||
A URI string is a non-root base with no alias::
|
||||
|
||||
await table.add_bases("s3://bucket/media/")
|
||||
"""
|
||||
await self._inner.add_bases(_normalize_bases(bases))
|
||||
|
||||
async def fetch_blobs(
|
||||
self, column: str, row_ids: Union[list[int], pa.Table]
|
||||
) -> pa.LargeBinaryArray:
|
||||
@@ -6484,6 +6530,30 @@ class AsyncTable:
|
||||
await self._inner.replace_field_metadata(field_name, new_metadata)
|
||||
|
||||
|
||||
def _normalize_bases(
|
||||
base_inputs: Union[str, TableBase, Iterable[Union[str, TableBase]]],
|
||||
) -> list[TableBase]:
|
||||
if isinstance(base_inputs, (str, TableBase)):
|
||||
items: Iterable[Union[str, TableBase]] = [base_inputs]
|
||||
elif isinstance(base_inputs, Mapping):
|
||||
raise TypeError(
|
||||
"Expected a URI string, TableBase, or an iterable of those values"
|
||||
)
|
||||
else:
|
||||
items = base_inputs
|
||||
normalized_bases: list[TableBase] = []
|
||||
for base in items:
|
||||
if isinstance(base, str):
|
||||
normalized_bases.append(TableBase(path=base))
|
||||
elif isinstance(base, TableBase):
|
||||
normalized_bases.append(base)
|
||||
else:
|
||||
raise TypeError(
|
||||
f"Expected a URI string or TableBase, got {type(base).__name__}"
|
||||
)
|
||||
return normalized_bases
|
||||
|
||||
|
||||
@dataclass
|
||||
class IndexStatistics:
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
|
||||
|
||||
def test_add_bases_accepts_named_and_dataset_root(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
parent = tmp_path / "parent"
|
||||
media.mkdir()
|
||||
parent.mkdir()
|
||||
db = lancedb.connect(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases(
|
||||
[
|
||||
lancedb.TableBase(path=media.as_uri(), name="media", is_dataset_root=False),
|
||||
lancedb.TableBase(
|
||||
path=parent.as_uri(), name="parent", is_dataset_root=True
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_add_bases_accepts_two_unnamed_paths(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
other = tmp_path / "other"
|
||||
media.mkdir()
|
||||
other.mkdir()
|
||||
db = lancedb.connect(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases([media.as_uri(), other.as_uri()])
|
||||
|
||||
|
||||
def test_add_bases_rejects_dict_input(tmp_path):
|
||||
db = lancedb.connect(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
with pytest.raises(TypeError, match="TableBase"):
|
||||
table.add_bases({"path": "s3://bucket/media/"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_add_bases_accepts_file_uri(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
media.mkdir()
|
||||
db = await lancedb.connect_async(tmp_path / "db")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = await db.create_table("photos", schema=schema)
|
||||
await table.add_bases(media.as_uri())
|
||||
|
||||
|
||||
def test_memory_add_bases_accepts_file_uri(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
media.mkdir()
|
||||
db = lancedb.connect("memory:///")
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases(media.as_uri())
|
||||
|
||||
|
||||
def test_namespace_add_bases_accepts_file_uri(tmp_path):
|
||||
media = tmp_path / "media"
|
||||
media.mkdir()
|
||||
db = lancedb.connect_namespace("dir", {"root": str(tmp_path / "ns")})
|
||||
schema = pa.schema([pa.field("id", pa.int64())])
|
||||
table = db.create_table("photos", schema=schema)
|
||||
table.add_bases(media.as_uri())
|
||||
@@ -2306,3 +2306,36 @@ def test_remote_connection_jobs_surface():
|
||||
assert job.status() == "failed"
|
||||
with pytest.raises(JobFailedError, match="worker died"):
|
||||
job.wait(timeout=timedelta(seconds=5))
|
||||
|
||||
|
||||
def test_remote_add_bases_posts_the_bases_array():
|
||||
captured_body = {}
|
||||
|
||||
def handler(request):
|
||||
if 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(BLOB_DESCRIBE_RESPONSE).encode())
|
||||
elif request.path == "/v1/table/test/bases/":
|
||||
content_len = int(request.headers.get("Content-Length", 0))
|
||||
captured_body.update(json.loads(request.rfile.read(content_len)))
|
||||
request.send_response(200)
|
||||
request.send_header("Content-Type", "application/json")
|
||||
request.end_headers()
|
||||
request.wfile.write(b'{"version": 2}')
|
||||
else:
|
||||
request.send_response(404)
|
||||
request.end_headers()
|
||||
|
||||
with mock_lancedb_connection(handler) as db:
|
||||
table = db.open_table("test")
|
||||
table.add_bases(lancedb.TableBase(path="s3://bucket/media/"))
|
||||
|
||||
assert captured_body["bases"] == [
|
||||
{
|
||||
"path": "s3://bucket/media/",
|
||||
"isDatasetRoot": False,
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@@ -22,6 +22,7 @@ use lancedb::index::scalar::FtsIndexBuilder;
|
||||
use lancedb::table::{
|
||||
AddDataMode, ColumnAlteration, Duration, FieldMetadataUpdate, FtsToken as LanceDbFtsToken,
|
||||
NewColumnTransform, OptimizeAction, OptimizeOptions, Ref, Table as LanceDbTable,
|
||||
TableBase as LanceTableBase,
|
||||
};
|
||||
use lancedb::tokenize as lancedb_tokenize;
|
||||
use pyo3::{
|
||||
@@ -94,6 +95,13 @@ fn lsm_stats_to_py(py: Python<'_>, stats: &lancedb::table::LsmStats) -> PyResult
|
||||
Ok(out.unbind())
|
||||
}
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
pub(crate) struct PyTableBase {
|
||||
path: String,
|
||||
name: Option<String>,
|
||||
is_dataset_root: bool,
|
||||
}
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
enum PredicateArg {
|
||||
Expr(PyExpr),
|
||||
@@ -1238,6 +1246,25 @@ impl Table {
|
||||
})
|
||||
}
|
||||
|
||||
#[pyo3(signature = (bases))]
|
||||
pub fn add_bases(
|
||||
self_: PyRef<'_, Self>,
|
||||
bases: Vec<PyTableBase>,
|
||||
) -> PyResult<Bound<'_, PyAny>> {
|
||||
let inner = self_.inner_ref()?.clone();
|
||||
let bases: Vec<LanceTableBase> = bases
|
||||
.into_iter()
|
||||
.map(|base| LanceTableBase {
|
||||
path: base.path,
|
||||
name: base.name,
|
||||
is_dataset_root: base.is_dataset_root,
|
||||
})
|
||||
.collect();
|
||||
future_into_py(self_.py(), async move {
|
||||
inner.add_bases(bases).await.infer_error()
|
||||
})
|
||||
}
|
||||
|
||||
/// Read blob bytes for `row_ids` from blob v2 column `column`.
|
||||
#[pyo3(signature = (column, row_ids))]
|
||||
pub fn fetch_blobs(
|
||||
|
||||
Reference in New Issue
Block a user