mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-28 00:48:40 +00:00
Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| bc3837c4fe | |||
| 2a4f4f338b |
@@ -296,18 +296,16 @@ jobs:
|
||||
cargo update -p aws-types --precise 1.3.9
|
||||
cargo update -p aws-sigv4 --precise 1.3.5
|
||||
cargo update -p aws-credential-types --precise 1.2.8
|
||||
# aws-smithy-checksums must stay at or above 0.63.13: OpenDAL's S3
|
||||
# service needs crc-fast ~1.9, and older releases pin it to ~1.3.
|
||||
cargo update -p aws-smithy-checksums --precise 0.63.13
|
||||
cargo update -p aws-smithy-checksums --precise 0.63.9
|
||||
cargo update -p aws-smithy-runtime --precise 1.9.3
|
||||
cargo update -p aws-smithy-http --precise 0.62.6
|
||||
cargo update -p aws-smithy-eventstream --precise 0.60.14
|
||||
cargo update -p aws-smithy-http --precise 0.62.4
|
||||
cargo update -p aws-smithy-eventstream --precise 0.60.12
|
||||
cargo update -p aws-smithy-http-client --precise 1.1.3
|
||||
cargo update -p aws-smithy-observability --precise 0.1.4
|
||||
cargo update -p aws-smithy-query --precise 0.60.8
|
||||
cargo update -p aws-smithy-runtime-api --precise 1.9.3
|
||||
cargo update -p aws-smithy-async --precise 1.2.7
|
||||
cargo update -p aws-smithy-types --precise 1.3.6
|
||||
cargo update -p aws-smithy-runtime-api --precise 1.9.1
|
||||
cargo update -p aws-smithy-async --precise 1.2.6
|
||||
cargo update -p aws-smithy-types --precise 1.3.5
|
||||
cargo update -p aws-smithy-xml --precise 0.60.11
|
||||
cargo update -p home --precise 0.5.9
|
||||
- name: cargo +${{ matrix.msrv }} check
|
||||
|
||||
@@ -152,54 +152,3 @@ Please consider the following when reviewing code contributions.
|
||||
### Documentation
|
||||
* New features must include updates to the rust documentation comments. Link to
|
||||
relevant structs and methods to increase the value of documentation.
|
||||
|
||||
## Cursor Cloud specific instructions
|
||||
|
||||
The VM snapshot already has the Rust `1.97.0` toolchain (auto-selected by
|
||||
`rust-toolchain.toml`), `protoc`, `uv` (on `PATH` via `~/.bashrc`), the Rust
|
||||
debug build artifacts, the Python editable extension, and `nodejs/node_modules`.
|
||||
The startup update script only refreshes dependencies (`uv sync` for Python and
|
||||
`pnpm install` for Node); it deliberately does NOT rebuild the native
|
||||
extensions. After changing Rust or PyO3/napi binding code you must rebuild the
|
||||
affected binding yourself (see per-binding rebuild commands below).
|
||||
|
||||
Non-obvious caveats discovered during setup:
|
||||
|
||||
* The documented Python bootstrap `uv run --extra tests --extra dev maturin
|
||||
develop --extras tests,dev` does not work as-is here: `maturin` is not
|
||||
installed as a CLI in the uv environment, and `maturin develop --extras`
|
||||
runs its own dependency resolution that cannot find the prerelease
|
||||
`pylance==9.0.0rc1` (it lacks the extra package index that `uv` uses via
|
||||
`uv.lock`). Because `uv run --extra tests --extra dev` already installs those
|
||||
extras, the working command is:
|
||||
`cd python && uv run --extra tests --extra dev --with maturin maturin develop`
|
||||
(note: `--with maturin`, and no `--extras`). This is the Python binding
|
||||
rebuild command.
|
||||
* Rust core, the Python extension (maturin), and the Node addon (napi) all
|
||||
compile into the SHARED `/workspace/target`. Cargo feature unification differs
|
||||
between `maturin develop` and `pnpm build`, so alternating between building
|
||||
the Python and Node bindings forces a full recompile of shared crates
|
||||
(`lancedb`, `datafusion`, `lance-*`) — roughly 6-7 min each way on this
|
||||
4-core VM. Build one binding at a time to avoid the churn.
|
||||
* The `_lancedb` release build (triggered when `uv run`/`uv sync` installs the
|
||||
`lancedb` project itself) uses `lto = "fat"` + `opt-level = 3`, needs ~11 GB
|
||||
RAM, and takes ~20 min cold on this VM. To avoid it, the update script uses
|
||||
`uv sync --no-install-project --inexact` (the `--inexact` flag is required so
|
||||
the sync does not uninstall the editable extension). Prefer the debug
|
||||
`maturin develop` (~6 min cold, seconds when warm) for iteration.
|
||||
* `cargo check` only produces metadata, so the first `cargo run --example ...`
|
||||
or `cargo test` after a check triggers a large codegen/link compile.
|
||||
* Node binding rebuild: `cd nodejs && pnpm build` (napi debug build + `tsc`).
|
||||
The native addon lands at `nodejs/dist/lancedb.linux-x64-gnu.node`.
|
||||
|
||||
Verified working (local backend, no cloud credentials needed):
|
||||
|
||||
* Rust: `cargo check/clippy --features remote --tests --examples`,
|
||||
`cargo test --features remote -p lancedb --lib`, `cargo run --features remote
|
||||
--example simple`.
|
||||
* Python: `cd python && uv run --extra tests pytest python/tests/test_table.py`,
|
||||
`uv run --directory python --extra dev ruff check python`.
|
||||
* Node: `cd nodejs && pnpm lint`, `pnpm test __test__/connection.test.ts`.
|
||||
|
||||
Java (`java/`) is optional; its integration tests need LanceDB Cloud
|
||||
credentials (`LANCEDB_DB`, `LANCEDB_API_KEY`) and were not set up here.
|
||||
|
||||
Generated
+233
-257
File diff suppressed because it is too large
Load Diff
+14
-14
@@ -13,20 +13,20 @@ categories = ["database-implementations"]
|
||||
rust-version = "1.91.0"
|
||||
|
||||
[workspace.dependencies]
|
||||
lance = { "version" = "=11.0.0-beta.2", default-features = false, "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { "version" = "=11.0.0-beta.2", default-features = false, "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { "version" = "=11.0.0-beta.2", default-features = false, "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { "version" = "=11.0.0-beta.2", "tag" = "v11.0.0-beta.2", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance = { "version" = "=10.1.0-beta.1", default-features = false, "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-core = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datagen = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-file = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-io = { "version" = "=10.1.0-beta.1", default-features = false, "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-index = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-linalg = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-namespace-impls = { "version" = "=10.1.0-beta.1", default-features = false, "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-table = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-testing = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-datafusion = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-encoding = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
lance-arrow = { "version" = "=10.1.0-beta.1", "tag" = "v10.1.0-beta.1", "git" = "https://github.com/lance-format/lance.git" }
|
||||
ahash = "0.8"
|
||||
# Note that this one does not include pyarrow
|
||||
arrow = { version = "58.0.0", optional = false }
|
||||
|
||||
+1
-1
@@ -28,7 +28,7 @@
|
||||
<properties>
|
||||
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
|
||||
<arrow.version>15.0.0</arrow.version>
|
||||
<lance-core.version>11.0.0-beta.2</lance-core.version>
|
||||
<lance-core.version>10.1.0-beta.1</lance-core.version>
|
||||
<spotless.skip>false</spotless.skip>
|
||||
<spotless.version>2.30.0</spotless.version>
|
||||
<spotless.java.googlejavaformat.version>1.7</spotless.java.googlejavaformat.version>
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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) => {
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -261,6 +261,7 @@ class Table:
|
||||
def name(self) -> str: ...
|
||||
def __repr__(self) -> str: ...
|
||||
def is_open(self) -> bool: ...
|
||||
def _is_native(self) -> bool: ...
|
||||
def close(self) -> None: ...
|
||||
async def schema(self) -> pa.Schema: ...
|
||||
async def add(
|
||||
|
||||
@@ -101,7 +101,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]])
|
||||
|
||||
@@ -314,8 +314,7 @@ class HnswPq:
|
||||
m: int = 20
|
||||
ef_construction: int = 300
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -422,8 +421,7 @@ class HnswSq:
|
||||
m: int = 20
|
||||
ef_construction: int = 300
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -618,8 +616,7 @@ class IvfFlat:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -651,8 +648,7 @@ class IvfSq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
@@ -784,7 +780,7 @@ class IvfPq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# Name of the accelerator ("cuda" or "mps") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
@@ -840,8 +836,7 @@ class IvfRq:
|
||||
max_iterations: int = 50
|
||||
sample_rate: int = 256
|
||||
target_partition_size: Optional[int] = None
|
||||
# Name of the accelerator (e.g. "cuda") to use for IVF training. When set,
|
||||
# create_index() dispatches to pylance to build the index on the accelerator.
|
||||
# Reserved for future accelerator support. create_index() currently raises if set.
|
||||
accelerator: Optional[str] = None
|
||||
|
||||
|
||||
|
||||
@@ -68,6 +68,14 @@ from ..table import AsyncTable, BlobMode, Branches, IndexStatistics, Query, Tabl
|
||||
from ..types import BaseTokenizerType
|
||||
|
||||
|
||||
def _reject_index_accelerator(
|
||||
config: Optional[IndexConfigType] = None,
|
||||
accelerator: Optional[str] = None,
|
||||
) -> None:
|
||||
if accelerator is not None or getattr(config, "accelerator", None) is not None:
|
||||
raise ValueError("Index accelerators are not supported on LanceDB Cloud.")
|
||||
|
||||
|
||||
class RemoteTable(Table):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -457,6 +465,8 @@ class RemoteTable(Table):
|
||||
... "l2", vector_column_name="vector"
|
||||
... )
|
||||
"""
|
||||
_reject_index_accelerator(config, accelerator)
|
||||
|
||||
# Detect whether this is a legacy API call
|
||||
is_legacy = self._is_legacy_create_index_call(
|
||||
metric,
|
||||
@@ -484,12 +494,6 @@ class RemoteTable(Table):
|
||||
|
||||
column = vector_column_name
|
||||
|
||||
if accelerator is not None:
|
||||
logging.warning(
|
||||
"GPU accelerator is not yet supported on LanceDB cloud."
|
||||
"If you have 100M+ vectors to index,"
|
||||
"please contact us at contact@lancedb.com"
|
||||
)
|
||||
if replace is not None:
|
||||
logging.warning(
|
||||
"replace is not supported on LanceDB cloud."
|
||||
@@ -557,6 +561,8 @@ class RemoteTable(Table):
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the index exists until :meth:`Job.wait` returns.
|
||||
"""
|
||||
_reject_index_accelerator(config)
|
||||
|
||||
return Job(
|
||||
LOOP.run(
|
||||
self._table.create_index_async(
|
||||
|
||||
@@ -214,6 +214,45 @@ IndexConfigType = Union[
|
||||
# Known distance metrics for legacy API detection
|
||||
KNOWN_METRICS = {"l2", "cosine", "dot", "hamming"}
|
||||
|
||||
_PYLANCE_ACCELERATED_INDEX_TYPE = "IVF_PQ"
|
||||
|
||||
|
||||
def _pylance_accelerated_index_options(
|
||||
config: IndexConfigType,
|
||||
*,
|
||||
accelerator: Optional[str] = None,
|
||||
index_type: Optional[str] = None,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Translate an accelerated vector config into PyLance index options."""
|
||||
if accelerator is None:
|
||||
accelerator = getattr(config, "accelerator", None)
|
||||
if accelerator is None:
|
||||
return None
|
||||
|
||||
if index_type is None:
|
||||
index_type = (
|
||||
_PYLANCE_ACCELERATED_INDEX_TYPE
|
||||
if isinstance(config, IvfPq)
|
||||
else type(config).__name__
|
||||
)
|
||||
if index_type.upper() != _PYLANCE_ACCELERATED_INDEX_TYPE:
|
||||
raise ValueError(
|
||||
f"Index type {index_type} does not support an accelerator; "
|
||||
f"only {_PYLANCE_ACCELERATED_INDEX_TYPE} supports acceleration"
|
||||
)
|
||||
|
||||
return {
|
||||
"index_type": index_type,
|
||||
"metric": getattr(config, "distance_type", "l2"),
|
||||
"num_partitions": getattr(config, "num_partitions", None),
|
||||
"num_sub_vectors": getattr(config, "num_sub_vectors", None),
|
||||
"accelerator": accelerator,
|
||||
"num_bits": getattr(config, "num_bits", 8),
|
||||
"m": getattr(config, "m", 20),
|
||||
"ef_construction": getattr(config, "ef_construction", 300),
|
||||
"target_partition_size": getattr(config, "target_partition_size", None),
|
||||
}
|
||||
|
||||
|
||||
def _into_pyarrow_reader(
|
||||
data, schema: Optional[pa.Schema] = None
|
||||
@@ -2737,20 +2776,17 @@ class LanceTable(Table):
|
||||
)
|
||||
|
||||
# Handle accelerator through pylance
|
||||
if accelerator is not None:
|
||||
accelerated_options = _pylance_accelerated_index_options(
|
||||
config, accelerator=accelerator, index_type=index_type
|
||||
)
|
||||
if accelerated_options is not None:
|
||||
self.to_lance().create_index(
|
||||
column=column,
|
||||
index_type=index_type,
|
||||
metric=metric,
|
||||
num_partitions=num_partitions,
|
||||
num_sub_vectors=num_sub_vectors,
|
||||
replace=replace,
|
||||
accelerator=accelerator,
|
||||
index_cache_size=index_cache_size,
|
||||
num_bits=num_bits,
|
||||
m=m,
|
||||
ef_construction=ef_construction,
|
||||
target_partition_size=target_partition_size,
|
||||
name=name,
|
||||
train=train,
|
||||
**accelerated_options,
|
||||
)
|
||||
self.checkout_latest()
|
||||
return
|
||||
@@ -2758,39 +2794,21 @@ class LanceTable(Table):
|
||||
# New API: metric is the column name
|
||||
column = metric
|
||||
|
||||
# Check if config has accelerator set and dispatch to pylance
|
||||
if config is not None and hasattr(config, "accelerator"):
|
||||
acc = getattr(config, "accelerator", None)
|
||||
if acc is not None:
|
||||
# Dispatch to pylance for GPU acceleration
|
||||
index_type_map = {
|
||||
"IvfFlat": "IVF_FLAT",
|
||||
"IvfSq": "IVF_SQ",
|
||||
"IvfPq": "IVF_PQ",
|
||||
"IvfRq": "IVF_RQ",
|
||||
"HnswPq": "IVF_HNSW_PQ",
|
||||
"HnswSq": "IVF_HNSW_SQ",
|
||||
}
|
||||
cfg_type = type(config).__name__
|
||||
lance_index_type = index_type_map.get(cfg_type, "IVF_PQ")
|
||||
|
||||
self.to_lance().create_index(
|
||||
column=column,
|
||||
index_type=lance_index_type,
|
||||
metric=getattr(config, "distance_type", "l2"),
|
||||
num_partitions=getattr(config, "num_partitions", None),
|
||||
num_sub_vectors=getattr(config, "num_sub_vectors", None),
|
||||
replace=replace,
|
||||
accelerator=acc,
|
||||
num_bits=getattr(config, "num_bits", 8),
|
||||
m=getattr(config, "m", 20),
|
||||
ef_construction=getattr(config, "ef_construction", 300),
|
||||
target_partition_size=getattr(
|
||||
config, "target_partition_size", None
|
||||
),
|
||||
)
|
||||
self.checkout_latest()
|
||||
return
|
||||
accelerated_options = (
|
||||
_pylance_accelerated_index_options(config)
|
||||
if config is not None
|
||||
else None
|
||||
)
|
||||
if accelerated_options is not None:
|
||||
self.to_lance().create_index(
|
||||
column=column,
|
||||
replace=replace,
|
||||
name=name,
|
||||
train=train,
|
||||
**accelerated_options,
|
||||
)
|
||||
self.checkout_latest()
|
||||
return
|
||||
|
||||
return LOOP.run(
|
||||
self._table.create_index(
|
||||
@@ -2818,6 +2836,11 @@ class LanceTable(Table):
|
||||
The job may already be complete when returned; callers must not assume
|
||||
the index exists until :meth:`Job.wait` returns.
|
||||
"""
|
||||
if _pylance_accelerated_index_options(config) is not None:
|
||||
raise ValueError(
|
||||
"Accelerated index creation does not support create_index_async; "
|
||||
"use create_index instead."
|
||||
)
|
||||
return Job(
|
||||
LOOP.run(
|
||||
self._table.create_index_async(
|
||||
@@ -4830,6 +4853,7 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4900,6 +4924,23 @@ class AsyncTable:
|
||||
" BTree, Bitmap, LabelList, Fm, or FTS, but got "
|
||||
+ str(type(config))
|
||||
)
|
||||
accelerated_options = (
|
||||
_pylance_accelerated_index_options(config) if config is not None else None
|
||||
)
|
||||
if accelerated_options is not None:
|
||||
if not self._inner._is_native():
|
||||
raise ValueError("GPU accelerator is not supported on LanceDB Cloud.")
|
||||
dataset = await self.to_lance()
|
||||
await asyncio.to_thread(
|
||||
dataset.create_index,
|
||||
column=column,
|
||||
replace=True if replace is None else replace,
|
||||
name=name,
|
||||
train=train,
|
||||
**accelerated_options,
|
||||
)
|
||||
await self.checkout_latest()
|
||||
return
|
||||
try:
|
||||
await self._inner.create_index(
|
||||
column,
|
||||
@@ -4926,6 +4967,7 @@ class AsyncTable:
|
||||
config: Optional[
|
||||
Union[
|
||||
IvfFlat,
|
||||
IvfSq,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
HnswPq,
|
||||
@@ -4948,6 +4990,11 @@ class AsyncTable:
|
||||
be complete when returned; callers must not assume the index exists
|
||||
until :meth:`AsyncJob.wait` resolves.
|
||||
"""
|
||||
if config is not None and _pylance_accelerated_index_options(config):
|
||||
raise ValueError(
|
||||
"Accelerated index creation does not support create_index_async; "
|
||||
"use create_index instead."
|
||||
)
|
||||
job = await self._inner.create_index_async(
|
||||
column,
|
||||
index=config,
|
||||
|
||||
@@ -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,7 +2,6 @@
|
||||
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||
|
||||
|
||||
import inspect
|
||||
import re
|
||||
import sys
|
||||
from datetime import timedelta
|
||||
@@ -63,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})"
|
||||
|
||||
|
||||
@@ -64,23 +64,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 +115,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 +150,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 +164,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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
@@ -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,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"
|
||||
@@ -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(
|
||||
@@ -881,6 +875,25 @@ def test_remote_create_index_async_returns_job():
|
||||
job.cancel()
|
||||
|
||||
|
||||
def test_remote_create_index_rejects_accelerator():
|
||||
from lancedb.index import IvfPq
|
||||
from lancedb.remote.table import RemoteTable
|
||||
|
||||
inner = MagicMock()
|
||||
inner.name = "test"
|
||||
table = RemoteTable(inner, "dev")
|
||||
|
||||
with pytest.raises(ValueError, match="not supported on LanceDB Cloud"):
|
||||
table.create_index(accelerator="mps")
|
||||
with pytest.raises(ValueError, match="not supported on LanceDB Cloud"):
|
||||
table.create_index("vector", config=IvfPq(accelerator="mps"))
|
||||
with pytest.raises(ValueError, match="not supported on LanceDB Cloud"):
|
||||
table.create_index_async("vector", config=IvfPq(accelerator="mps"))
|
||||
|
||||
inner.create_index.assert_not_called()
|
||||
inner.create_index_async.assert_not_called()
|
||||
|
||||
|
||||
def test_remote_job_wait_raises_on_failure():
|
||||
from lancedb.exceptions import JobFailedError
|
||||
from lancedb.index import BTree
|
||||
|
||||
+183
-187
@@ -2,22 +2,29 @@
|
||||
# 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 time import sleep
|
||||
from typing import List
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import lancedb
|
||||
from lancedb.dependencies import _PANDAS_AVAILABLE
|
||||
from lancedb.index import BTree, FTS, HnswFlat, HnswPq, HnswSq, IvfPq
|
||||
from lancedb.index import (
|
||||
BTree,
|
||||
FTS,
|
||||
HnswFlat,
|
||||
HnswPq,
|
||||
HnswSq,
|
||||
IvfFlat,
|
||||
IvfPq,
|
||||
IvfRq,
|
||||
IvfSq,
|
||||
)
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import pyarrow as pa
|
||||
@@ -28,7 +35,7 @@ from lancedb.db import AsyncConnection, DBConnection
|
||||
from lancedb.embeddings import EmbeddingFunctionConfig, EmbeddingFunctionRegistry
|
||||
from lancedb.expr import col, lit
|
||||
from lancedb.pydantic import LanceModel, Vector
|
||||
from lancedb.table import LanceTable
|
||||
from lancedb.table import AsyncTable, LanceTable
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
@@ -102,30 +109,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"]})
|
||||
@@ -462,38 +445,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)
|
||||
@@ -1471,6 +1422,174 @@ def test_create_index_async_returns_done_job(mem_db: DBConnection):
|
||||
job.cancel()
|
||||
|
||||
|
||||
def test_create_index_dispatches_mps_to_pylance(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"mps_sync",
|
||||
data=[
|
||||
{"vector": [3.1, 4.1]},
|
||||
{"vector": [5.9, 26.5]},
|
||||
],
|
||||
)
|
||||
dataset = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(table, "to_lance", return_value=dataset),
|
||||
patch.object(table, "checkout_latest") as checkout_latest,
|
||||
):
|
||||
with pytest.warns(DeprecationWarning, match="create_index"):
|
||||
table.create_index(
|
||||
metric="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
replace=False,
|
||||
name="vector_mps",
|
||||
)
|
||||
|
||||
dataset.create_index.assert_called_once_with(
|
||||
column="vector",
|
||||
replace=False,
|
||||
index_cache_size=None,
|
||||
name="vector_mps",
|
||||
train=True,
|
||||
index_type="IVF_PQ",
|
||||
metric="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
num_bits=8,
|
||||
m=20,
|
||||
ef_construction=300,
|
||||
target_partition_size=None,
|
||||
)
|
||||
checkout_latest.assert_called_once_with()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"config",
|
||||
[
|
||||
IvfFlat(accelerator="mps"),
|
||||
IvfSq(accelerator="mps"),
|
||||
IvfRq(accelerator="mps"),
|
||||
HnswPq(accelerator="mps"),
|
||||
HnswSq(accelerator="mps"),
|
||||
],
|
||||
)
|
||||
def test_create_index_rejects_unsupported_accelerated_format(
|
||||
mem_db: DBConnection, config
|
||||
):
|
||||
table = mem_db.create_table(
|
||||
"unsupported_accelerator",
|
||||
data=[{"vector": [3.1, 4.1]}, {"vector": [5.9, 26.5]}],
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(table, "to_lance") as to_lance,
|
||||
pytest.raises(ValueError, match="only IVF_PQ supports acceleration"),
|
||||
):
|
||||
table.create_index("vector", config=config)
|
||||
|
||||
to_lance.assert_not_called()
|
||||
|
||||
|
||||
def test_legacy_create_index_rejects_unsupported_accelerated_format(
|
||||
mem_db: DBConnection,
|
||||
):
|
||||
table = mem_db.create_table(
|
||||
"unsupported_legacy_accelerator",
|
||||
data=[{"vector": [3.1, 4.1]}, {"vector": [5.9, 26.5]}],
|
||||
)
|
||||
|
||||
with (
|
||||
pytest.warns(DeprecationWarning, match="create_index"),
|
||||
patch.object(table, "to_lance") as to_lance,
|
||||
pytest.raises(ValueError, match="only IVF_PQ supports acceleration"),
|
||||
):
|
||||
table.create_index(index_type="IVF_FLAT", accelerator="mps")
|
||||
|
||||
to_lance.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_create_index_dispatches_mps_to_pylance():
|
||||
inner = MagicMock()
|
||||
inner._is_native.return_value = True
|
||||
inner.checkout_latest = AsyncMock()
|
||||
table = AsyncTable(inner)
|
||||
dataset = MagicMock()
|
||||
|
||||
with patch.object(table, "to_lance", AsyncMock(return_value=dataset)):
|
||||
await table.create_index(
|
||||
"vector",
|
||||
config=IvfPq(
|
||||
distance_type="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
),
|
||||
name="vector_mps",
|
||||
)
|
||||
|
||||
dataset.create_index.assert_called_once_with(
|
||||
column="vector",
|
||||
replace=True,
|
||||
name="vector_mps",
|
||||
train=True,
|
||||
index_type="IVF_PQ",
|
||||
metric="cosine",
|
||||
num_partitions=4,
|
||||
num_sub_vectors=2,
|
||||
accelerator="mps",
|
||||
num_bits=8,
|
||||
m=20,
|
||||
ef_construction=300,
|
||||
target_partition_size=None,
|
||||
)
|
||||
inner.create_index.assert_not_called()
|
||||
inner.checkout_latest.assert_awaited_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_create_index_rejects_unsupported_accelerated_format():
|
||||
inner = MagicMock()
|
||||
inner._is_native.return_value = True
|
||||
table = AsyncTable(inner)
|
||||
|
||||
with (
|
||||
patch.object(table, "to_lance", AsyncMock()) as to_lance,
|
||||
pytest.raises(ValueError, match="only IVF_PQ supports acceleration"),
|
||||
):
|
||||
await table.create_index("vector", config=IvfFlat(accelerator="mps"))
|
||||
|
||||
to_lance.assert_not_awaited()
|
||||
inner.create_index.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_background_index_rejects_accelerator():
|
||||
inner = MagicMock()
|
||||
inner.create_index_async = AsyncMock()
|
||||
table = AsyncTable(inner)
|
||||
|
||||
with pytest.raises(ValueError, match="Accelerated index creation does not support"):
|
||||
await table.create_index_async("vector", config=IvfPq(accelerator="mps"))
|
||||
|
||||
inner.create_index_async.assert_not_awaited()
|
||||
|
||||
|
||||
def test_background_index_rejects_accelerator(mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
"mps_background",
|
||||
data=[
|
||||
{"vector": [3.1, 4.1]},
|
||||
{"vector": [5.9, 26.5]},
|
||||
],
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="Accelerated index creation does not support"):
|
||||
table.create_index_async("vector", config=IvfPq(accelerator="mps"))
|
||||
|
||||
|
||||
@patch("lancedb.table.AsyncTable.create_index")
|
||||
def test_create_index_method(mock_create_index, mem_db: DBConnection):
|
||||
table = mem_db.create_table(
|
||||
@@ -1884,33 +2003,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
|
||||
@@ -2282,20 +2374,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",
|
||||
@@ -2463,55 +2541,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",
|
||||
@@ -2612,36 +2641,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"]})
|
||||
@@ -3668,8 +3667,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
|
||||
@@ -3681,9 +3680,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",
|
||||
|
||||
@@ -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",
|
||||
[
|
||||
|
||||
@@ -622,6 +622,10 @@ impl Table {
|
||||
self.inner.is_some()
|
||||
}
|
||||
|
||||
pub fn _is_native(&self) -> PyResult<bool> {
|
||||
Ok(self.inner_ref()?.as_native().is_some())
|
||||
}
|
||||
|
||||
/// Closes the table, releasing any resources associated with it.
|
||||
pub fn close(&mut self) {
|
||||
self.inner.take();
|
||||
|
||||
@@ -49,8 +49,8 @@ lance-namespace = { workspace = true }
|
||||
lance-namespace-impls = { workspace = true }
|
||||
metrics = { workspace = true, optional = true }
|
||||
metrics-util = { workspace = true, optional = true }
|
||||
# Pin the GooseFS SDK to the version required by Lance's OpenDAL dependency.
|
||||
goosefs-sdk = { version = "=0.1.9", optional = true }
|
||||
# Pin the transitive GooseFS SDK until the 0.1.6 compile break is fixed upstream.
|
||||
goosefs-sdk = { version = "=0.1.5", optional = true }
|
||||
moka = { workspace = true }
|
||||
pin-project = { workspace = true }
|
||||
tokio = { version = "1.23", features = ["rt-multi-thread", "sync"] }
|
||||
@@ -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
|
||||
|
||||
@@ -17,7 +17,7 @@ use arrow_array::builder::LargeBinaryBuilder;
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use lance::dataset::{BlobRangeRequest as LanceBlobRangeRequest, Dataset, WriteParams};
|
||||
use lance_arrow::FieldExt;
|
||||
use lance_file::version::LanceFileVersion;
|
||||
use lance_encoding::version::LanceFileVersion;
|
||||
use lance_io::object_store::ObjectStore;
|
||||
use object_store::path::Path;
|
||||
|
||||
|
||||
@@ -34,7 +34,7 @@ use crate::remote::{
|
||||
db::{OPT_REMOTE_API_KEY, OPT_REMOTE_HOST_OVERRIDE, OPT_REMOTE_REGION},
|
||||
};
|
||||
use lance::io::ObjectStoreParams;
|
||||
pub use lance_file::version::LanceFileVersion;
|
||||
pub use lance_encoding::version::LanceFileVersion;
|
||||
#[cfg(feature = "remote")]
|
||||
use lance_io::object_store::StorageOptions;
|
||||
use lance_io::object_store::{StorageOptionsAccessor, StorageOptionsProvider};
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -12,7 +12,7 @@ use lance::dataset::refs::Ref;
|
||||
use lance::dataset::{ReadParams, WriteMode, builder::DatasetBuilder};
|
||||
use lance::io::{ObjectStore, ObjectStoreParams, WrappingObjectStore};
|
||||
use lance_datafusion::utils::StreamingWriteSource;
|
||||
use lance_file::version::LanceFileVersion;
|
||||
use lance_encoding::version::LanceFileVersion;
|
||||
use lance_io::object_store::{StorageOptionsAccessor, StorageOptionsProvider};
|
||||
use lance_table::io::commit::commit_handler_from_url;
|
||||
use object_store::local::LocalFileSystem;
|
||||
@@ -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;
|
||||
|
||||
@@ -201,7 +201,7 @@ impl LanceNamespaceDatabase {
|
||||
&self,
|
||||
request: &DbCreateTableRequest,
|
||||
) -> Result<(
|
||||
Option<lance_file::version::LanceFileVersion>,
|
||||
Option<lance_encoding::version::LanceFileVersion>,
|
||||
Option<bool>,
|
||||
Option<bool>,
|
||||
)> {
|
||||
@@ -214,7 +214,7 @@ impl LanceNamespaceDatabase {
|
||||
|
||||
let storage_version_override = storage_options
|
||||
.and_then(|opts| opts.get(OPT_NEW_TABLE_STORAGE_VERSION))
|
||||
.map(|s| s.parse::<lance_file::version::LanceFileVersion>())
|
||||
.map(|s| s.parse::<lance_encoding::version::LanceFileVersion>())
|
||||
.transpose()?;
|
||||
|
||||
let v2_manifest_override = storage_options
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
@@ -2942,7 +2939,7 @@ impl<S: HttpSend> BaseTable for RemoteTable<S> {
|
||||
}
|
||||
|
||||
#[derive(Serialize, Clone, Debug)]
|
||||
pub struct MergeInsertRequest {
|
||||
pub(crate) struct MergeInsertRequest {
|
||||
on: String,
|
||||
when_matched_update_all: bool,
|
||||
when_matched_update_all_filt: Option<String>,
|
||||
@@ -5907,18 +5904,16 @@ mod tests {
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Positions are relative to the first retained token, so dropping the
|
||||
// leading "hello" stop word does not shift the remaining tokens.
|
||||
assert_eq!(
|
||||
tokens,
|
||||
vec![
|
||||
FtsToken {
|
||||
text: "こんにちは".to_string(),
|
||||
position: 0,
|
||||
position: 1,
|
||||
},
|
||||
FtsToken {
|
||||
text: "世界".to_string(),
|
||||
position: 1,
|
||||
position: 2,
|
||||
},
|
||||
]
|
||||
);
|
||||
@@ -6494,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| {
|
||||
|
||||
@@ -90,7 +90,7 @@ struct RemoteBlobState {
|
||||
|
||||
/// Seekable Cloud blob handle over HTTP Range.
|
||||
#[derive(Debug)]
|
||||
pub struct RemoteBlobFile {
|
||||
pub(crate) struct RemoteBlobFile {
|
||||
requester: Arc<dyn BlobRangeRequester>,
|
||||
state: Mutex<RemoteBlobState>,
|
||||
closed: AtomicBool,
|
||||
|
||||
@@ -33,7 +33,7 @@ use crate::table::{AddResult, MergeResult};
|
||||
/// same Arrow-IPC streaming body and error side-channel; only the target
|
||||
/// endpoint, query parameters, and parsed result type differ.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum WriteOp {
|
||||
pub(crate) enum WriteOp {
|
||||
/// `add`: stream to `/v1/table/{id}/insert/`, optionally overwriting.
|
||||
Insert { overwrite: bool },
|
||||
/// `merge_insert`: stream to `/v1/table/{id}/merge_insert/` with the merge
|
||||
@@ -49,7 +49,7 @@ pub enum WriteOp {
|
||||
/// The parsed server response for a completed write, discriminated by the
|
||||
/// operation that produced it.
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum WriteResult {
|
||||
pub(crate) enum WriteResult {
|
||||
Add(AddResult),
|
||||
Merge(MergeResult),
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -10,7 +10,7 @@ use arrow_array::{
|
||||
use arrow_schema::{DataType, Field, Fields, Schema};
|
||||
use futures::TryStreamExt;
|
||||
use lance::Dataset;
|
||||
use lance_file::version::LanceFileVersion;
|
||||
use lance_encoding::version::LanceFileVersion;
|
||||
use lancedb::{
|
||||
Connection, Error, Result, Table,
|
||||
blob::{BlobRangeRequest, blob},
|
||||
|
||||
Reference in New Issue
Block a user