mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-18 12:08:35 +00:00
Merge remote-tracking branch 'origin/main' into gatekeeper/fix-1869-1
# Conflicts: # python/python/tests/test_table.py
This commit is contained in:
@@ -197,6 +197,35 @@ 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),
|
||||
|
||||
@@ -170,6 +170,38 @@ 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) => {
|
||||
|
||||
@@ -86,6 +86,44 @@ 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
@@ -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-py39", "chrono"] }
|
||||
pyo3 = { version = "0.28", features = ["extension-module", "abi3-py310", "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-py39",
|
||||
"abi3-py310",
|
||||
] }
|
||||
|
||||
[features]
|
||||
|
||||
@@ -101,8 +101,7 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
|
||||
|
||||
@weak_lru(maxsize=1)
|
||||
def ndims(self):
|
||||
model = self.get_model()
|
||||
return model.encode("foo").shape[0]
|
||||
return len(self.generate_embeddings([[self.source_instruction, "foo"]])[0])
|
||||
|
||||
def compute_query_embeddings(self, query: str, *args, **kwargs) -> List[np.array]:
|
||||
return self.generate_embeddings([[self.query_instruction, query]])
|
||||
|
||||
@@ -395,6 +395,11 @@ 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,6 +2,7 @@
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
|
||||
import inspect
|
||||
import re
|
||||
import sys
|
||||
from datetime import timedelta
|
||||
@@ -62,17 +63,23 @@ def test_basic(tmp_path):
|
||||
assert db.open_table("test").name == db["test"].name
|
||||
|
||||
|
||||
def test_sync_repr_does_not_use_background_loop(tmp_path, monkeypatch):
|
||||
def test_sync_debugger_inspection_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("repr should not use the Python background loop")
|
||||
raise AssertionError("debugger inspection should not use the 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})"
|
||||
|
||||
|
||||
@@ -64,6 +64,23 @@ 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):
|
||||
@@ -115,34 +132,16 @@ def test_embedding_function_variables():
|
||||
assert func.safe_model_dump()["secret_key"] == "$var:secret"
|
||||
|
||||
|
||||
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]
|
||||
|
||||
def test_openai_variables_survive_metadata_round_trip():
|
||||
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("variable-parsing-test").create(
|
||||
api_key="$var:test_api_key", base_url="$var:test_base_url"
|
||||
function=registry.get("openai").create(
|
||||
api_key="$var:test_api_key", base_url="https://api.example.com"
|
||||
),
|
||||
)
|
||||
|
||||
@@ -150,7 +149,10 @@ def test_parse_functions_with_variables():
|
||||
|
||||
# Create a mock arrow table with the metadata
|
||||
schema = pa.schema(
|
||||
[pa.field("text", pa.string()), pa.field("vector", pa.list_(pa.float32(), 10))]
|
||||
[
|
||||
pa.field("text", pa.string()),
|
||||
pa.field("vector", pa.list_(pa.float32(), 1536)),
|
||||
]
|
||||
)
|
||||
table = pa.table({"text": [], "vector": []}, schema=schema)
|
||||
table = table.replace_schema_metadata(metadata)
|
||||
@@ -164,13 +166,15 @@ def test_parse_functions_with_variables():
|
||||
|
||||
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")
|
||||
|
||||
@@ -12,7 +12,7 @@ import pyarrow.compute as pc
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from lancedb.index import FTS
|
||||
from lancedb.index import BTree, FTS, IvfPq
|
||||
from lancedb.table import AsyncTable, Table
|
||||
|
||||
|
||||
@@ -99,6 +99,86 @@ 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
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
# 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}"
|
||||
)
|
||||
@@ -372,6 +372,31 @@ 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
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
# 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"
|
||||
@@ -35,6 +35,12 @@ 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(
|
||||
|
||||
@@ -2,10 +2,13 @@
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
|
||||
import ctypes
|
||||
import gc
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import warnings
|
||||
import weakref
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import date, datetime, timedelta
|
||||
from decimal import Decimal
|
||||
@@ -100,6 +103,30 @@ 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"]})
|
||||
@@ -451,6 +478,38 @@ 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)
|
||||
@@ -1841,6 +1900,33 @@ 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
|
||||
@@ -2281,6 +2367,20 @@ def test_update_expr_filter_preserves_typed_semantics(mem_db: DBConnection):
|
||||
assert result.rows_updated == 2
|
||||
|
||||
|
||||
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",
|
||||
@@ -2448,6 +2548,55 @@ 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",
|
||||
@@ -2548,6 +2697,36 @@ 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"]})
|
||||
@@ -3574,8 +3753,8 @@ def test_create_table_empty_list_no_schema_error(mem_db: DBConnection):
|
||||
mem_db.create_table("test_empty_no_schema", data=[])
|
||||
|
||||
|
||||
def test_add_table_with_empty_embeddings(tmp_path):
|
||||
"""Test exact scenario from issue #1968
|
||||
def test_create_table_without_data_with_vector_schema(tmp_path):
|
||||
"""Test exact scenario from issue #1968.
|
||||
|
||||
Regression test for issue #1968:
|
||||
https://github.com/lancedb/lancedb/issues/1968
|
||||
@@ -3587,6 +3766,9 @@ def test_add_table_with_empty_embeddings(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",
|
||||
|
||||
@@ -75,6 +75,22 @@ 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",
|
||||
[
|
||||
|
||||
@@ -75,6 +75,8 @@ 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,6 +202,17 @@ 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,
|
||||
|
||||
@@ -1376,6 +1376,68 @@ 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;
|
||||
|
||||
@@ -132,9 +132,14 @@ impl ObjectStore for MirroringObjectStore {
|
||||
if to.primary_only() {
|
||||
self.primary.copy_opts(from, to, options).await
|
||||
} else {
|
||||
self.secondary.copy_opts(from, to, options.clone()).await?;
|
||||
self.primary.copy_opts(from, to, options).await?;
|
||||
Ok(())
|
||||
// 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
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -192,7 +197,8 @@ mod test {
|
||||
use futures::TryStreamExt;
|
||||
use lance::{dataset::WriteParams, io::ObjectStoreParams};
|
||||
use lance_testing::datagen::{BatchGenerator, IncrementingInt32, RandomVector};
|
||||
use object_store::local::LocalFileSystem;
|
||||
use object_store::{local::LocalFileSystem, memory::InMemory};
|
||||
use std::time::Duration;
|
||||
use tempfile;
|
||||
|
||||
use crate::{
|
||||
@@ -201,6 +207,139 @@ 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.
|
||||
|
||||
@@ -1661,14 +1661,8 @@ 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(uri).execute().await.unwrap();
|
||||
let conn = connect("memory://foo").execute().await.unwrap();
|
||||
let table = conn
|
||||
.create_table("my_table", batches)
|
||||
.execute()
|
||||
@@ -1763,14 +1757,8 @@ 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(uri).execute().await.unwrap();
|
||||
let conn = connect("memory://foo").execute().await.unwrap();
|
||||
let table = conn
|
||||
.create_table("my_table", batches)
|
||||
.execute()
|
||||
@@ -1889,14 +1877,8 @@ 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(uri).execute().await.unwrap();
|
||||
let conn = connect("memory://foo").execute().await.unwrap();
|
||||
let table = conn
|
||||
.create_table("my_table", batches)
|
||||
.execute()
|
||||
@@ -1993,15 +1975,9 @@ 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(uri).execute().await.unwrap();
|
||||
let conn = connect("memory://foo").execute().await.unwrap();
|
||||
let table = conn
|
||||
.create_table("my_table", batches)
|
||||
.execute()
|
||||
|
||||
@@ -373,6 +373,37 @@ 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 {
|
||||
@@ -480,7 +511,11 @@ impl RestfulLanceDbClient<Sender> {
|
||||
|
||||
let host = match host_override {
|
||||
Some(host_override) => host_override,
|
||||
None => format!("https://{}.{}.api.lancedb.com", parsed_url.db_name, region),
|
||||
None => {
|
||||
let hostname = format!("{}.{}.api.lancedb.com", parsed_url.db_name, region);
|
||||
validate_dns_hostname(&hostname)?;
|
||||
format!("https://{hostname}")
|
||||
}
|
||||
};
|
||||
debug!("Created client for host: {}", host);
|
||||
let retry_config = client_config.retry_config.clone().try_into()?;
|
||||
@@ -1157,6 +1192,29 @@ 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 {
|
||||
|
||||
@@ -2791,9 +2791,10 @@ 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/{}/stats/",
|
||||
self.identifier, index_name
|
||||
"/v1/table/{}/index/{encoded_name}/stats/",
|
||||
self.identifier
|
||||
));
|
||||
let version = self.current_version().await;
|
||||
let mut body = serde_json::json!({ "version": version });
|
||||
@@ -2820,9 +2821,10 @@ 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/{}/drop/",
|
||||
self.identifier, index_name
|
||||
"/v1/table/{}/index/{encoded_name}/drop/",
|
||||
self.identifier
|
||||
)));
|
||||
let (request_id, response) = self.send(request, true).await?;
|
||||
if response.status() == StatusCode::NOT_FOUND {
|
||||
@@ -2835,9 +2837,10 @@ 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/{}/prewarm/",
|
||||
self.identifier, index_name
|
||||
"/v1/table/{}/index/{encoded_name}/prewarm/",
|
||||
self.identifier
|
||||
));
|
||||
let (request_id, response) = self.send(request, true).await?;
|
||||
if response.status() == StatusCode::NOT_FOUND {
|
||||
@@ -6489,6 +6492,41 @@ 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| {
|
||||
|
||||
@@ -315,7 +315,10 @@ pub(crate) async fn execute_merge_insert(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator, RecordBatchReader};
|
||||
use arrow_array::builder::FixedSizeBinaryBuilder;
|
||||
use arrow_array::{
|
||||
Int32Array, RecordBatch, RecordBatchIterator, RecordBatchReader, StringArray, UInt64Array,
|
||||
};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -337,6 +340,42 @@ 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();
|
||||
@@ -388,6 +427,36 @@ 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();
|
||||
|
||||
@@ -214,12 +214,17 @@ pub(crate) async fn execute_optimize(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use arrow_array::{Int32Array, RecordBatch, StringArray};
|
||||
use arrow_array::{
|
||||
Array, FixedSizeListArray, Float32Array, 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};
|
||||
@@ -304,6 +309,96 @@ 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();
|
||||
@@ -442,6 +537,58 @@ 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();
|
||||
|
||||
Reference in New Issue
Block a user