Compare commits

..

2 Commits

Author SHA1 Message Date
Gatefixer 88d8a69a99 fix(python): scope Instructor compatibility shim 2026-08-06 01:30:35 +00:00
Gatefixer bd779bb7d5 fix(python): support legacy InstructorEmbedding downloads 2026-08-06 01:12:01 +00:00
30 changed files with 202 additions and 1853 deletions
-29
View File
@@ -197,35 +197,6 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
expect(table.getChild("d")?.toJSON()).toEqual([9n, 10n, null]);
});
it("will use a provided FixedSizeList schema with typed array values", function () {
const schema = new Schema([
new Field("text", new Utf8(), false),
new Field(
"vector",
new FixedSizeList(3, new Field("item", new Float32(), false)),
false,
),
]);
const table = makeArrowTable(
[
{
text: "foo",
vector: new Float32Array([1, 2, 3]),
},
],
{ schema },
);
expect(table.getChild("text")?.toJSON()).toEqual(["foo"]);
expect(
table
.getChild("vector")
?.toJSON()
.map((value) => value.toJSON()),
).toEqual([[1, 2, 3]]);
});
it("will assume the column `vector` is FixedSizeList<Float32> by default", async function () {
const schema = new Schema([
new Field("a", new Float(Precision.DOUBLE), true),
-32
View File
@@ -170,38 +170,6 @@ describe("remote connection", () => {
);
});
it("surfaces JSON server errors from remote table operations", async () => {
await withMockDatabase(
(req, res) => {
const path = req.url ?? "";
if (path.endsWith("/describe/")) {
res.writeHead(200, { "Content-Type": "application/json" }).end(
JSON.stringify({
name: "broken_table",
version: 1,
schema: { fields: [] },
}),
);
return;
}
if (path.endsWith("/count_rows/")) {
res
.writeHead(400, { "Content-Type": "application/json" })
.end(JSON.stringify({ error: "count rows failed" }));
return;
}
res.writeHead(404).end();
},
async (db) => {
const table = await db.openTable("broken_table");
await expect(table.countRows()).rejects.toThrow("count rows failed");
},
);
});
it("should pass on requested extra headers", async () => {
await withMockDatabase(
(req, res) => {
-38
View File
@@ -86,44 +86,6 @@ describe.each([arrow15, arrow16, arrow17, arrow18])(
await expect(table.countRows()).resolves.toBe(3);
});
it("should support a foreign Float64 vector schema end to end", async () => {
const conn = await connect(tmpDir.name);
const schema = new arrow.Schema([
new arrow.Field("resource_id", new arrow.Int32(), false),
new arrow.Field(
"vector",
new arrow.FixedSizeList(
3,
new arrow.Field("value", new arrow.Float64(), true),
),
false,
),
]);
const data = [
{
// biome-ignore lint/style/useNamingConvention: matches the reported schema
resource_id: 0,
vector: [0.1, 0.1, 0.1],
},
];
const resources = await conn.createTable("resources", data, { schema });
const existing = await resources
.query()
.where("resource_id = 0")
.limit(1)
.toArray();
expect(existing).toHaveLength(1);
const matched = await resources
.search(Float64Array.from(data[0].vector))
.limit(1)
.toArray();
expect(matched).toHaveLength(1);
expect(matched[0]["resource_id"]).toBe(0);
});
it("should support branches", async () => {
await table.add([{ id: 1 }]);
expect(await table.countRows()).toBe(1);
+2 -2
View File
@@ -26,7 +26,7 @@ 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"] }
pyo3 = { version = "0.28", features = ["extension-module", "abi3-py39", "chrono"] }
chrono = { version = "0.4", default-features = false, features = ["clock"] }
pyo3-async-runtimes = { version = "0.28", features = [
"attributes",
@@ -43,7 +43,7 @@ libc = "0.2"
[build-dependencies]
pyo3-build-config = { version = "0.28", features = [
"extension-module",
"abi3-py310",
"abi3-py39",
] }
[features]
+1 -1
View File
@@ -804,7 +804,7 @@ class LanceDBConnection(DBConnection):
"manifest_enabled": self._manifest_enabled,
"namespace_client_properties": self._namespace_client_properties,
"read_consistency_interval_seconds": (
rci.total_seconds() if rci is not None else None
rci.total_seconds() if rci else None
),
}
)
+58 -4
View File
@@ -3,6 +3,7 @@
from typing import List
from urllib.parse import unquote, urlparse
import numpy as np
@@ -101,7 +102,8 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
@weak_lru(maxsize=1)
def ndims(self):
return len(self.generate_embeddings([[self.source_instruction, "foo"]])[0])
model = self.get_model()
return model.encode("foo").shape[0]
def compute_query_embeddings(self, query: str, *args, **kwargs) -> List[np.array]:
return self.generate_embeddings([[self.query_instruction, query]])
@@ -124,9 +126,20 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
@weak_lru(maxsize=1)
def get_model(self):
instructor_embedding = attempt_import_or_raise(
"InstructorEmbedding", "InstructorEmbedding"
)
huggingface_hub = attempt_import_or_raise("huggingface_hub", "huggingface-hub")
missing = object()
original_cached_download = getattr(huggingface_hub, "cached_download", missing)
if original_cached_download is missing:
huggingface_hub.cached_download = _cached_download(huggingface_hub)
try:
instructor_embedding = attempt_import_or_raise(
"InstructorEmbedding", "InstructorEmbedding"
)
finally:
if original_cached_download is missing:
del huggingface_hub.cached_download
torch = attempt_import_or_raise("torch", "torch")
model = instructor_embedding.INSTRUCTOR(self.name)
@@ -139,3 +152,44 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
model, {torch.nn.Linear}, dtype=torch.qint8
)
return model
def _cached_download(huggingface_hub):
"""Provide the legacy download API used by sentence-transformers 2.2.x."""
def cached_download(
*,
url,
cache_dir=None,
force_filename=None,
library_name=None,
library_version=None,
user_agent=None,
use_auth_token=None,
**_,
):
path = urlparse(url).path.lstrip("/")
try:
repo_id, resolved_path = path.split("/resolve/", maxsplit=1)
revision, filename = resolved_path.split("/", maxsplit=1)
except ValueError as err:
raise ValueError(f"Unsupported Hugging Face Hub URL: {url}") from err
repo_id = unquote(repo_id)
revision = unquote(revision)
filename = unquote(filename)
# sentence-transformers derives force_filename from this Hub path with
# os.path.join. Using the URL path beneath local_dir produces the same
# local destination without sending Windows separators to the Hub.
return huggingface_hub.hf_hub_download(
repo_id=repo_id,
filename=filename,
revision=revision,
local_dir=cache_dir,
library_name=library_name,
library_version=library_version,
user_agent=user_agent,
token=use_auth_token,
)
return cached_download
+1 -17
View File
@@ -482,16 +482,6 @@ class LanceNamespaceDBConnection(DBConnection):
def serialize(self) -> str:
import json
if (
self._namespace_client_impl is None
or self._namespace_client_properties is None
):
raise ValueError(
"Cannot serialize a namespace connection constructed from an "
"opaque namespace client. Pass namespace_client_impl and "
"namespace_client_properties when constructing the connection."
)
return json.dumps(
{
"connection_type": "namespace",
@@ -503,7 +493,7 @@ class LanceNamespaceDBConnection(DBConnection):
"storage_options": self.storage_options or None,
"read_consistency_interval_seconds": (
self.read_consistency_interval.total_seconds()
if self.read_consistency_interval is not None
if self.read_consistency_interval
else None
),
}
@@ -579,7 +569,6 @@ class LanceNamespaceDBConnection(DBConnection):
self,
name,
namespace_path=namespace_path,
storage_options=storage_options,
namespace_client=self._namespace_client,
pushdown_operations=self._namespace_client_pushdown_operations,
route_pushdown_to_rust=self._route_pushdown_to_rust,
@@ -618,8 +607,6 @@ class LanceNamespaceDBConnection(DBConnection):
self,
name,
namespace_path=namespace_path,
storage_options=storage_options,
index_cache_size=index_cache_size,
namespace_client=self._namespace_client,
pushdown_operations=self._namespace_client_pushdown_operations,
route_pushdown_to_rust=self._route_pushdown_to_rust,
@@ -912,13 +899,10 @@ class LanceNamespaceDBConnection(DBConnection):
self,
name,
namespace_path=namespace_path,
storage_options=storage_options,
index_cache_size=index_cache_size,
location=table_uri,
namespace_client=namespace_client,
managed_versioning=managed_versioning,
pushdown_operations=self._namespace_client_pushdown_operations,
route_pushdown_to_rust=self._route_pushdown_to_rust,
_async=async_table,
)
-7
View File
@@ -591,13 +591,6 @@ class Permutation:
then the first split will be used.
"""
assert base_table is not None, "base_table is required"
# A PyTorch fork worker may construct its Permutation lazily from a
# table opened in the parent process. Reopen that table before the
# Rust reader clones its object-store clients and connection pools.
if hasattr(base_table, "_ensure_open"):
base_table._ensure_open()
if permutation_table is not None and hasattr(permutation_table, "_ensure_open"):
permutation_table._ensure_open()
if split is not None:
if permutation_table is None:
raise ValueError(
+9 -301
View File
@@ -6,8 +6,6 @@ from __future__ import annotations
import asyncio
import inspect
import deprecation
import os
import threading
import warnings
from abc import ABC, abstractmethod
from dataclasses import dataclass
@@ -165,7 +163,7 @@ def _maybe_add_fts_error_note(
if TYPE_CHECKING:
from .db import DBConnection, LanceDBConnection
from .db import LanceDBConnection
from ._lancedb import (
Table as LanceDBTable,
OptimizeStats,
@@ -2107,23 +2105,6 @@ class Table(ABC):
"""
@dataclass
class _LanceTableReopenState:
"""Process-independent coordinates for reopening a native table."""
connection_state: Optional[str]
can_reopen_after_fork: bool
fork_reopen_error: Optional[str]
name: str
namespace_path: List[str]
storage_options: Optional[Dict[str, str]]
index_cache_size: Optional[int]
location: Optional[str]
managed_versioning: Optional[bool]
branch: Optional[str]
checkout_version: Optional[int]
class LanceTable(Table):
"""
A table in a LanceDB database.
@@ -2158,10 +2139,7 @@ class LanceTable(Table):
namespace_path = []
self._conn = connection
self._namespace_path = namespace_path
self._storage_options = storage_options
self._index_cache_size = index_cache_size
self._location = location # Store location for use in _dataset_path
self._managed_versioning = managed_versioning
self._namespace_client = namespace_client
self._pushdown_operations = pushdown_operations or set()
# When the connection built the namespace client natively (e.g. an
@@ -2186,206 +2164,9 @@ class LanceTable(Table):
managed_versioning=managed_versioning,
)
)
self._initialize_reopen_state(name)
def _initialize_reopen_state(self, name: str) -> None:
"""Capture the state needed to replace inherited native handles."""
self._name = name
self._pid = os.getpid()
self._native_state_guard = (self._pid, threading.RLock())
# A native table owns object-store clients and connection pools. Those
# handles must not be used after fork, so retain a process-independent
# connection description while it is still safe to inspect the parent
# connection. Connections without reconstructible metadata are not
# safe to reuse in a forked child, so retain a clear diagnostic rather
# than advertising them as reopenable based on JSON encoding alone.
try:
connection_uri: Optional[str] = self._conn.uri
except Exception:
connection_uri = None
fork_reopen_error: Optional[str] = None
try:
connection_state: Optional[str] = self._conn.serialize()
can_reopen_after_fork = connection_uri is not None and not (
connection_uri.startswith("memory://")
)
except Exception as error:
connection_state = None
can_reopen_after_fork = False
if connection_uri is not None and not connection_uri.startswith(
"memory://"
):
fork_reopen_error = (
f"Cannot reopen table {name!r} in a forked process: {error}"
)
self._reopen_state = _LanceTableReopenState(
connection_state=connection_state,
can_reopen_after_fork=can_reopen_after_fork,
fork_reopen_error=fork_reopen_error,
name=name,
namespace_path=list(self._namespace_path),
storage_options=(
dict(self._storage_options)
if self._storage_options is not None
else None
),
index_cache_size=self._index_cache_size,
location=self._location,
managed_versioning=self._managed_versioning,
branch=self._table.current_branch(),
checkout_version=None,
)
@property
def _connection_state(self) -> Optional[str]:
"""Serialized connection retained for worker reconstruction."""
return self._reopen_state.connection_state
@property
def _can_reopen_after_fork(self) -> bool:
return self._reopen_state.can_reopen_after_fork
@property
def _branch(self) -> Optional[str]:
state = getattr(self, "_reopen_state", None)
if state is not None:
return state.branch
return getattr(self, "_legacy_branch", None)
@_branch.setter
def _branch(self, value: Optional[str]) -> None:
state = getattr(self, "_reopen_state", None)
if state is not None:
state.branch = value
else:
self._legacy_branch = value
@property
def _checkout_version(self) -> Optional[int]:
state = getattr(self, "_reopen_state", None)
if state is not None:
return state.checkout_version
return getattr(self, "_legacy_checkout_version", None)
@_checkout_version.setter
def _checkout_version(self, value: Optional[int]) -> None:
state = getattr(self, "_reopen_state", None)
if state is not None:
state.checkout_version = value
else:
self._legacy_checkout_version = value
def _native_state_lock(self):
"""Return the per-process lock coordinating native mode and reopen state."""
pid = os.getpid()
guard = getattr(self, "_native_state_guard", None)
if guard is None:
candidate = (pid, threading.RLock())
guard = self.__dict__.setdefault("_native_state_guard", candidate)
elif guard[0] != pid:
# A lock inherited while another parent thread held it cannot be
# safely acquired in the child. Child state starts single-threaded,
# so replace it before coordinating the first reopen.
guard = (pid, threading.RLock())
self._native_state_guard = guard
return guard[1]
@classmethod
def _open_from_reopen_state(
cls,
connection: "DBConnection",
state: "_LanceTableReopenState",
) -> "LanceTable":
"""Open a table from its complete process-independent descriptor."""
async_connection = getattr(connection, "_conn", None)
if async_connection is None:
async_connection = connection._inner
namespace_client = getattr(connection, "_namespace_client", None)
async_table = LOOP.run(
async_connection.open_table(
state.name,
namespace_path=state.namespace_path,
storage_options=state.storage_options,
index_cache_size=state.index_cache_size,
location=state.location,
namespace_client=namespace_client,
managed_versioning=state.managed_versioning,
)
)
table = cls(
connection,
state.name,
namespace_path=state.namespace_path,
storage_options=state.storage_options,
index_cache_size=state.index_cache_size,
location=state.location,
namespace_client=namespace_client,
managed_versioning=state.managed_versioning,
pushdown_operations=getattr(
connection, "_namespace_client_pushdown_operations", None
),
route_pushdown_to_rust=getattr(
connection, "_route_pushdown_to_rust", False
),
_async=async_table,
)
if state.branch is not None:
table = table.branches.checkout(state.branch, state.checkout_version)
elif state.checkout_version is not None:
table.checkout(state.checkout_version)
return table
def _ensure_open(self) -> None:
"""Reopen native table handles inherited from another process."""
with self._native_state_lock():
pid = os.getpid()
if getattr(self, "_pid", pid) == pid:
return
state = getattr(self, "_reopen_state", None)
fork_reopen_error = getattr(state, "fork_reopen_error", None)
if fork_reopen_error is not None:
raise RuntimeError(fork_reopen_error)
if (
state is None
or not state.can_reopen_after_fork
or state.connection_state is None
):
# In-memory and opaque Rust-only connections cannot be recreated
# from connection metadata. Their local handles retain the prior
# best-effort fork behavior.
self._pid = pid
return
from lancedb import deserialize_conn
connection = deserialize_conn(state.connection_state, for_worker=True)
reopened = self._open_from_reopen_state(
connection,
state,
)
# Keep this Python object stable because user datasets commonly retain
# it across fork. Replace every process-bound component with the fresh
# child's equivalent.
self._conn = reopened._conn
self._table = reopened._table
self._namespace_client = reopened._namespace_client
self._pushdown_operations = reopened._pushdown_operations
self._route_pushdown_to_rust = reopened._route_pushdown_to_rust
self._reopen_state = reopened._reopen_state
self._pid = pid
@property
def name(self) -> str:
if hasattr(self, "_name"):
return self._name
# Preserve compatibility with lightweight / legacy instances that
# were constructed without running ``LanceTable.__init__``.
return self._table.name
@property
@@ -2602,75 +2383,18 @@ class LanceTable(Table):
def _wrap_branch_handle(
self, async_table: "AsyncTable", version: Optional[int] = None
) -> "LanceTable":
table = LanceTable(
# version is unused locally: the pin already lives on async_table and a
# local handle is not reopened via a serialized connection.
return LanceTable(
self._conn,
async_table.name,
namespace_path=self._namespace_path,
storage_options=self._storage_options,
index_cache_size=self._index_cache_size,
namespace_client=self._namespace_client,
pushdown_operations=self._pushdown_operations,
route_pushdown_to_rust=self._route_pushdown_to_rust,
location=self._location,
managed_versioning=self._managed_versioning,
_async=async_table,
)
table._checkout_version = version
return table
def _resolve_checkout_version(self, version: Union[int, str]) -> int:
if isinstance(version, int):
return version
try:
return self.tags.get_version(version)
except RuntimeError as err:
# Native checkout historically exposes an unknown tag as ValueError.
# Preserve that contract while resolving tags before mutating the table.
if "Ref not found" in str(err) and "does not exist" in str(err):
raise ValueError(str(err)) from err
raise
async def _commit_native_state(
self,
transition,
version: Optional[int],
started: threading.Event,
finished: threading.Event,
):
"""Commit a native transition and its fork coordinate as one task."""
started.set()
try:
task = asyncio.ensure_future(transition)
try:
result = await asyncio.shield(task)
except asyncio.CancelledError:
# BackgroundEventLoop cancels its submitted task when the
# waiting caller is interrupted. Let an already-started native
# transition reach its authoritative terminal state before the
# per-table boundary is released.
result = await task
self._checkout_version = version
raise
self._checkout_version = version
return result
finally:
finished.set()
def _run_native_state_transition(self, transition, version: Optional[int]):
started = threading.Event()
finished = threading.Event()
try:
return LOOP.run(
self._commit_native_state(transition, version, started, finished)
)
except BaseException:
if started.is_set():
while not finished.is_set():
try:
finished.wait()
except BaseException: # noqa: PERF203
continue
raise
def checkout(self, version: Union[int, str]):
"""Checkout a version of the table. This is an in-place operation.
@@ -2708,14 +2432,7 @@ class LanceTable(Table):
vector type
0 [1.1, 0.9] vector
"""
# Resolve tags before mutating the native handle. This leaves the live
# handle and reopen descriptor aligned if tag lookup fails, and avoids a
# second fallible version lookup after checkout succeeds.
with self._native_state_lock():
resolved_version = self._resolve_checkout_version(version)
self._run_native_state_transition(
self._table.checkout(resolved_version), resolved_version
)
LOOP.run(self._table.checkout(version))
def checkout_latest(self):
"""Checkout the latest version of the table. This is an in-place operation.
@@ -2723,8 +2440,7 @@ class LanceTable(Table):
The table will be set back into standard mode, and will track the latest
version of the table.
"""
with self._native_state_lock():
self._run_native_state_transition(self._table.checkout_latest(), None)
LOOP.run(self._table.checkout_latest())
def restore(self, version: Optional[Union[int, str]] = None):
"""Restore a version of the table. This is an in-place operation.
@@ -2770,13 +2486,9 @@ class LanceTable(Table):
>>> len(table.list_versions())
4
"""
with self._native_state_lock():
if version is not None:
resolved_version = self._resolve_checkout_version(version)
self._run_native_state_transition(
self._table.checkout(resolved_version), resolved_version
)
self._run_native_state_transition(self._table.restore(), None)
if version is not None:
LOOP.run(self._table.checkout(version))
LOOP.run(self._table.restore())
def count_rows(self, filter: Optional[str] = None) -> int:
return LOOP.run(self._table.count_rows(filter))
@@ -3887,9 +3599,7 @@ class LanceTable(Table):
self = cls.__new__(cls)
self._conn = db
self._namespace_path = namespace_path
self._index_cache_size = None
self._location = location
self._managed_versioning = None
self._namespace_client = namespace_client
self._pushdown_operations = pushdown_operations or set()
self._route_pushdown_to_rust = route_pushdown_to_rust
@@ -3917,7 +3627,6 @@ class LanceTable(Table):
enable_v2_manifest_paths
)
self._storage_options = storage_options
self._table = LOOP.run(
self._conn._conn.create_table(
name,
@@ -3934,7 +3643,6 @@ class LanceTable(Table):
namespace_client=namespace_client,
)
)
self._initialize_reopen_state(name)
return self
def delete(self, where: Union[str, Expr]) -> DeleteResult:
-5
View File
@@ -395,11 +395,6 @@ def _(value: dict):
)
@value_to_sql.register(pa.Scalar)
def _(value: pa.Scalar):
return value_to_sql(value.as_py())
@value_to_sql.register(np.ndarray)
def _(value: np.ndarray):
return value_to_sql(value.tolist())
+2 -21
View File
@@ -2,8 +2,6 @@
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import json
import inspect
import re
import sys
from datetime import timedelta
@@ -64,23 +62,17 @@ def test_basic(tmp_path):
assert db.open_table("test").name == db["test"].name
def test_sync_debugger_inspection_does_not_use_background_loop(tmp_path, monkeypatch):
def test_sync_repr_does_not_use_background_loop(tmp_path, monkeypatch):
from lancedb.background_loop import LOOP
db = lancedb.connect(tmp_path)
table = db.create_table("test", data=[{"id": 1}])
def fail_run(*args, **kwargs):
raise AssertionError("debugger inspection should not use the background loop")
raise AssertionError("repr should not use the Python background loop")
monkeypatch.setattr(LOOP, "run", fail_run)
# Debuggers enumerate and evaluate every exposed attribute when expanding a
# variable. This must remain safe while their breakpoint suspends LOOP's thread.
members = dict(inspect.getmembers(db))
assert members["uri"] == str(tmp_path)
assert members["read_consistency_interval"] is None
assert repr(db) == f"LanceDBConnection(uri={str(tmp_path)!r})"
assert repr(table) == f"LanceTable(name='test', _conn={db!r})"
@@ -102,17 +94,6 @@ def test_read_consistency_interval_does_not_use_background_loop(tmp_path, monkey
assert db_from_inner.read_consistency_interval == consistency_interval
def test_serialize_preserves_zero_read_consistency_interval(tmp_path):
db = lancedb.connect(tmp_path, read_consistency_interval=timedelta(0))
table = db.create_table("items", pa.table({"x": [1]}))
encoded = json.loads(table._connection_state)
assert encoded["read_consistency_interval_seconds"] == 0.0
restored = lancedb.deserialize_conn(table._connection_state)
assert restored.read_consistency_interval == timedelta(0)
def test_ingest_pd(tmp_path):
db = lancedb.connect(tmp_path)
+83 -31
View File
@@ -1,8 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import ntpath
import os
import pickle
import sys
from types import ModuleType
from typing import List, Optional, Union
from unittest.mock import MagicMock, patch
@@ -64,23 +67,6 @@ def test_embedding_function(tmp_path):
assert np.allclose(actual, expected)
def test_instructor_ndims_uses_instruction():
instructor = get_registry().get("instructor").create()
model = MagicMock()
model.encode.return_value = np.zeros((1, 384))
with patch.object(type(instructor), "get_model", return_value=model):
assert instructor.ndims() == 384
model.encode.assert_called_once_with(
[[instructor.source_instruction, "foo"]],
batch_size=instructor.batch_size,
show_progress_bar=instructor.show_progress_bar,
normalize_embeddings=instructor.normalize_embeddings,
device=instructor.device,
)
def test_embedding_function_variables():
@register("variable-testing")
class VariableTestingFunction(TextEmbeddingFunction):
@@ -132,16 +118,34 @@ def test_embedding_function_variables():
assert func.safe_model_dump()["secret_key"] == "$var:secret"
def test_openai_variables_survive_metadata_round_trip():
def test_parse_functions_with_variables():
@register("variable-parsing-test")
class VariableParsingFunction(TextEmbeddingFunction):
api_key: str
base_url: Optional[str] = None
@staticmethod
def sensitive_keys():
return ["api_key"]
def ndims(self):
return 10
def generate_embeddings(self, texts):
# Mock implementation that just returns random embeddings
# In real usage, this would use the api_key to call an API
return [np.random.rand(self.ndims()).tolist() for _ in texts]
registry = EmbeddingFunctionRegistry.get_instance()
registry.set_var("test_api_key", "sk-test-key-12345")
registry.set_var("test_base_url", "https://api.example.com")
conf = EmbeddingFunctionConfig(
source_column="text",
vector_column="vector",
function=registry.get("openai").create(
api_key="$var:test_api_key", base_url="https://api.example.com"
function=registry.get("variable-parsing-test").create(
api_key="$var:test_api_key", base_url="$var:test_base_url"
),
)
@@ -149,10 +153,7 @@ def test_openai_variables_survive_metadata_round_trip():
# Create a mock arrow table with the metadata
schema = pa.schema(
[
pa.field("text", pa.string()),
pa.field("vector", pa.list_(pa.float32(), 1536)),
]
[pa.field("text", pa.string()), pa.field("vector", pa.list_(pa.float32(), 10))]
)
table = pa.table({"text": [], "vector": []}, schema=schema)
table = table.replace_schema_metadata(metadata)
@@ -166,15 +167,13 @@ def test_openai_variables_survive_metadata_round_trip():
assert parsed_func.api_key == "sk-test-key-12345"
assert parsed_func.base_url == "https://api.example.com"
embeddings = parsed_func.generate_embeddings(["test text"])
assert len(embeddings) == 1
assert len(embeddings[0]) == 10
assert parsed_func.safe_model_dump()["api_key"] == "$var:test_api_key"
with patch("lancedb.embeddings.openai.attempt_import_or_raise") as import_openai:
parsed_func._openai_client
import_openai.return_value.OpenAI.assert_called_once_with(
api_key="sk-test-key-12345", base_url="https://api.example.com"
)
def test_embedding_with_bad_results(tmp_path):
@register("null-embedding")
@@ -526,6 +525,59 @@ def test_embedding_function_safe_model_dump(embedding_type):
)
def test_instructor_embedding_supports_huggingface_hub_without_cached_download(
tmp_path, monkeypatch
):
from lancedb.embeddings.instructor import InstructorEmbeddingFunction
hub_download = MagicMock(return_value="/cache/1_Pooling/config.json")
huggingface_hub = ModuleType("huggingface_hub")
huggingface_hub.hf_hub_download = hub_download
torch = ModuleType("torch")
monkeypatch.setitem(sys.modules, "huggingface_hub", huggingface_hub)
monkeypatch.setitem(sys.modules, "torch", torch)
monkeypatch.delitem(sys.modules, "InstructorEmbedding", raising=False)
monkeypatch.syspath_prepend(str(tmp_path))
(tmp_path / "InstructorEmbedding.py").write_text(
"from huggingface_hub import cached_download\n\n"
"class INSTRUCTOR:\n"
" def __init__(self, name):\n"
" self.name = name\n"
)
embedding = InstructorEmbeddingFunction.create(show_progress_bar=False)
instructor_model = embedding.get_model()
assert instructor_model.name == "hkunlp/instructor-base"
assert not hasattr(huggingface_hub, "cached_download")
instructor_embedding = sys.modules["InstructorEmbedding"]
path = instructor_embedding.cached_download(
url=(
"https://huggingface.co/hkunlp/instructor-base/resolve/abc123/"
"1_Pooling/config.json"
),
cache_dir="/cache",
force_filename=ntpath.join("1_Pooling", "config.json"),
library_name="sentence-transformers",
library_version="2.2.2",
use_auth_token="token",
)
assert path == "/cache/1_Pooling/config.json"
hub_download.assert_called_once_with(
repo_id="hkunlp/instructor-base",
filename="1_Pooling/config.json",
revision="abc123",
local_dir="/cache",
library_name="sentence-transformers",
library_version="2.2.2",
user_agent=None,
token="token",
)
@patch("time.sleep")
def test_retry(mock_sleep):
test_function = MagicMock(side_effect=[Exception] * 9 + ["result"])
+1 -81
View File
@@ -12,7 +12,7 @@ import pyarrow.compute as pc
import pytest
import pytest_asyncio
from lancedb.index import BTree, FTS, IvfPq
from lancedb.index import FTS
from lancedb.table import AsyncTable, Table
@@ -99,86 +99,6 @@ async def test_async_hybrid_query_filters(table: AsyncTable):
assert result["text"].to_pylist() == ["cat", "b"]
@pytest.mark.asyncio
async def test_hybrid_query_with_stale_fixed_size_binary_prefilter(
tmpdir_factory,
):
tmp_path = str(tmpdir_factory.mktemp("stale_scalar_prefilter"))
db = await lancedb.connect_async(tmp_path)
def fixed_size_binary(value: int) -> bytes:
return value.to_bytes(16, byteorder="big")
num_rows = 1000
data = pa.table(
{
"space_id": pa.array(
[fixed_size_binary(i) for i in range(num_rows)],
type=pa.binary(16),
),
"text": ["book"] * num_rows,
"vector": pa.array(
[[float(i), float(i)] for i in range(num_rows)],
type=pa.list_(pa.float32(), 2),
),
}
)
table = await db.create_table("test", data)
await table.create_index(
"vector", config=IvfPq(num_partitions=4, num_sub_vectors=2)
)
await table.create_index("space_id", config=BTree())
await table.create_index("text", config=FTS(with_position=False))
# Advance the search indices without advancing the scalar index. This is the
# state that previously let hybrid search use an incomplete scalar prefilter.
await table.add(data)
lance_dataset = await table.to_lance()
lance_dataset.optimize.optimize_indices(index_names=["vector_idx", "text_idx"])
await table.checkout_latest()
scalar_stats = await table.index_stats("space_id_idx")
assert scalar_stats is not None
assert scalar_stats.num_indexed_rows == num_rows
assert scalar_stats.num_unindexed_rows == num_rows
for index_name in ["vector_idx", "text_idx"]:
search_stats = await table.index_stats(index_name)
assert search_stats is not None
assert search_stats.num_indexed_rows == num_rows * 2
assert search_stats.num_unindexed_rows == 0
matching_ids = [5, 10, 15, 20, 25, 30]
literals = [
f"arrow_cast(0x{fixed_size_binary(i).hex()}, 'FixedSizeBinary(16)')"
for i in matching_ids
]
predicate = f"space_id IN ({', '.join(literals)})"
expected_ids = sorted(fixed_size_binary(i) for i in matching_ids for _ in range(2))
vector_query = (
table.query().where(predicate).nearest_to([5.0, 5.0]).limit(num_rows * 2)
)
vector_results = await vector_query.to_arrow()
assert sorted(vector_results["space_id"].to_pylist()) == expected_ids
fts_query = (
table.query().where(predicate).nearest_to_text("book").limit(num_rows * 2)
)
fts_results = await fts_query.to_arrow()
assert sorted(fts_results["space_id"].to_pylist()) == expected_ids
hybrid_results = await (
table.query()
.where(predicate)
.nearest_to([5.0, 5.0])
.nearest_to_text("book")
.limit(num_rows * 2)
.to_arrow()
)
assert sorted(hybrid_results["space_id"].to_pylist()) == expected_ids
@pytest.mark.asyncio
async def test_async_hybrid_query_default_limit(table: AsyncTable):
# add 10 new rows
-33
View File
@@ -1,33 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import re
import shutil
import subprocess
import sys
import lancedb._lancedb as _lancedb
import pytest
@pytest.mark.skipif(sys.platform != "linux", reason="ldd is Linux-specific")
def test_native_extension_does_not_link_openssl():
"""OpenSSL-linked wheels abort when imported on RHEL hosts in FIPS mode."""
ldd = shutil.which("ldd")
if ldd is None:
pytest.skip("ldd is not installed")
result = subprocess.run(
[ldd, _lancedb.__file__],
check=True,
capture_output=True,
text=True,
)
openssl_libraries = re.findall(
r"^\s*(lib(?:crypto|ssl)\S*)\s+=>", result.stdout, flags=re.MULTILINE
)
assert not openssl_libraries, (
"the LanceDB native extension must use rustls instead of linking OpenSSL: "
f"{openssl_libraries}"
)
-25
View File
@@ -372,31 +372,6 @@ async def test_create_vector_index(some_table: AsyncTable):
assert stats.num_indices == 1
@pytest.mark.asyncio
async def test_create_ivf_index_reports_unsplittable_partitions(db_async):
dim = 8
num_partitions = 300 # More than 256 selects hierarchical k-means.
base_vectors = [[float(row == column) for column in range(dim)] for row in range(5)]
vectors = pa.array(base_vectors * 200, pa.list_(pa.float32(), dim))
table = await db_async.create_table(
"unsplittable_partitions",
pa.table({"vector": vectors}),
)
error_pattern = (
rf"Cannot create {num_partitions} IVF partitions: k-means could only form"
)
with pytest.raises(RuntimeError, match=error_pattern):
await table.create_index(
"vector",
config=IvfFlat(
distance_type="dot",
num_partitions=num_partitions,
max_iterations=10,
),
)
@pytest.mark.asyncio
async def test_create_4bit_ivfpq_index(some_table: AsyncTable):
# Can create
+1 -103
View File
@@ -6,13 +6,9 @@
import tempfile
import shutil
import importlib
import multiprocessing as mp
import sys
from datetime import timedelta
import pytest
import pyarrow as pa
import lancedb
from lance_namespace import connect as namespace_connect
from lance_namespace.errors import NamespaceNotEmptyError, TableNotFoundError
from lancedb.namespace import _MAX_QUERY_K
from lancedb.table import AsyncTable, LanceTable
@@ -76,16 +72,6 @@ def _namespace_lance_table(namespace_client: _NamespaceClient) -> LanceTable:
return table
def _direct_namespace_fork_child(table, result_queue):
from lancedb.permutation import Permutation
try:
permutation = Permutation.identity(table)
result_queue.put(("ok", permutation.num_rows))
except Exception as error:
result_queue.put((type(error).__name__, str(error)))
class TestNamespaceConnection:
"""Test namespace-based LanceDB connection using DirectoryNamespace."""
@@ -433,95 +419,7 @@ class TestNamespaceConnection:
pa.field("vector", pa.list_(pa.float32(), 2)),
]
)
created = db.create_table(
"test_table", schema=schema, storage_options=table_opts
)
assert created._storage_options == table_opts
opened = db.open_table(
"test_table",
storage_options={"allow_http": "true"},
index_cache_size=17,
)
assert opened._storage_options == {"allow_http": "true"}
assert opened._index_cache_size == 17
opened._pid = -1
opened._ensure_open()
assert opened.count_rows() == 0
def test_serialize_preserves_zero_read_consistency_interval(self):
db = lancedb.connect_namespace(
"dir",
{"root": self.temp_dir},
read_consistency_interval=timedelta(0),
)
restored = lancedb.deserialize_conn(db.serialize())
assert restored.read_consistency_interval == timedelta(0)
@pytest.mark.skipif(
sys.platform != "linux",
reason="fork() is only supported safely for this test on Linux",
)
def test_direct_namespace_with_descriptor_reopens_after_fork(self):
properties = {"root": self.temp_dir}
namespace = namespace_connect("dir", properties)
db = lancedb.LanceNamespaceDBConnection(
namespace,
namespace_client_impl="dir",
namespace_client_properties=properties,
)
table = db.create_table("items", pa.table({"id": [1]}))
ctx = mp.get_context("fork")
result_queue = ctx.Queue()
process = ctx.Process(
target=_direct_namespace_fork_child,
args=(table, result_queue),
)
process.start()
process.join(10)
if process.is_alive():
process.terminate()
process.join(5)
pytest.fail("Direct namespace table hung while reopening after fork")
assert process.exitcode == 0
assert result_queue.get(timeout=2) == ("ok", 1)
@pytest.mark.skipif(
sys.platform != "linux",
reason="fork() is only supported safely for this test on Linux",
)
def test_opaque_direct_namespace_reports_unsupported_fork(self):
namespace = namespace_connect("dir", {"root": self.temp_dir})
db = lancedb.LanceNamespaceDBConnection(namespace)
table = db.create_table("items", pa.table({"id": [1]}))
with pytest.raises(ValueError, match="opaque namespace client"):
db.serialize()
assert not table._can_reopen_after_fork
ctx = mp.get_context("fork")
result_queue = ctx.Queue()
process = ctx.Process(
target=_direct_namespace_fork_child,
args=(table, result_queue),
)
process.start()
process.join(10)
if process.is_alive():
process.terminate()
process.join(5)
pytest.fail("Opaque namespace table hung after fork")
assert process.exitcode == 0
error_type, message = result_queue.get(timeout=2)
assert error_type == "RuntimeError"
assert "Cannot reopen table 'items' in a forked process" in message
assert "namespace_client_impl and namespace_client_properties" in message
db.create_table("test_table", schema=schema, storage_options=table_opts)
def test_namespace_operations(self):
"""Test namespace management operations."""
@@ -1,42 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import importlib
import re
import sys
from pathlib import Path
import pytest
def test_pyo3_abi_matches_minimum_supported_python():
project_dir = Path(__file__).parents[2]
pyproject = (project_dir / "pyproject.toml").read_text()
cargo_manifest = (project_dir / "Cargo.toml").read_text()
minimum_python = re.search(
r'^requires-python\s*=\s*">=(\d+)\.(\d+)"$', pyproject, re.MULTILINE
)
assert minimum_python is not None
major, minor = minimum_python.groups()
expected_abi = f"abi3-py{major}{minor}"
configured_abis = re.findall(r'"(abi3-py\d+)"', cargo_manifest)
assert configured_abis == [expected_abi, expected_abi], (
"the pyo3 runtime and build ABI features must both match requires-python"
)
@pytest.mark.skipif(sys.platform != "win32", reason="Windows wheel regression test")
def test_windows_wheel_tag_and_native_import():
project_dir = Path(__file__).parents[2]
wheels = list((project_dir.parent / "target" / "wheels").glob("lancedb-*.whl"))
if not wheels:
pytest.skip("no wheel artifact is available in this development environment")
assert len(wheels) == 1
assert wheels[0].name.endswith("-cp310-abi3-win_amd64.whl")
native_module = importlib.import_module("lancedb._lancedb")
assert Path(native_module.__file__).suffix == ".pyd"
-6
View File
@@ -35,12 +35,6 @@ def make_mock_http_handler(handler):
return MockLanceDBHandler
@pytest.mark.parametrize("db_name", ["a" * 64, "invalid..database"])
def test_connect_rejects_invalid_cloud_dns_hostname(db_name):
with pytest.raises(ValueError, match="DNS labels must contain 1 to 63 bytes"):
lancedb.connect(f"db://{db_name}", api_key="fake")
@contextlib.contextmanager
def mock_lancedb_connection(handler):
with http.server.HTTPServer(
+3 -440
View File
@@ -2,15 +2,11 @@
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import asyncio
import ctypes
import gc
import os
import sys
import threading
import warnings
import weakref
from concurrent.futures import CancelledError, ThreadPoolExecutor
from concurrent.futures import ThreadPoolExecutor
from datetime import date, datetime, timedelta
from time import sleep
from typing import List
@@ -103,30 +99,6 @@ def test_basic(mem_db: DBConnection):
assert table.to_arrow() == expected_data
def test_search_preserves_nulls_from_sliced_arrow_table(mem_db: DBConnection):
data = pa.table(
{
"id": [0, 1, 2, 3, 4],
"score_cn": [None, 22, None, 5, 8],
"score_mt": [None, 42, None, 5, 8],
"vector": [
[20, 19, -1, -1],
[41, 38, 22, 42],
[10, 10, -1, -1],
[5, 5, 5, 5],
[8, 8, 8, 8],
],
}
).slice(1)
table = mem_db.create_table("sliced_nullable", data=data)
result = table.search([41, 38, 22, 42]).limit(1).to_arrow()
assert result["id"].to_pylist() == [1]
assert result["score_cn"].to_pylist() == [22]
assert result["score_mt"].to_pylist() == [42]
def test_table_to_pandas_default_matches_arrow(tmp_db: DBConnection):
pd = pytest.importorskip("pandas")
data = pa.table({"id": [1, 2], "text": ["one", "two"]})
@@ -463,38 +435,6 @@ def test_add(mem_db: DBConnection):
_add(table, schema)
def test_add_releases_arrow_buffers_without_gc(mem_db: DBConnection):
"""Regression test for https://github.com/lancedb/lancedb/issues/2512."""
schema = pa.schema([pa.field("x", pa.int64())])
table = mem_db.create_table("test_add_releases_arrow_buffers", schema=schema)
class BufferOwner:
def __init__(self, size: int):
self.memory = ctypes.create_string_buffer(size)
owner_refs = []
gc_was_enabled = gc.isenabled()
gc.disable()
try:
for _ in range(3):
size = 8 * 1024
owner = BufferOwner(size)
arrow_buffer = pa.foreign_buffer(
ctypes.addressof(owner.memory), size, owner
)
array = pa.Array.from_buffers(pa.int64(), 1024, [None, arrow_buffer])
batch = pa.RecordBatch.from_arrays([array], schema=schema)
owner_refs.append(weakref.ref(owner))
table.add(batch)
del batch, array, arrow_buffer, owner
assert all(owner_ref() is None for owner_ref in owner_refs)
finally:
if gc_was_enabled:
gc.enable()
def test_add_write_parallelism(mem_db: DBConnection):
schema = pa.schema([pa.field("id", pa.int64())])
table = mem_db.create_table("test", schema=schema)
@@ -1885,33 +1825,6 @@ def test_add_nullable_struct_with_none(mem_db: DBConnection):
assert result.column("data").to_pylist() == [{"x": 1.0}, None]
def test_read_mostly_null_list_v2_2_page_boundary(tmp_path):
# Regression test for #3194. This row/value count crosses a v2.2 structural
# encoding page boundary where Lance 3.0.0 sliced repetition/definition
# levels by row offset and decoded child arrays at different lengths.
num_rows = 64_885
num_values = 217
list_type = pa.list_(pa.float32())
source = pa.table(
{
"id": np.arange(num_rows, dtype=np.int64),
"coords": pa.array(
[[1.0, 2.0, 3.0, 4.0]] * num_values + [None] * (num_rows - num_values),
type=list_type,
),
}
)
db = lancedb.connect(
tmp_path,
storage_options={"new_table_data_storage_version": "2.2"},
)
table = db.create_table("test_sparse_nullable_list", data=source)
result = table.search().select(["id", "coords"]).limit(num_rows).to_arrow()
assert result.equals(source)
def test_add_with_integer_embeddings_preserves_casting(mem_db: DBConnection):
class Schema(LanceModel):
text: str
@@ -2146,260 +2059,6 @@ def test_restore(mem_db: DBConnection):
table.restore(0)
def test_restore_tracks_checkout_when_restore_fails():
class FailingRestore:
def __init__(self):
self.live_version = None
async def checkout(self, version):
self.live_version = version
async def restore(self):
raise RuntimeError("injected restore failure")
inner = FailingRestore()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = None
with pytest.raises(RuntimeError, match="injected restore failure"):
table.restore(7)
assert table._checkout_version == inner.live_version
@pytest.mark.parametrize(
("operation", "expected_descriptor", "expected_restore_calls"),
[("checkout", 11, 0), ("restore", None, 1)],
)
def test_string_tag_resolves_before_checkout(
operation, expected_descriptor, expected_restore_calls
):
class Tags:
async def get_version(self, tag):
assert tag == "tag-v1"
return 11
class NoPostCheckoutVersionLookup:
def __init__(self):
self.tags = Tags()
self.checkout_versions = []
self.restore_calls = 0
async def checkout(self, version):
self.checkout_versions.append(version)
async def version(self):
raise RuntimeError("post-checkout version lookup must not run")
async def restore(self):
self.restore_calls += 1
inner = NoPostCheckoutVersionLookup()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = 3
getattr(table, operation)("tag-v1")
assert inner.checkout_versions == [11]
assert table._checkout_version == expected_descriptor
assert inner.restore_calls == expected_restore_calls
@pytest.mark.parametrize("operation", ["checkout", "restore"])
def test_string_tag_resolution_failure_does_not_mutate_handle(operation):
class FailingTags:
async def get_version(self, tag):
assert tag == "missing-tag"
raise RuntimeError("injected tag lookup failure")
class UnchangedTable:
def __init__(self):
self.tags = FailingTags()
self.checkout_calls = 0
self.restore_calls = 0
async def checkout(self, version):
self.checkout_calls += 1
async def restore(self):
self.restore_calls += 1
inner = UnchangedTable()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = 3
with pytest.raises(RuntimeError, match="injected tag lookup failure"):
getattr(table, operation)("missing-tag")
assert table._checkout_version == 3
assert inner.checkout_calls == 0
assert inner.restore_calls == 0
def test_native_state_transitions_are_serialized(monkeypatch):
from lancedb.background_loop import LOOP
class Inner:
def __init__(self):
self.live_version = None
async def checkout(self, version):
self.live_version = version
inner = Inner()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = None
first_native_done = threading.Event()
release_first_call = threading.Event()
second_call_started = threading.Event()
second_call_done = threading.Event()
errors = []
original_run = LOOP.run
def delayed_delivery(awaitable):
result = original_run(awaitable)
if threading.current_thread().name == "checkout-1":
first_native_done.set()
assert release_first_call.wait(5)
return result
monkeypatch.setattr(LOOP, "run", delayed_delivery)
def checkout(version):
if version == 2:
second_call_started.set()
try:
table.checkout(version)
except BaseException as err:
errors.append(err)
finally:
if version == 2:
second_call_done.set()
first = threading.Thread(target=checkout, args=(1,), name="checkout-1")
first.start()
assert first_native_done.wait(5)
second = threading.Thread(target=checkout, args=(2,), name="checkout-2")
second.start()
assert second_call_started.wait(5)
assert not second_call_done.wait(0.1)
release_first_call.set()
first.join(5)
second.join(5)
assert not first.is_alive()
assert not second.is_alive()
assert errors == []
assert inner.live_version == 2
assert table._checkout_version == 2
@pytest.mark.parametrize(
("operation", "args", "initial_version", "expected_version"),
[
("checkout", (11,), 3, 11),
("checkout_latest", (), 3, None),
("restore", (), 11, None),
],
)
def test_native_state_commits_before_success_delivery(
monkeypatch, operation, args, initial_version, expected_version
):
from lancedb.background_loop import LOOP
class Inner:
def __init__(self, live_version):
self.live_version = live_version
async def checkout(self, version):
self.live_version = version
async def checkout_latest(self):
self.live_version = None
async def restore(self):
self.live_version = None
inner = Inner(initial_version)
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = initial_version
original_run = LOOP.run
def success_then_interrupt(awaitable):
original_run(awaitable)
raise KeyboardInterrupt("injected after native success")
monkeypatch.setattr(LOOP, "run", success_then_interrupt)
with pytest.raises(KeyboardInterrupt, match="injected after native success"):
getattr(table, operation)(*args)
assert inner.live_version == expected_version
assert table._checkout_version == expected_version
def test_native_state_waits_for_cancelled_delivery(monkeypatch):
from lancedb.background_loop import LOOP
class Inner:
def __init__(self):
self.live_version = 3
async def checkout(self, version):
await asyncio.sleep(0.01)
self.live_version = version
inner = Inner()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = 3
original_run = LOOP.run
def cancel_while_running(awaitable):
async def cancel_after_start():
task = asyncio.create_task(awaitable)
await asyncio.sleep(0)
task.cancel()
return await task
return original_run(cancel_after_start())
monkeypatch.setattr(LOOP, "run", cancel_while_running)
with pytest.raises(CancelledError):
table.checkout(11)
assert inner.live_version == 11
assert table._checkout_version == 11
def test_reopen_preserves_explicit_table_location(tmp_path):
db = lancedb.connect(tmp_path / "db")
location = str(tmp_path / "physical-table")
table = LanceTable.create(
db,
"items",
pa.table({"x": [1]}),
location=location,
)
table._pid = -1
table._ensure_open()
assert table.count_rows() == 1
assert table._location == location
def test_restore_with_tags(mem_db: DBConnection):
table = mem_db.create_table(
"my_table",
@@ -2537,20 +2196,6 @@ def test_update(mem_db: DBConnection):
assert np.allclose(v, np.array([[1.2, 1.9], [1.1, 1.1]]))
def test_update_with_arrow_scalar(mem_db: DBConnection):
schema = pa.schema({"id": pa.int64(), "vector": pa.list_(pa.float32(), 4)})
table = mem_db.create_table("my_table", schema=schema)
table.add([{"id": 1, "vector": [1.0, 2.0, 3.0, 4.0]}])
value = table.search().select(["vector"]).limit(1).to_arrow()["vector"][0]
assert isinstance(value, pa.FixedSizeListScalar)
result = table.update(where="id == 1", values={"vector": value})
assert result.rows_updated == 1
assert table.to_arrow()["vector"].to_pylist() == [[1.0, 2.0, 3.0, 4.0]]
def test_update_types(mem_db: DBConnection):
table = mem_db.create_table(
"my_table",
@@ -2718,55 +2363,6 @@ def test_merge_insert(mem_db: DBConnection):
)
def test_merge_insert_nullable_pandas_into_pydantic_schema(mem_db: DBConnection):
# Regression test for https://github.com/lancedb/lancedb/issues/2366
pd = pytest.importorskip("pandas")
class Document(LanceModel):
id: int
title: str
content: str
table = mem_db.create_table("documents", schema=Document)
table.add(
pd.DataFrame(
{
"title": ["Old title", "Unchanged"],
"id": [2, 3],
"content": ["Old content", "Keep this"],
}
)
)
# Pandas produces nullable Arrow fields, in an order that differs from the
# non-nullable Pydantic schema. This is valid as long as the data has no nulls.
new_data = pd.DataFrame(
{
"title": ["Inserted", "Updated"],
"id": [1, 2],
"content": ["New row", "New content"],
}
)
result = (
table.merge_insert("id")
.when_matched_update_all()
.when_not_matched_insert_all()
.execute(new_data)
)
assert result.num_inserted_rows == 1
assert result.num_updated_rows == 1
expected = pa.Table.from_pylist(
[
{"id": 1, "title": "Inserted", "content": "New row"},
{"id": 2, "title": "Updated", "content": "New content"},
{"id": 3, "title": "Unchanged", "content": "Keep this"},
],
schema=Document.to_arrow_schema(),
)
assert table.to_arrow().sort_by("id") == expected
def test_merge_insert_by_source_delete_expr(mem_db: DBConnection):
table = mem_db.create_table(
"my_table",
@@ -2867,36 +2463,6 @@ def test_merge_insert_subschema(mem_db: DBConnection, data_format):
assert table.to_arrow().sort_by("id") == expected
def test_repeated_partial_merge_insert_with_scalar_index(mem_db: DBConnection):
def make_batch(start: int) -> pa.Table:
return pa.table(
{
"id": [f"id-{i:04}" for i in range(start, start + 100)],
"category": ["A"] * 100,
"value_a": [float(i) for i in range(start, start + 100)],
"value_b": [float(i) / 10 for i in range(100)],
}
)
table = mem_db.create_table("my_table", data=make_batch(0))
table.add(make_batch(100))
table.add(make_batch(200))
table.create_index("id", config=BTree())
ids = [f"id-{i:04}" for i in range(100, 200)]
for value in (999.0, 888.0):
result = (
table.merge_insert("id")
.when_matched_update_all()
.execute(pa.table({"id": ids, "value_a": [value] * 100}))
)
assert result.num_updated_rows == 100
actual = table.to_arrow().sort_by("id")
assert actual.num_rows == 300
assert actual["value_a"].to_pylist()[100:200] == [888.0] * 100
@pytest.mark.asyncio
async def test_merge_insert_async(mem_db_async: AsyncConnection):
data = pa.table({"a": [1, 2, 3], "b": ["a", "b", "c"]})
@@ -3923,8 +3489,8 @@ def test_create_table_empty_list_no_schema_error(mem_db: DBConnection):
mem_db.create_table("test_empty_no_schema", data=[])
def test_create_table_without_data_with_vector_schema(tmp_path):
"""Test exact scenario from issue #1968.
def test_add_table_with_empty_embeddings(tmp_path):
"""Test exact scenario from issue #1968
Regression test for issue #1968:
https://github.com/lancedb/lancedb/issues/1968
@@ -3936,9 +3502,6 @@ def test_create_table_without_data_with_vector_schema(tmp_path):
embedding: Vector(16)
table = db.create_table("test", schema=MySchema)
assert table.count_rows() == 0
assert table.schema == MySchema.to_arrow_schema()
table.add(
[{"text": "bar", "embedding": [0.1] * 16}],
on_bad_vectors="drop",
-76
View File
@@ -342,42 +342,6 @@ def _multiworker_dataloader_target(db_uri: str, result_queue):
result_queue.put(count)
class _LazyPermutationDataset(torch.utils.data.Dataset):
"""Match applications that create their Permutation inside a fork worker."""
def __init__(self, table):
self._table = table
self._permutation = None
self._length = table.count_rows()
def __len__(self):
return self._length
def __getitems__(self, indices):
if self._permutation is None:
inherited_connection = self._table._conn
self._permutation = Permutation.identity(self._table)
if self._table._conn is inherited_connection:
raise RuntimeError("Permutation reused a connection inherited by fork")
return self._permutation.__getitems__(indices)
def _lazy_multiworker_dataloader_target(db_uri: str, result_queue):
table = lancedb.connect(db_uri).open_table("test_table")
dataset = _LazyPermutationDataset(table)
dataloader = torch.utils.data.DataLoader(
dataset,
batch_size=10,
num_workers=2,
multiprocessing_context="fork",
)
count = 0
for batch in dataloader:
assert batch["a"].size(0) == 10
count += 1
result_queue.put(count)
def _remote_multiworker_dataloader_target(port: int, result_queue):
import lancedb
from lancedb.permutation import Permutation
@@ -446,46 +410,6 @@ def test_permutation_dataloader_fork_workers(tmp_path):
assert queue.get() == 100
@pytest.mark.skipif(
sys.platform != "linux",
reason=(
"fork() is unavailable on Windows and unsafe on macOS "
"(Apple frameworks/TLS are not fork-safe)"
),
)
def test_lazy_permutation_reopens_inherited_table_in_fork_worker(tmp_path):
"""A lazily built Permutation must not reuse an inherited table client.
Object-store table handles contain HTTP connection pools that are unsafe
after fork. The local table makes the handle replacement deterministic
without requiring an S3 service in the unit-test environment.
"""
db_uri = str(tmp_path / "db")
db = lancedb.connect(db_uri)
db.create_table("test_table", pa.table({"a": list(range(1000))}))
ctx = mp.get_context("spawn")
queue = ctx.Queue()
proc = ctx.Process(
target=_lazy_multiworker_dataloader_target,
args=(db_uri, queue),
)
proc.start()
proc.join(timeout=30)
if proc.is_alive():
proc.terminate()
proc.join(timeout=5)
if proc.is_alive():
proc.kill()
proc.join()
pytest.fail("Lazy Permutation hung in a fork-based DataLoader worker")
assert proc.exitcode == 0, f"child exited with code {proc.exitcode}"
assert not queue.empty(), "child produced no batches"
assert queue.get() == 100
@pytest.mark.skipif(
sys.platform != "linux",
reason=(
@@ -75,22 +75,6 @@ class TestVoyageAIModelRegistration:
with pytest.raises(ValueError, match="not supported"):
func.ndims()
def test_voyage3_source_embeddings_use_text_api(self, mock_voyageai_client):
"""Regression test for text table data being sent to the multimodal API."""
mock_voyageai_client.tokenize.return_value = [["hello", "world"]]
mock_voyageai_client.embed.return_value.embeddings = [[0.1] * 1024]
registry = get_registry()
func = registry.get("voyageai").create(name="voyage-3")
embeddings = func.compute_source_embeddings("hello world")
assert embeddings == [[0.1] * 1024]
mock_voyageai_client.embed.assert_called_once_with(
texts=["hello world"], model="voyage-3", input_type="document"
)
mock_voyageai_client.multimodal_embed.assert_not_called()
@pytest.mark.parametrize(
"model_name",
[
-2
View File
@@ -75,8 +75,6 @@ reqwest = { version = "0.12.0", default-features = false, features = [
"http2",
"json",
"macos-system-configuration",
# Avoid linking OpenSSL into Python wheels, which breaks on FIPS hosts.
"rustls-tls-native-roots",
"stream",
], optional = true }
http = { version = "1", optional = true } # Matching what is in reqwest
@@ -202,17 +202,6 @@ mod tests {
assert_eq!(table.count_rows(None).await.unwrap(), 0);
}
#[tokio::test]
async fn create_table_in_named_memory_database() {
let db = connect("memory://foo").execute().await.unwrap();
let batch = record_batch!(("id", Int64, [1, 2, 3])).unwrap();
let table = db.create_table("my_table", batch).execute().await.unwrap();
assert_eq!(table.uri().await.unwrap(), "memory://foo/my_table.lance");
assert_eq!(table.count_rows(None).await.unwrap(), 3);
}
async fn test_create_table_with_data<T>(data: T)
where
T: Scannable + 'static,
-62
View File
@@ -1376,68 +1376,6 @@ mod tests {
assert!(!tempdir.path().join("__manifest").exists());
}
/// Regression test for https://github.com/lancedb/lancedb/issues/1600.
///
/// Opening a table used to create a separate object-store client instead of
/// reusing the one that successfully connected to the database. Repeating
/// credential discovery made S3 table opens intermittent, especially in AWS
/// Lambda, and the failed open was reported as `TableNotFound`.
#[tokio::test]
async fn test_open_table_reuses_connection_object_store() {
let tempdir = tempdir().unwrap();
let uri = tempdir.path().to_str().unwrap();
let registry = Arc::new(lance_io::object_store::ObjectStoreRegistry::default());
let session = Arc::new(lance::session::Session::new(16, 16, registry.clone()));
let request = ConnectRequest {
uri: uri.to_string(),
#[cfg(feature = "remote")]
client_config: Default::default(),
options: Default::default(),
namespace_client_properties: Default::default(),
manifest_enabled: false,
read_consistency_interval: None,
session: Some(session),
};
let db = ListingDatabase::connect_with_options(&request)
.await
.unwrap();
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
db.create_table(CreateTableRequest {
name: "test".to_string(),
namespace_path: vec![],
data: Box::new(RecordBatch::new_empty(schema)) as Box<dyn Scannable>,
mode: CreateTableMode::Create,
write_options: Default::default(),
location: None,
namespace_client: None,
})
.await
.unwrap();
let before_open = registry.stats();
for _ in 0..3 {
let table = db
.open_table(OpenTableRequest {
name: "test".to_string(),
namespace_path: vec![],
index_cache_size: None,
lance_read_params: None,
location: None,
namespace_client: None,
managed_versioning: None,
})
.await
.unwrap();
assert_eq!(table.count_rows(None).await.unwrap(), 0);
}
let after_open = registry.stats();
assert_eq!(after_open.misses, before_open.misses);
assert!(after_open.hits >= before_open.hits + 3);
}
#[tokio::test]
async fn test_clone_table_basic() {
let (_tempdir, db) = setup_database().await;
+4 -143
View File
@@ -132,14 +132,9 @@ impl ObjectStore for MirroringObjectStore {
if to.primary_only() {
self.primary.copy_opts(from, to, options).await
} else {
// The secondary store can be process-local and less durable than the
// primary, so a source written by another process may not exist here
// or may be evicted before the copy begins.
match self.secondary.copy_opts(from, to, options.clone()).await {
Ok(()) | Err(Error::NotFound { .. }) => {}
Err(err) => return Err(err),
}
self.primary.copy_opts(from, to, options).await
self.secondary.copy_opts(from, to, options.clone()).await?;
self.primary.copy_opts(from, to, options).await?;
Ok(())
}
}
}
@@ -197,8 +192,7 @@ mod test {
use futures::TryStreamExt;
use lance::{dataset::WriteParams, io::ObjectStoreParams};
use lance_testing::datagen::{BatchGenerator, IncrementingInt32, RandomVector};
use object_store::{local::LocalFileSystem, memory::InMemory};
use std::time::Duration;
use object_store::local::LocalFileSystem;
use tempfile;
use crate::{
@@ -207,139 +201,6 @@ mod test {
table::WriteOptions,
};
#[derive(Debug)]
struct EvictBeforeCopyStore {
inner: Arc<dyn ObjectStore>,
}
impl std::fmt::Display for EvictBeforeCopyStore {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "EvictBeforeCopyStore")
}
}
#[async_trait]
impl ObjectStore for EvictBeforeCopyStore {
async fn put_opts(
&self,
location: &Path,
payload: PutPayload,
options: PutOptions,
) -> Result<PutResult> {
self.inner.put_opts(location, payload, options).await
}
async fn put_multipart_opts(
&self,
location: &Path,
options: PutMultipartOptions,
) -> Result<Box<dyn MultipartUpload>> {
self.inner.put_multipart_opts(location, options).await
}
async fn get_opts(&self, location: &Path, options: GetOptions) -> Result<GetResult> {
self.inner.get_opts(location, options).await
}
fn delete_stream(
&self,
locations: BoxStream<'static, Result<Path>>,
) -> BoxStream<'static, Result<Path>> {
self.inner.delete_stream(locations)
}
fn list(&self, prefix: Option<&Path>) -> BoxStream<'static, Result<ObjectMeta>> {
self.inner.list(prefix)
}
async fn list_with_delimiter(&self, prefix: Option<&Path>) -> Result<ListResult> {
self.inner.list_with_delimiter(prefix).await
}
async fn copy_opts(&self, from: &Path, to: &Path, options: CopyOptions) -> Result<()> {
self.inner.delete(from).await?;
self.inner.copy_opts(from, to, options).await
}
}
#[tokio::test]
async fn test_copy_when_source_is_missing_from_secondary() {
let primary_dir = tempfile::tempdir().unwrap();
let secondary_dir = tempfile::tempdir().unwrap();
let primary: Arc<dyn ObjectStore> =
Arc::new(LocalFileSystem::new_with_prefix(primary_dir.path()).unwrap());
let secondary: Arc<dyn ObjectStore> =
Arc::new(LocalFileSystem::new_with_prefix(secondary_dir.path()).unwrap());
let store = MirroringObjectStore {
primary: primary.clone(),
secondary: secondary.clone(),
};
let staging = Path::from("_versions/1.manifest-staging");
let finalized = Path::from("_versions/1.manifest");
primary
.put(&staging, "manifest contents".into())
.await
.unwrap();
tokio::time::timeout(Duration::from_secs(5), store.copy(&staging, &finalized))
.await
.expect("copy should not hang when the secondary source is missing")
.unwrap();
let copied = primary
.get(&finalized)
.await
.unwrap()
.bytes()
.await
.unwrap();
assert_eq!(copied, "manifest contents");
assert!(matches!(
secondary.head(&finalized).await,
Err(Error::NotFound { .. })
));
}
#[tokio::test]
async fn test_copy_when_secondary_source_disappears_after_head() {
let primary: Arc<dyn ObjectStore> = Arc::new(InMemory::new());
let secondary_inner: Arc<dyn ObjectStore> = Arc::new(InMemory::new());
let secondary: Arc<dyn ObjectStore> = Arc::new(EvictBeforeCopyStore {
inner: secondary_inner.clone(),
});
let store = MirroringObjectStore {
primary: primary.clone(),
secondary,
};
let staging = Path::from("_versions/1.manifest-staging");
let finalized = Path::from("_versions/1.manifest");
primary
.put(&staging, "manifest contents".into())
.await
.unwrap();
secondary_inner
.put(&staging, "manifest contents".into())
.await
.unwrap();
store.copy(&staging, &finalized).await.unwrap();
let copied = primary
.get(&finalized)
.await
.unwrap()
.bytes()
.await
.unwrap();
assert_eq!(copied, "manifest contents");
assert!(matches!(
secondary_inner.head(&finalized).await,
Err(Error::NotFound { .. })
));
}
// This test is ignored because lance 3.0 introduced LocalWriter optimization
// that bypasses the object store wrapper for local writes. The mirroring feature
// still works for remote/cloud storage, but can't be tested with local storage.
+28 -4
View File
@@ -1661,8 +1661,14 @@ mod tests {
#[tokio::test]
async fn test_setters_getters() {
// TODO: Switch back to memory://foo after https://github.com/lancedb/lancedb/issues/1051
// is fixed
let tmp_dir = tempdir().unwrap();
let dataset_path = tmp_dir.path().join("test.lance");
let uri = dataset_path.to_str().unwrap();
let batches = make_test_batches();
let conn = connect("memory://foo").execute().await.unwrap();
let conn = connect(uri).execute().await.unwrap();
let table = conn
.create_table("my_table", batches)
.execute()
@@ -1757,8 +1763,14 @@ mod tests {
#[tokio::test]
async fn test_execute() {
// TODO: Switch back to memory://foo after https://github.com/lancedb/lancedb/issues/1051
// is fixed
let tmp_dir = tempdir().unwrap();
let dataset_path = tmp_dir.path().join("test.lance");
let uri = dataset_path.to_str().unwrap();
let batches = make_non_empty_batches();
let conn = connect("memory://foo").execute().await.unwrap();
let conn = connect(uri).execute().await.unwrap();
let table = conn
.create_table("my_table", batches)
.execute()
@@ -1877,8 +1889,14 @@ mod tests {
#[tokio::test]
async fn test_select_with_transform() {
// TODO: Switch back to memory://foo after https://github.com/lancedb/lancedb/issues/1051
// is fixed
let tmp_dir = tempdir().unwrap();
let dataset_path = tmp_dir.path().join("test.lance");
let uri = dataset_path.to_str().unwrap();
let batches = make_non_empty_batches();
let conn = connect("memory://foo").execute().await.unwrap();
let conn = connect(uri).execute().await.unwrap();
let table = conn
.create_table("my_table", batches)
.execute()
@@ -1975,9 +1993,15 @@ mod tests {
#[tokio::test]
async fn test_execute_no_vector() {
// TODO: Switch back to memory://foo after https://github.com/lancedb/lancedb/issues/1051
// is fixed
let tmp_dir = tempdir().unwrap();
let dataset_path = tmp_dir.path().join("test.lance");
let uri = dataset_path.to_str().unwrap();
// test that it's ok to not specify a query vector (just filter / limit)
let batches = make_non_empty_batches();
let conn = connect("memory://foo").execute().await.unwrap();
let conn = connect(uri).execute().await.unwrap();
let table = conn
.create_table("my_table", batches)
.execute()
+1 -59
View File
@@ -373,37 +373,6 @@ pub fn parse_db_url(db_url: &str) -> Result<ParsedDbUrl> {
Ok(ParsedDbUrl { db_name, db_prefix })
}
fn validate_dns_hostname(hostname: &str) -> Result<()> {
let ascii_hostname = match url::Host::parse(hostname) {
Ok(url::Host::Domain(hostname)) => hostname,
Ok(_) => {
return Err(Error::InvalidInput {
message: "LanceDB Cloud database URI or region produced a non-DNS hostname"
.to_string(),
});
}
Err(err) => {
return Err(Error::InvalidInput {
message: format!(
"LanceDB Cloud database URI or region produced an invalid hostname: {err}"
),
});
}
};
if ascii_hostname.len() > 253
|| ascii_hostname
.split('.')
.any(|label| label.is_empty() || label.len() > 63)
{
return Err(Error::InvalidInput {
message: "LanceDB Cloud database URI or region produced an invalid hostname: DNS labels must contain 1 to 63 bytes and the full hostname must not exceed 253 bytes".to_string(),
});
}
Ok(())
}
impl RestfulLanceDbClient<Sender> {
fn get_timeout(passed: Option<Duration>, env_var: &str) -> Result<Option<Duration>> {
if let Some(passed) = passed {
@@ -511,11 +480,7 @@ impl RestfulLanceDbClient<Sender> {
let host = match host_override {
Some(host_override) => host_override,
None => {
let hostname = format!("{}.{}.api.lancedb.com", parsed_url.db_name, region);
validate_dns_hostname(&hostname)?;
format!("https://{hostname}")
}
None => format!("https://{}.{}.api.lancedb.com", parsed_url.db_name, region),
};
debug!("Created client for host: {}", host);
let retry_config = client_config.retry_config.clone().try_into()?;
@@ -1192,29 +1157,6 @@ mod tests {
assert_eq!(headers.get("x-api-key").unwrap(), "api-key");
}
#[test]
fn test_rejects_invalid_cloud_dns_hostname() {
let invalid_database_names = ["a".repeat(64), "invalid..database".to_string()];
for db_name in invalid_database_names {
let parsed_url = parse_db_url(&format!("db://{db_name}")).unwrap();
let error = RestfulLanceDbClient::<Sender>::try_new(
&parsed_url,
"us-east-1",
None,
HeaderMap::new(),
ClientConfig::default(),
None,
)
.unwrap_err();
assert!(
matches!(error, Error::InvalidInput { ref message } if message.contains("DNS labels must contain 1 to 63 bytes")),
"unexpected error: {error}"
);
}
}
// Test implementation of HeaderProvider
#[derive(Debug, Clone)]
struct TestHeaderProvider {
+6 -44
View File
@@ -2791,10 +2791,9 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
}
async fn index_stats(&self, index_name: &str) -> Result<Option<IndexStatistics>> {
let encoded_name = urlencoding::encode(index_name);
let mut request = self.post_read(&format!(
"/v1/table/{}/index/{encoded_name}/stats/",
self.identifier
"/v1/table/{}/index/{}/stats/",
self.identifier, index_name
));
let version = self.current_version().await;
let mut body = serde_json::json!({ "version": version });
@@ -2821,10 +2820,9 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
}
async fn drop_index(&self, index_name: &str) -> Result<()> {
let encoded_name = urlencoding::encode(index_name);
let request = self.apply_branch_query(self.client.post(&format!(
"/v1/table/{}/index/{encoded_name}/drop/",
self.identifier
"/v1/table/{}/index/{}/drop/",
self.identifier, index_name
)));
let (request_id, response) = self.send(request, true).await?;
if response.status() == StatusCode::NOT_FOUND {
@@ -2837,10 +2835,9 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
}
async fn prewarm_index(&self, index_name: &str) -> Result<()> {
let encoded_name = urlencoding::encode(index_name);
let request = self.client.post(&format!(
"/v1/table/{}/index/{encoded_name}/prewarm/",
self.identifier
"/v1/table/{}/index/{}/prewarm/",
self.identifier, index_name
));
let (request_id, response) = self.send(request, true).await?;
if response.status() == StatusCode::NOT_FOUND {
@@ -6492,41 +6489,6 @@ mod tests {
assert!(matches!(e, Error::IndexNotFound { .. }));
}
/// Index names are unvalidated, so reserved characters must be
/// percent-encoded or they restructure the request path.
#[tokio::test]
async fn test_per_index_paths_encode_reserved_characters() {
const NAME: &str = "my/index?a#b c";
const PREFIX: &str = "/v1/table/my_table/index/my%2Findex%3Fa%23b%20c";
let table = Table::new_with_handler("my_table", |request| {
assert_eq!(request.url().path(), format!("{PREFIX}/stats/"));
let body = serde_json::json!({
"num_indexed_rows": 1,
"num_unindexed_rows": 0,
"index_type": "IVF_PQ",
"distance_type": "l2"
});
http::Response::builder()
.status(200)
.body(serde_json::to_string(&body).unwrap())
.unwrap()
});
assert!(table.index_stats(NAME).await.unwrap().is_some());
let table = Table::new_with_handler("my_table", |request| {
assert_eq!(request.url().path(), format!("{PREFIX}/drop/"));
http::Response::builder().status(200).body("{}").unwrap()
});
table.drop_index(NAME).await.unwrap();
let table = Table::new_with_handler("my_table", |request| {
assert_eq!(request.url().path(), format!("{PREFIX}/prewarm/"));
http::Response::builder().status(200).body("{}").unwrap()
});
table.prewarm_index(NAME).await.unwrap();
}
#[tokio::test]
async fn test_set_lsm_write_spec_unsharded() {
let table = Table::new_with_handler("my_table", |request| {
+1 -70
View File
@@ -315,10 +315,7 @@ pub(crate) async fn execute_merge_insert(
#[cfg(test)]
mod tests {
use arrow_array::builder::FixedSizeBinaryBuilder;
use arrow_array::{
Int32Array, RecordBatch, RecordBatchIterator, RecordBatchReader, StringArray, UInt64Array,
};
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator, RecordBatchReader};
use arrow_schema::{DataType, Field, Schema};
use std::sync::Arc;
@@ -340,42 +337,6 @@ mod tests {
Box::new(RecordBatchIterator::new(vec![Ok(batch)], schema))
}
fn fixed_size_binary_merge_batch(
id_range: std::ops::Range<u64>,
price: u64,
) -> Box<dyn RecordBatchReader + Send> {
let ids = id_range.collect::<Vec<_>>();
let mut id_builder = FixedSizeBinaryBuilder::new(16);
for id in &ids {
let mut bytes = [0; 16];
bytes[..8].copy_from_slice(&id.to_le_bytes());
id_builder.append_value(bytes).unwrap();
}
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::FixedSizeBinary(16), false),
Field::new("id_as_int", DataType::UInt64, false),
Field::new("name", DataType::Utf8, false),
Field::new("market", DataType::Utf8, false),
]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(id_builder.finish()),
Arc::new(UInt64Array::from_iter_values(ids.iter().copied())),
Arc::new(StringArray::from_iter_values(
ids.iter().map(|id| format!("name{id}")),
)),
Arc::new(StringArray::from_iter_values(std::iter::repeat_n(
format!("market_{price}"),
ids.len(),
))),
],
)
.unwrap();
Box::new(RecordBatchIterator::new(vec![Ok(batch)], schema))
}
#[tokio::test]
async fn test_merge_insert() {
let conn = connect("memory://").execute().await.unwrap();
@@ -427,36 +388,6 @@ mod tests {
);
}
#[tokio::test]
async fn test_merge_insert_fixed_size_binary_non_nullable() {
// Regression test for #2869: an unrelated FixedSizeBinary column used to corrupt the
// outer join that implements when_not_matched_by_source_delete.
let conn = connect("memory://").execute().await.unwrap();
let table = conn
.create_table(
"fixed_size_binary_merge",
fixed_size_binary_merge_batch(0..256, 100),
)
.execute()
.await
.unwrap();
let mut merge_insert = table.merge_insert(&["id_as_int"]);
merge_insert
.when_matched_update_all(None)
.when_not_matched_insert_all()
.when_not_matched_by_source_delete(None);
let result = merge_insert
.execute(fixed_size_binary_merge_batch(100..356, 200))
.await
.unwrap();
assert_eq!(result.num_updated_rows, 156);
assert_eq!(result.num_inserted_rows, 100);
assert_eq!(result.num_deleted_rows, 100);
assert_eq!(table.count_rows(None).await.unwrap(), 256);
}
#[tokio::test]
async fn test_merge_insert_use_index() {
let conn = connect("memory://").execute().await.unwrap();
+1 -148
View File
@@ -214,17 +214,12 @@ pub(crate) async fn execute_optimize(
#[cfg(test)]
mod tests {
use arrow_array::{
Array, FixedSizeListArray, Float32Array, Int32Array, RecordBatch, StringArray,
};
use arrow_array::{Int32Array, RecordBatch, StringArray};
use arrow_schema::{DataType, Field, Schema};
use lance_arrow::FixedSizeListArrayExt;
use rstest::rstest;
use std::sync::Arc;
use crate::connect;
use crate::database::listing::OPT_NEW_TABLE_ENABLE_STABLE_ROW_IDS;
use crate::index::vector::IvfRqIndexBuilder;
use crate::index::{Index, scalar::BTreeIndexBuilder};
use crate::query::ExecutableQuery;
use crate::table::{CompactionOptions, OptimizeAction, OptimizeStats};
@@ -309,96 +304,6 @@ mod tests {
assert_eq!(all_values, expected);
}
#[tokio::test]
async fn test_compact_with_concurrent_add() {
const NUM_FRAGMENTS: usize = 5;
const ROWS_PER_FRAGMENT: i32 = 300;
let tmpdir = tempfile::tempdir().unwrap();
let conn = connect(tmpdir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let schema = Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]));
let batch = RecordBatch::try_new(
schema,
vec![Arc::new(Int32Array::from_iter_values(0..ROWS_PER_FRAGMENT))],
)
.unwrap();
let table = conn
.create_table("test_concurrent_compact", batch.clone())
.execute()
.await
.unwrap();
table
.create_index(&["id"], Index::BTree(BTreeIndexBuilder::default()))
.execute()
.await
.unwrap();
for _ in 0..NUM_FRAGMENTS {
table.add(batch.clone()).execute().await.unwrap();
}
// Use separate handles so the two writes actually overlap, as they can
// when different Node connections operate on the same S3 table.
let compact_table = conn
.open_table("test_concurrent_compact")
.execute()
.await
.unwrap();
let append_table = conn
.open_table("test_concurrent_compact")
.execute()
.await
.unwrap();
let compact_task = tokio::spawn(async move {
compact_table
.optimize(OptimizeAction::Compact {
options: CompactionOptions {
target_rows_per_fragment: 1_000,
..Default::default()
},
remap_options: None,
})
.await
});
tokio::task::yield_now().await;
for _ in 0..NUM_FRAGMENTS {
append_table.add(batch.clone()).execute().await.unwrap();
}
compact_task.await.unwrap().unwrap();
let table = conn
.open_table("test_concurrent_compact")
.execute()
.await
.unwrap();
let dataset = table.dataset().unwrap().get().await.unwrap();
let fragment_ids = dataset
.get_fragments()
.iter()
.map(|fragment| fragment.id())
.collect::<Vec<_>>();
assert!(fragment_ids.windows(2).all(|ids| ids[0] < ids[1]));
// A second compaction exposed the original out-of-order row-id bug.
table
.optimize(OptimizeAction::Compact {
options: CompactionOptions {
target_rows_per_fragment: 1_000,
..Default::default()
},
remap_options: None,
})
.await
.unwrap();
assert_eq!(
table.count_rows(None).await.unwrap(),
ROWS_PER_FRAGMENT as usize * (NUM_FRAGMENTS * 2 + 1)
);
}
#[tokio::test]
async fn test_optimize_prune_versions() {
let conn = connect("memory://").execute().await.unwrap();
@@ -537,58 +442,6 @@ mod tests {
assert_eq!(final_row_count, 200);
}
#[tokio::test]
async fn test_optimize_vector_index_after_delete_with_stable_row_ids() {
const NUM_ROWS: i32 = 400;
const DIMENSION: i32 = 32;
let conn = connect("memory://").execute().await.unwrap();
let vectors = FixedSizeListArray::try_new_from_values(
Float32Array::from_iter_values((0..NUM_ROWS).flat_map(|id| {
(0..DIMENSION).map(move |offset| ((id as f32 * 0.1) + (offset as f32 * 0.3)).sin())
})),
DIMENSION,
)
.unwrap();
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new("vector", vectors.data_type().clone(), false),
]));
let batch = RecordBatch::try_new(
schema,
vec![
Arc::new(Int32Array::from_iter_values(0..NUM_ROWS)),
Arc::new(vectors),
],
)
.unwrap();
let table = conn
.create_table("test_vector_index_optimize_after_delete", batch)
.storage_option(OPT_NEW_TABLE_ENABLE_STABLE_ROW_IDS, "true")
.execute()
.await
.unwrap();
table
.create_index(
&["vector"],
Index::IvfRq(IvfRqIndexBuilder::default().num_partitions(4)),
)
.execute()
.await
.unwrap();
table.delete("id % 3 = 0").await.unwrap();
// Regression test for #3330: deleted stable row IDs used to become
// misaligned with row addresses while joining small IVF partitions.
table
.optimize(OptimizeAction::Index(Default::default()))
.await
.unwrap();
assert_eq!(table.count_rows(None).await.unwrap(), 266);
}
#[tokio::test]
async fn test_optimize_all() {
let conn = connect("memory://").execute().await.unwrap();