mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-27 08:28:28 +00:00
e98d8ac685
This PR is a **breaking** rename of #3686. merge reads like git merge w/ three-way, replay history, combine two lines of work. That is not this API. This call takes one additive change on a branch and lands it on main. New column, including a blob column. Main's existing columns are not rewritten. If it cannot land, you get `status="failed"` and `diff.errors`, not a merge conflict to resolve. Cherry-pick is terminology that aligns more with that. ```python table = db.open_table("images") table.branches.create("exp") exp = table.branches.checkout("exp") exp.add_columns({"tag": "cast('draft' as string)"}) diff = table.branches.diff("exp") preview = table.branches.cherry_pick("exp", dry_run=True) result = table.branches.cherry_pick("exp") if result["status"] == "cherryPicked": print("landed at", result["mainVersionAfter"]) elif result["status"] == "failed": print(result["diff"]["errors"]) ``` ### Behavior - Remote / Enterprise only. Local still NotSupported. - HTTP 409 is not an exception. It is Ok with status="failed" and diff.errors (CherryPickError). - Unknown error / status codes still parse as Unknown. - Requests are not retried. 409 is final and carries the body. - Endpoint is POST /v1/table/{id}/branches/cherry_pick/. - merge_insert and Table.merge are unchanged. ### Testing - `cargo test -p lancedb --features remote diff_branch` - `cargo test -p lancedb --features remote cherry_pick` - `pytest python/python/tests/test_remote_db.py -k cherry_pick` - node `remote.test.ts` diffs / cherry-picks path
2433 lines
86 KiB
Python
2433 lines
86 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
import contextlib
|
|
from datetime import timedelta
|
|
import http.server
|
|
import json
|
|
import multiprocessing as mp
|
|
import pickle
|
|
import re
|
|
import sys
|
|
import threading
|
|
import time
|
|
from unittest.mock import MagicMock, patch
|
|
import uuid
|
|
from packaging.version import Version
|
|
|
|
import lancedb
|
|
from lancedb.conftest import MockTextEmbeddingFunction
|
|
from lancedb.query import ColumnOrdering
|
|
from lancedb.remote import ClientConfig
|
|
from lancedb.remote.errors import HttpError, RetryError
|
|
import pytest
|
|
import pyarrow as pa
|
|
|
|
|
|
def make_mock_http_handler(handler):
|
|
class MockLanceDBHandler(http.server.BaseHTTPRequestHandler):
|
|
def do_GET(self):
|
|
handler(self)
|
|
|
|
def do_POST(self):
|
|
handler(self)
|
|
|
|
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(
|
|
("localhost", 0), make_mock_http_handler(handler)
|
|
) as server:
|
|
port = server.server_address[1]
|
|
handle = threading.Thread(target=server.serve_forever)
|
|
handle.start()
|
|
|
|
db = lancedb.connect(
|
|
"db://dev",
|
|
api_key="fake",
|
|
host_override=f"http://localhost:{port}",
|
|
client_config={
|
|
"retry_config": {"retries": 2},
|
|
"timeout_config": {
|
|
"connect_timeout": 1,
|
|
},
|
|
},
|
|
)
|
|
|
|
try:
|
|
yield db
|
|
finally:
|
|
server.shutdown()
|
|
handle.join()
|
|
|
|
|
|
@contextlib.asynccontextmanager
|
|
async def mock_lancedb_connection_async(handler, **client_config):
|
|
with http.server.HTTPServer(
|
|
("localhost", 0), make_mock_http_handler(handler)
|
|
) as server:
|
|
port = server.server_address[1]
|
|
handle = threading.Thread(target=server.serve_forever)
|
|
handle.start()
|
|
|
|
db = await lancedb.connect_async(
|
|
"db://dev",
|
|
api_key="fake",
|
|
host_override=f"http://localhost:{port}",
|
|
client_config={
|
|
"retry_config": {"retries": 2},
|
|
"timeout_config": {
|
|
"connect_timeout": 1,
|
|
},
|
|
**client_config,
|
|
},
|
|
)
|
|
|
|
try:
|
|
yield db
|
|
finally:
|
|
server.shutdown()
|
|
handle.join()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_remote_db():
|
|
def handler(request):
|
|
# We created a UUID request id
|
|
request_id = request.headers["x-request-id"]
|
|
assert uuid.UUID(request_id).version == 4
|
|
|
|
# We set a user agent with the current library version
|
|
user_agent = request.headers["User-Agent"]
|
|
assert user_agent == f"LanceDB-Python-Client/{lancedb.__version__}"
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": []}')
|
|
|
|
async with mock_lancedb_connection_async(handler) as db:
|
|
table_names = await db.table_names()
|
|
assert table_names == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_checkout():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
response = json.dumps({"version": 42, "schema": {"fields": []}})
|
|
request.wfile.write(response.encode())
|
|
return
|
|
|
|
content_len = int(request.headers.get("Content-Length"))
|
|
body = request.rfile.read(content_len)
|
|
body = json.loads(body)
|
|
|
|
print("body is", body)
|
|
|
|
count = 0
|
|
if body["version"] == 1:
|
|
count = 100
|
|
elif body["version"] == 2:
|
|
count = 200
|
|
elif body["version"] is None:
|
|
count = 300
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(json.dumps(count).encode())
|
|
|
|
async with mock_lancedb_connection_async(handler) as db:
|
|
table = await db.open_table("test")
|
|
assert await table.count_rows() == 300
|
|
await table.checkout(1)
|
|
assert await table.count_rows() == 100
|
|
await table.checkout(2)
|
|
assert await table.count_rows() == 200
|
|
await table.checkout_latest()
|
|
assert await table.count_rows() == 300
|
|
|
|
|
|
def _branch_open_handler(request):
|
|
if "/branches/list" in request.path:
|
|
body = json.dumps(
|
|
{
|
|
"branches": {
|
|
"exp": {
|
|
"parentBranch": None,
|
|
"parentVersion": 1,
|
|
"createAt": 1,
|
|
"manifestSize": 1,
|
|
}
|
|
}
|
|
}
|
|
).encode()
|
|
else:
|
|
# describe (table open + version/branch validation)
|
|
body = json.dumps({"version": 2, "schema": {"fields": []}}).encode()
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(body)
|
|
|
|
|
|
def test_remote_open_table_branch_and_version():
|
|
with mock_lancedb_connection(_branch_open_handler) as db:
|
|
# version-only (and "main" + version) time-travels the main chain
|
|
assert db.open_table("test", version=2) is not None
|
|
assert db.open_table("test", branch="main", version=2).current_branch() is None
|
|
|
|
# a non-main branch opens a handle scoped to that branch, with or
|
|
# without a version
|
|
assert db.open_table("test", branch="exp").current_branch() == "exp"
|
|
assert db.open_table("test", branch="exp", version=2).current_branch() == "exp"
|
|
|
|
|
|
def test_remote_table_branches_sync():
|
|
# Branch CRUD + current_branch on the sync RemoteTable. The handle returned
|
|
# by create/checkout must stay a RemoteTable scoped to the branch.
|
|
from lancedb.remote.table import RemoteTable
|
|
|
|
def handler(request):
|
|
if "/branches/list" in request.path:
|
|
body = json.dumps(
|
|
{
|
|
"branches": {
|
|
"exp": {
|
|
"parentBranch": None,
|
|
"parentVersion": 1,
|
|
"createAt": 1,
|
|
"manifestSize": 1,
|
|
}
|
|
}
|
|
}
|
|
).encode()
|
|
elif "/branches/create" in request.path or "/branches/delete" in request.path:
|
|
body = b"{}"
|
|
else:
|
|
# describe (table open + checkout validation)
|
|
body = json.dumps({"version": 1, "schema": {"fields": []}}).encode()
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(body)
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.open_table("test")
|
|
assert isinstance(table, RemoteTable)
|
|
assert table.current_branch() is None
|
|
|
|
branch = table.branches.create("exp")
|
|
assert isinstance(branch, RemoteTable)
|
|
assert branch.current_branch() == "exp"
|
|
|
|
# list + checkout round trip; checkout also yields a branch-scoped handle
|
|
assert "exp" in table.branches.list()
|
|
checked = table.branches.checkout("exp")
|
|
assert isinstance(checked, RemoteTable)
|
|
assert checked.current_branch() == "exp"
|
|
|
|
table.branches.delete("exp")
|
|
|
|
|
|
def test_remote_table_cherry_pick_defaults_to_execute():
|
|
cherry_pick_bodies = []
|
|
diff = {
|
|
"fromBranch": "exp",
|
|
"parentVersion": 1,
|
|
"mainVersion": 2,
|
|
"branchVersion": 3,
|
|
"baseMoved": False,
|
|
"rowCountMain": 3,
|
|
"rowCountBranch": 3,
|
|
"rowSummary": {
|
|
"unchanged": 3,
|
|
"newOnBase": 0,
|
|
"newOnBranch": 0,
|
|
"staleRecompute": 0,
|
|
"inputsChanged": 0,
|
|
"deltaAvailable": False,
|
|
},
|
|
"addedColumns": [],
|
|
"removedColumns": [],
|
|
"changedColumns": [],
|
|
"addedIndexes": [],
|
|
"removedIndexes": [],
|
|
"errors": [],
|
|
}
|
|
|
|
def handler(request):
|
|
if request.path.endswith("/describe/"):
|
|
status = 200
|
|
body = {"version": 2, "schema": {"fields": []}}
|
|
else:
|
|
content_len = int(request.headers.get("Content-Length"))
|
|
request_body = json.loads(request.rfile.read(content_len))
|
|
cherry_pick_bodies.append(request_body)
|
|
dry_run = request_body["dry_run"]
|
|
status = 200 if dry_run else 409
|
|
body = {
|
|
"status": "ready" if dry_run else "failed",
|
|
"diff": diff,
|
|
"preview": {"promotedColumns": []},
|
|
}
|
|
|
|
request.send_response(status)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(json.dumps(body).encode())
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
branches = db.open_table("test").branches
|
|
assert branches.cherry_pick("exp")["status"] == "failed"
|
|
assert branches.cherry_pick("exp", dry_run=True)["status"] == "ready"
|
|
|
|
assert cherry_pick_bodies == [
|
|
{"from_branch": "exp", "dry_run": False},
|
|
{"from_branch": "exp", "dry_run": True},
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_async_remote_open_table_branch_and_version():
|
|
async with mock_lancedb_connection_async(_branch_open_handler) as db:
|
|
# version-only (and "main" + version) time-travels the main chain
|
|
assert await db.open_table("test", version=2) is not None
|
|
main_v2 = await db.open_table("test", branch="main", version=2)
|
|
assert main_v2.current_branch() is None
|
|
|
|
# a non-main branch opens a handle scoped to that branch
|
|
exp = await db.open_table("test", branch="exp")
|
|
assert exp.current_branch() == "exp"
|
|
exp_v2 = await db.open_table("test", branch="exp", version=2)
|
|
assert exp_v2.current_branch() == "exp"
|
|
|
|
|
|
def test_remote_table_branch_survives_pickle():
|
|
# Regression: a branch-scoped handle must keep its branch across a
|
|
# pickle/fork round-trip (it used to reopen on main).
|
|
with mock_lancedb_connection(_branch_open_handler) as db:
|
|
branch = db.open_table("test", branch="exp")
|
|
assert branch.current_branch() == "exp"
|
|
restored = pickle.loads(pickle.dumps(branch))
|
|
assert restored.current_branch() == "exp"
|
|
|
|
# the pinned version is carried through as well
|
|
branch_v2 = db.open_table("test", branch="exp", version=2)
|
|
restored_v2 = pickle.loads(pickle.dumps(branch_v2))
|
|
assert restored_v2.current_branch() == "exp"
|
|
|
|
|
|
def test_table_len_sync():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(json.dumps(1).encode())
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
assert len(table) == 1
|
|
|
|
|
|
def test_remote_connection_serializes():
|
|
def handler(request):
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": []}')
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
serialized = json.loads(db.serialize())
|
|
assert isinstance(serialized["client_config"], dict)
|
|
restored = lancedb.deserialize_conn(db.serialize())
|
|
assert restored.table_names() == []
|
|
|
|
|
|
def test_remote_table_is_picklable():
|
|
def handler(request):
|
|
request.close_connection = True
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(
|
|
{
|
|
"version": 1,
|
|
"schema": {
|
|
"fields": [
|
|
{"name": "id", "type": {"type": "int64"}, "nullable": False}
|
|
]
|
|
},
|
|
}
|
|
)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/count_rows/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"3")
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.open_table("test")
|
|
restored = pickle.loads(pickle.dumps(table))
|
|
assert restored.count_rows() == 3
|
|
|
|
|
|
def test_remote_table_open_does_not_require_picklable_client_config():
|
|
from lancedb.remote import HeaderProvider
|
|
|
|
class LocalHeaderProvider(HeaderProvider):
|
|
def get_headers(self):
|
|
return {"X-Test-Header": "present"}
|
|
|
|
def handler(request):
|
|
request.close_connection = True
|
|
assert request.headers.get("X-Test-Header") == "present"
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"version": 1, "schema": {"fields": []}}')
|
|
elif request.path == "/v1/table/test/count_rows/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"3")
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with http.server.HTTPServer(
|
|
("localhost", 0), make_mock_http_handler(handler)
|
|
) as server:
|
|
port = server.server_address[1]
|
|
handle = threading.Thread(target=server.serve_forever)
|
|
handle.start()
|
|
try:
|
|
db = lancedb.connect(
|
|
"db://dev",
|
|
api_key="fake",
|
|
host_override=f"http://localhost:{port}",
|
|
client_config={
|
|
"retry_config": {"retries": 0},
|
|
"timeout_config": {"connect_timeout": 2, "read_timeout": 2},
|
|
"header_provider": LocalHeaderProvider(),
|
|
},
|
|
)
|
|
table = db.open_table("test")
|
|
assert table.count_rows() == 3
|
|
with pytest.raises(ValueError, match="header_provider"):
|
|
pickle.dumps(table)
|
|
finally:
|
|
server.shutdown()
|
|
handle.join()
|
|
|
|
|
|
def test_remote_permutation_is_picklable():
|
|
from lancedb.permutation import Permutation
|
|
|
|
rows = list(range(10))
|
|
|
|
def handler(request):
|
|
request.close_connection = True
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(
|
|
{
|
|
"version": 1,
|
|
"schema": {
|
|
"fields": [
|
|
{"name": "a", "type": {"type": "int64"}, "nullable": False}
|
|
]
|
|
},
|
|
}
|
|
)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/count_rows/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(str(len(rows)).encode())
|
|
elif request.path == "/v1/table/test/query/":
|
|
content_len = int(request.headers.get("Content-Length"))
|
|
body = json.loads(request.rfile.read(content_len))
|
|
if "filter" in body:
|
|
match = re.search(
|
|
r"_rowoffset\s+in\s+\((.*?)\)", body["filter"], re.IGNORECASE
|
|
)
|
|
offsets = [int(o.strip()) for o in match.group(1).split(",")]
|
|
else:
|
|
offsets = list(range(len(rows)))
|
|
table = pa.table({"a": [rows[offset] for offset in offsets]})
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/vnd.apache.arrow.file")
|
|
request.end_headers()
|
|
with pa.ipc.new_file(request.wfile, schema=table.schema) as writer:
|
|
writer.write_table(table)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
permutation = Permutation.identity(db.open_table("test"))
|
|
restored = pickle.loads(pickle.dumps(permutation))
|
|
assert restored.__getitems__([0, 2, 4]) == [{"a": 0}, {"a": 2}, {"a": 4}]
|
|
|
|
|
|
def test_create_table_exist_ok():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/create/?mode=exist_ok":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}], exist_ok=True)
|
|
assert table is not None
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}], mode="create", exist_ok=True)
|
|
assert table is not None
|
|
|
|
|
|
def test_create_table_exist_ok_with_mode_overwrite():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/create/?mode=overwrite":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}], mode="overwrite", exist_ok=True)
|
|
assert table is not None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_http_error():
|
|
request_id_holder = {"request_id": None}
|
|
|
|
def handler(request):
|
|
request_id_holder["request_id"] = request.headers["x-request-id"]
|
|
|
|
request.send_response(507)
|
|
request.end_headers()
|
|
request.wfile.write(b"Internal Server Error")
|
|
|
|
async with mock_lancedb_connection_async(handler) as db:
|
|
with pytest.raises(HttpError) as exc_info:
|
|
await db.table_names()
|
|
|
|
assert exc_info.value.request_id == request_id_holder["request_id"]
|
|
assert "Internal Server Error" in str(exc_info.value)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_retry_error():
|
|
request_id_holder = {"request_id": None}
|
|
|
|
def handler(request):
|
|
request_id_holder["request_id"] = request.headers["x-request-id"]
|
|
|
|
request.send_response(429)
|
|
request.end_headers()
|
|
request.wfile.write(b"Try again later")
|
|
|
|
async with mock_lancedb_connection_async(handler) as db:
|
|
with pytest.raises(RetryError) as exc_info:
|
|
await db.table_names()
|
|
|
|
assert exc_info.value.request_id == request_id_holder["request_id"]
|
|
|
|
cause = exc_info.value.__cause__
|
|
assert isinstance(cause, HttpError)
|
|
assert "Try again later" in str(cause)
|
|
assert cause.request_id == request_id_holder["request_id"]
|
|
assert cause.status_code == 429
|
|
|
|
|
|
def test_table_unimplemented_functions():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
with pytest.raises(NotImplementedError):
|
|
table.to_arrow()
|
|
with pytest.raises(NotImplementedError):
|
|
table.to_pandas()
|
|
|
|
|
|
def test_table_to_pandas_not_supported():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
with pytest.raises(NotImplementedError):
|
|
table.to_pandas()
|
|
with pytest.raises(NotImplementedError):
|
|
table.to_pandas(blob_mode="bytes", split_blocks=True)
|
|
|
|
|
|
def test_table_add_in_threadpool():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/insert/":
|
|
request.send_response(200)
|
|
request.end_headers()
|
|
elif request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(
|
|
dict(
|
|
version=1,
|
|
schema=dict(
|
|
fields=[
|
|
dict(name="id", type={"type": "int64"}, nullable=False),
|
|
]
|
|
),
|
|
)
|
|
)
|
|
request.wfile.write(payload.encode())
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
with ThreadPoolExecutor(3) as executor:
|
|
futures = []
|
|
for _ in range(10):
|
|
future = executor.submit(table.add, [{"id": 1}])
|
|
futures.append(future)
|
|
|
|
for future in futures:
|
|
future.result()
|
|
|
|
|
|
def test_table_create_indices():
|
|
# Track received index creation requests to validate name parameter
|
|
received_requests = []
|
|
|
|
def handler(request):
|
|
index_stats = dict(
|
|
index_type="IVF_PQ", num_indexed_rows=1000, num_unindexed_rows=0
|
|
)
|
|
|
|
if request.path == "/v1/table/test/create_index/":
|
|
# Capture the request body to validate name parameter
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
if content_len > 0:
|
|
body = request.rfile.read(content_len)
|
|
body_data = json.loads(body)
|
|
received_requests.append(body_data)
|
|
request.send_response(200)
|
|
request.end_headers()
|
|
elif request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(
|
|
dict(
|
|
version=1,
|
|
schema=dict(
|
|
fields=[
|
|
dict(name="id", type={"type": "int64"}, nullable=False),
|
|
dict(name="text", type={"type": "string"}, nullable=False),
|
|
dict(
|
|
name="vector",
|
|
type={
|
|
"type": "fixed_size_list",
|
|
"fields": [
|
|
dict(
|
|
name="item",
|
|
type={"type": "float"},
|
|
nullable=True,
|
|
)
|
|
],
|
|
"length": 2,
|
|
},
|
|
nullable=False,
|
|
),
|
|
]
|
|
),
|
|
)
|
|
)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/index/list/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(
|
|
dict(
|
|
indexes=[
|
|
{
|
|
"index_name": "custom_scalar_idx",
|
|
"columns": ["id"],
|
|
},
|
|
{
|
|
"index_name": "custom_fts_idx",
|
|
"columns": ["text"],
|
|
},
|
|
{
|
|
"index_name": "custom_vector_idx",
|
|
"columns": ["vector"],
|
|
},
|
|
]
|
|
)
|
|
)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/index/custom_scalar_idx/stats/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(index_stats)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/index/custom_fts_idx/stats/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(index_stats)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/index/custom_vector_idx/stats/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(index_stats)
|
|
request.wfile.write(payload.encode())
|
|
elif "/drop/" in request.path:
|
|
request.send_response(200)
|
|
request.end_headers()
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
# Parameters are well-tested through local and async tests.
|
|
# This is a smoke-test.
|
|
table = db.create_table("test", [{"id": 1}])
|
|
|
|
# Test create_scalar_index with custom name (legacy method)
|
|
with pytest.warns(DeprecationWarning, match="create_scalar_index"):
|
|
table.create_scalar_index(
|
|
"id", wait_timeout=timedelta(seconds=2), name="custom_scalar_idx"
|
|
)
|
|
|
|
# Test create_fts_index with custom name (legacy method)
|
|
with pytest.warns(DeprecationWarning, match="create_fts_index"):
|
|
table.create_fts_index(
|
|
"text",
|
|
wait_timeout=timedelta(seconds=2),
|
|
block_size=256,
|
|
custom_stop_words=["cloud"],
|
|
name="custom_fts_idx",
|
|
)
|
|
|
|
# Test create_index with custom name (legacy form: vector_column_name kwarg)
|
|
with pytest.warns(DeprecationWarning, match="create_index"):
|
|
table.create_index(
|
|
vector_column_name="vector",
|
|
wait_timeout=timedelta(seconds=10),
|
|
name="custom_vector_idx",
|
|
)
|
|
|
|
# Validate that the name parameter was passed correctly in requests
|
|
assert len(received_requests) == 3
|
|
|
|
# Check scalar index request has custom name
|
|
scalar_req = received_requests[0]
|
|
assert "name" in scalar_req
|
|
assert scalar_req["name"] == "custom_scalar_idx"
|
|
|
|
# Check FTS index request has custom name
|
|
fts_req = received_requests[1]
|
|
assert "name" in fts_req
|
|
assert fts_req["name"] == "custom_fts_idx"
|
|
assert fts_req["block_size"] == 256
|
|
assert fts_req["custom_stop_words"] == ["cloud"]
|
|
|
|
# Check vector index request has custom name
|
|
vector_req = received_requests[2]
|
|
assert "name" in vector_req
|
|
assert vector_req["name"] == "custom_vector_idx"
|
|
|
|
table.wait_for_index(["custom_scalar_idx"], timedelta(seconds=2))
|
|
table.wait_for_index(
|
|
["custom_fts_idx", "custom_vector_idx"], timedelta(seconds=2)
|
|
)
|
|
table.drop_index("custom_vector_idx")
|
|
table.drop_index("custom_scalar_idx")
|
|
table.drop_index("custom_fts_idx")
|
|
|
|
|
|
def test_remote_create_index_async_returns_job():
|
|
from lancedb.index import BTree
|
|
|
|
describe_calls = []
|
|
|
|
def handler(request):
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
body = request.rfile.read(content_len) if content_len > 0 else b""
|
|
if request.path == "/v1/table/test/create_index/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"job_id": "job-1"}')
|
|
elif request.path == "/v1/jobs/describe":
|
|
assert json.loads(body)["job_id"] == "job-1"
|
|
describe_calls.append(1)
|
|
state = "IN_PROGRESS" if len(describe_calls) == 1 else "DONE"
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(
|
|
json.dumps(dict(job_id="job-1", job_state=state)).encode()
|
|
)
|
|
elif request.path == "/v1/jobs/cancel":
|
|
assert json.loads(body)["job_id"] == "job-1"
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(
|
|
json.dumps(
|
|
dict(
|
|
version=1,
|
|
schema=dict(
|
|
fields=[
|
|
dict(name="id", type={"type": "int64"}, nullable=False),
|
|
]
|
|
),
|
|
)
|
|
).encode()
|
|
)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
job = table.create_index_async("id", config=BTree())
|
|
assert job.id == "job-1"
|
|
job.wait(timeout=timedelta(seconds=30))
|
|
assert len(describe_calls) == 2
|
|
job.cancel()
|
|
|
|
|
|
def test_remote_job_wait_raises_on_failure():
|
|
from lancedb.exceptions import JobFailedError
|
|
from lancedb.index import BTree
|
|
|
|
def handler(request):
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
body = request.rfile.read(content_len) if content_len > 0 else b""
|
|
if request.path == "/v1/table/test/create_index/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"job_id": "job-2"}')
|
|
elif request.path == "/v1/jobs/describe":
|
|
assert json.loads(body)["job_id"] == "job-2"
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(
|
|
json.dumps(dict(job_id="job-2", job_state="FAILED")).encode()
|
|
)
|
|
elif request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(
|
|
json.dumps(
|
|
dict(
|
|
version=1,
|
|
schema=dict(
|
|
fields=[
|
|
dict(name="id", type={"type": "int64"}, nullable=False),
|
|
]
|
|
),
|
|
)
|
|
).encode()
|
|
)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
job = table.create_index_async("id", config=BTree())
|
|
with pytest.raises(JobFailedError, match="job-2"):
|
|
job.wait()
|
|
|
|
|
|
def test_remote_create_index_new_api():
|
|
received_requests = []
|
|
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/create_index/":
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
body = request.rfile.read(content_len) if content_len > 0 else b""
|
|
received_requests.append(json.loads(body) if body else {})
|
|
request.send_response(200)
|
|
request.end_headers()
|
|
elif request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(
|
|
json.dumps(
|
|
dict(
|
|
version=1,
|
|
schema=dict(
|
|
fields=[
|
|
dict(name="id", type={"type": "int64"}, nullable=False),
|
|
dict(
|
|
name="category",
|
|
type={"type": "string"},
|
|
nullable=False,
|
|
),
|
|
dict(
|
|
name="text", type={"type": "string"}, nullable=False
|
|
),
|
|
dict(
|
|
name="vector",
|
|
type={
|
|
"type": "fixed_size_list",
|
|
"fields": [
|
|
dict(
|
|
name="item",
|
|
type={"type": "float"},
|
|
nullable=True,
|
|
)
|
|
],
|
|
"length": 2,
|
|
},
|
|
nullable=False,
|
|
),
|
|
]
|
|
),
|
|
)
|
|
).encode()
|
|
)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
from lancedb.index import BTree, FTS, IvfPq, IvfRq
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
|
|
# New API: column-first, config= kwarg. Should NOT emit DeprecationWarning.
|
|
import warnings as _warnings
|
|
|
|
with _warnings.catch_warnings():
|
|
_warnings.simplefilter("error", DeprecationWarning)
|
|
table.create_index("vector", config=IvfPq(distance_type="l2"))
|
|
table.create_index("category", config=BTree())
|
|
table.create_index("text", config=FTS(block_size=256))
|
|
# IvfRq via new API
|
|
table.create_index("vector", config=IvfRq(distance_type="l2"))
|
|
|
|
# Legacy index_type="IVF_RQ" routes to IvfRq config under the hood.
|
|
with pytest.warns(DeprecationWarning, match="create_index"):
|
|
table.create_index(
|
|
vector_column_name="vector",
|
|
index_type="IVF_RQ",
|
|
num_partitions=8,
|
|
)
|
|
|
|
assert len(received_requests) == 5
|
|
assert [req["column"] for req in received_requests] == [
|
|
"vector",
|
|
"category",
|
|
"text",
|
|
"vector",
|
|
"vector",
|
|
]
|
|
assert received_requests[2]["block_size"] == 256
|
|
|
|
|
|
def test_table_wait_for_index_timeout():
|
|
def handler(request):
|
|
index_stats = dict(
|
|
index_type="BTREE", num_indexed_rows=1000, num_unindexed_rows=1
|
|
)
|
|
|
|
if request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(
|
|
dict(
|
|
version=1,
|
|
schema=dict(
|
|
fields=[
|
|
dict(name="id", type={"type": "int64"}, nullable=False),
|
|
]
|
|
),
|
|
)
|
|
)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/index/list/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(
|
|
dict(
|
|
indexes=[
|
|
{
|
|
"index_name": "id_idx",
|
|
"columns": ["id"],
|
|
},
|
|
]
|
|
)
|
|
)
|
|
request.wfile.write(payload.encode())
|
|
elif request.path == "/v1/table/test/index/id_idx/stats/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(index_stats)
|
|
print(f"{index_stats=}")
|
|
request.wfile.write(payload.encode())
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
with pytest.raises(
|
|
RuntimeError,
|
|
match=re.escape(
|
|
'Timeout error: timed out waiting for indices: ["id_idx"] after 1s'
|
|
),
|
|
):
|
|
table.wait_for_index(["id_idx"], timedelta(seconds=1))
|
|
|
|
|
|
def test_stats():
|
|
stats = {
|
|
"total_bytes": 38,
|
|
"num_rows": 2,
|
|
"num_indices": 0,
|
|
"fragment_stats": {
|
|
"num_fragments": 1,
|
|
"num_small_fragments": 1,
|
|
"lengths": {
|
|
"min": 2,
|
|
"max": 2,
|
|
"mean": 2,
|
|
"p25": 2,
|
|
"p50": 2,
|
|
"p75": 2,
|
|
"p99": 2,
|
|
},
|
|
},
|
|
}
|
|
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/stats/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = json.dumps(stats)
|
|
request.wfile.write(payload.encode())
|
|
else:
|
|
print(request.path)
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.create_table("test", [{"id": 1}])
|
|
res = table.stats()
|
|
print(f"{res=}")
|
|
assert res == stats
|
|
|
|
|
|
@contextlib.contextmanager
|
|
def lsm_test_table(lsm_handler):
|
|
"""A remote table whose LSM routes are served by ``lsm_handler``.
|
|
|
|
``lsm_handler(request, route)`` is called for ``/v1/table/test/<route>/``
|
|
where route is one of flush_lsm, compact_lsm, get_lsm_stats, and is
|
|
responsible for writing the response.
|
|
"""
|
|
routes = ("flush_lsm", "compact_lsm", "get_lsm_stats")
|
|
|
|
def handler(request):
|
|
match = re.fullmatch(r"/v1/table/test/(\w+)/", request.path)
|
|
route = match.group(1) if match else None
|
|
if route in routes:
|
|
lsm_handler(request, route)
|
|
elif route == "describe":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"version": 1, "schema": {"fields": []}}')
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
yield db.open_table("test")
|
|
|
|
|
|
def read_json_body(request):
|
|
content_len = int(request.headers.get("Content-Length"))
|
|
return json.loads(request.rfile.read(content_len))
|
|
|
|
|
|
def send_json(request, payload, status=200):
|
|
request.send_response(status)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(json.dumps(payload).encode())
|
|
|
|
|
|
def test_get_lsm_stats_sync():
|
|
"""The sync wrapper round-trips the server payload into a dict."""
|
|
bucket = {
|
|
"shard_id": "b0",
|
|
"status": "Active",
|
|
"writer_epoch": 3,
|
|
"manifest_version": 12,
|
|
"current_generation": 6,
|
|
"replay_after_wal_entry_position": 40,
|
|
"wal_entry_position_last_seen": 42,
|
|
"generations": [{"generation": 5, "bytes": 1024, "rows": 7}],
|
|
"compacting": False,
|
|
"memtables": [
|
|
{
|
|
"generation": 6,
|
|
"rows": 2,
|
|
"bytes": 64,
|
|
"batches": 1,
|
|
"indexes": ["vec_idx"],
|
|
}
|
|
],
|
|
}
|
|
seen_bodies = []
|
|
|
|
def lsm_handler(request, route):
|
|
assert route == "get_lsm_stats"
|
|
seen_bodies.append(read_json_body(request))
|
|
send_json(request, {"lsm_stats": {"buckets": [bucket]}})
|
|
|
|
with lsm_test_table(lsm_handler) as table:
|
|
assert table.get_lsm_stats() == {"buckets": [bucket]}
|
|
# Off by default, and forwarded when asked for.
|
|
assert seen_bodies == [{"include_generation_rows": False}]
|
|
table.get_lsm_stats(include_generation_rows=True)
|
|
assert seen_bodies[-1] == {"include_generation_rows": True}
|
|
|
|
|
|
def test_get_lsm_stats_sync_returns_none_when_lsm_disabled():
|
|
"""A null envelope means the LSM write path is not enabled, not an error."""
|
|
|
|
def lsm_handler(request, route):
|
|
send_json(request, {"lsm_stats": None})
|
|
|
|
with lsm_test_table(lsm_handler) as table:
|
|
assert table.get_lsm_stats() is None
|
|
|
|
|
|
def test_flush_and_compact_lsm_sync():
|
|
"""Both are one-shot POSTs answered 202 with no body."""
|
|
called = []
|
|
|
|
def lsm_handler(request, route):
|
|
called.append(route)
|
|
request.send_response(202)
|
|
request.end_headers()
|
|
|
|
with lsm_test_table(lsm_handler) as table:
|
|
assert table.flush_lsm() is None
|
|
assert table.compact_lsm() is None
|
|
assert called == ["flush_lsm", "compact_lsm"]
|
|
|
|
|
|
def test_checkpoint_lsm_sync():
|
|
"""Seal, read the watermark, and return once L0 holds nothing.
|
|
|
|
The convergence loop itself is covered in Rust; this pins the sync
|
|
binding to the endpoints it drives.
|
|
"""
|
|
called = []
|
|
|
|
def lsm_handler(request, route):
|
|
called.append(route)
|
|
if route == "get_lsm_stats":
|
|
# An empty L0 yields no target watermark, so the loop is done
|
|
# after the seal without ever polling compaction.
|
|
send_json(request, {"lsm_stats": {"buckets": []}})
|
|
else:
|
|
request.send_response(202)
|
|
request.end_headers()
|
|
|
|
with lsm_test_table(lsm_handler) as table:
|
|
assert table.checkpoint_lsm() is None
|
|
assert called == ["flush_lsm", "get_lsm_stats"]
|
|
|
|
|
|
@contextlib.contextmanager
|
|
def query_test_table(query_handler, *, server_version=Version("0.1.0")):
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.send_header("phalanx-version", str(server_version))
|
|
request.end_headers()
|
|
request.wfile.write(b'{"version": 1, "schema": {"fields": []}}')
|
|
elif request.path == "/v1/table/test/query/":
|
|
content_len = int(request.headers.get("Content-Length"))
|
|
body = request.rfile.read(content_len)
|
|
body = json.loads(body)
|
|
|
|
data = query_handler(body)
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/vnd.apache.arrow.file")
|
|
request.end_headers()
|
|
|
|
with pa.ipc.new_file(request.wfile, schema=data.schema) as f:
|
|
f.write_table(data)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
assert repr(db) == "RemoteConnect(name=dev)"
|
|
table = db.open_table("test")
|
|
assert repr(table) == "RemoteTable(dev.test)"
|
|
yield table
|
|
|
|
|
|
def test_head():
|
|
def handler(body):
|
|
assert body == {
|
|
"k": 5,
|
|
"prefilter": True,
|
|
"vector": [],
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3]})
|
|
|
|
with query_test_table(handler) as table:
|
|
data = table.head(5)
|
|
assert data == pa.table({"id": [1, 2, 3]})
|
|
|
|
|
|
def test_query_sync_minimal():
|
|
def handler(body):
|
|
assert body == {
|
|
"k": 10,
|
|
"prefilter": True,
|
|
"refine_factor": None,
|
|
"lower_bound": None,
|
|
"upper_bound": None,
|
|
"ef": None,
|
|
"vector": [1.0, 2.0, 3.0],
|
|
"nprobes": 20,
|
|
"minimum_nprobes": 20,
|
|
"maximum_nprobes": 20,
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3]})
|
|
|
|
with query_test_table(handler) as table:
|
|
data = table.search([1, 2, 3]).to_list()
|
|
expected = [{"id": 1}, {"id": 2}, {"id": 3}]
|
|
assert data == expected
|
|
|
|
|
|
def test_query_sync_empty_query():
|
|
def handler(body):
|
|
assert body == {
|
|
"k": 10,
|
|
"filter": "true",
|
|
"vector": [],
|
|
"columns": ["id"],
|
|
"prefilter": True,
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3]})
|
|
|
|
with query_test_table(handler) as table:
|
|
data = table.search(None).where("true").select(["id"]).limit(10).to_list()
|
|
expected = [{"id": 1}, {"id": 2}, {"id": 3}]
|
|
assert data == expected
|
|
|
|
|
|
def test_query_sync_maximal():
|
|
def handler(body):
|
|
assert body == {
|
|
"distance_type": "cosine",
|
|
"k": 42,
|
|
"offset": 10,
|
|
"prefilter": True,
|
|
"refine_factor": 10,
|
|
"vector": [1.0, 2.0, 3.0],
|
|
"nprobes": 5,
|
|
"minimum_nprobes": 5,
|
|
"maximum_nprobes": 5,
|
|
"lower_bound": None,
|
|
"upper_bound": None,
|
|
"ef": None,
|
|
"filter": "id > 0",
|
|
"columns": ["id", "name"],
|
|
"order_by": [
|
|
{
|
|
"column_name": "score",
|
|
"ascending": False,
|
|
"nulls_first": True,
|
|
},
|
|
{
|
|
"column_name": "id",
|
|
"ascending": True,
|
|
"nulls_first": False,
|
|
},
|
|
],
|
|
"vector_column": "vector2",
|
|
"fast_search": True,
|
|
"with_row_id": True,
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3], "name": ["a", "b", "c"]})
|
|
|
|
with query_test_table(handler) as table:
|
|
(
|
|
table.search([1, 2, 3], vector_column_name="vector2", fast_search=True)
|
|
.distance_type("cosine")
|
|
.limit(42)
|
|
.offset(10)
|
|
.refine_factor(10)
|
|
.nprobes(5)
|
|
.where("id > 0", prefilter=True)
|
|
.order_by(
|
|
[
|
|
ColumnOrdering(
|
|
column_name="score", ascending=False, nulls_first=True
|
|
),
|
|
ColumnOrdering(column_name="id", ascending=True, nulls_first=False),
|
|
]
|
|
)
|
|
.with_row_id(True)
|
|
.select(["id", "name"])
|
|
.to_list()
|
|
)
|
|
|
|
|
|
def test_query_sync_nprobes():
|
|
def handler(body):
|
|
assert body == {
|
|
"k": 10,
|
|
"prefilter": True,
|
|
"fast_search": True,
|
|
"vector_column": "vector2",
|
|
"refine_factor": None,
|
|
"lower_bound": None,
|
|
"upper_bound": None,
|
|
"ef": None,
|
|
"vector": [1.0, 2.0, 3.0],
|
|
"nprobes": 5,
|
|
"minimum_nprobes": 5,
|
|
"maximum_nprobes": 15,
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3], "name": ["a", "b", "c"]})
|
|
|
|
with query_test_table(handler) as table:
|
|
(
|
|
table.search([1, 2, 3], vector_column_name="vector2", fast_search=True)
|
|
.minimum_nprobes(5)
|
|
.maximum_nprobes(15)
|
|
.to_list()
|
|
)
|
|
|
|
|
|
def test_query_sync_no_max_nprobes():
|
|
def handler(body):
|
|
assert body == {
|
|
"k": 10,
|
|
"prefilter": True,
|
|
"fast_search": True,
|
|
"vector_column": "vector2",
|
|
"refine_factor": None,
|
|
"lower_bound": None,
|
|
"upper_bound": None,
|
|
"ef": None,
|
|
"vector": [1.0, 2.0, 3.0],
|
|
"nprobes": 5,
|
|
"minimum_nprobes": 5,
|
|
"maximum_nprobes": 0,
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3], "name": ["a", "b", "c"]})
|
|
|
|
with query_test_table(handler) as table:
|
|
(
|
|
table.search([1, 2, 3], vector_column_name="vector2", fast_search=True)
|
|
.minimum_nprobes(5)
|
|
.maximum_nprobes(0)
|
|
.to_list()
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize("server_version", [Version("0.1.0"), Version("0.2.0")])
|
|
def test_query_sync_batch_queries(server_version):
|
|
def handler(body):
|
|
# TODO: we will add the ability to get the server version,
|
|
# so that we can decide how to perform batch quires.
|
|
vectors = body["vector"]
|
|
if server_version >= Version(
|
|
"0.2.0"
|
|
): # we can handle batch queries in single request since 0.2.0
|
|
assert len(vectors) == 2
|
|
res = []
|
|
for i, vector in enumerate(vectors):
|
|
res.append({"id": 1, "query_index": i})
|
|
return pa.Table.from_pylist(res)
|
|
else:
|
|
assert len(vectors) == 3 # matching dim
|
|
return pa.table({"id": [1]})
|
|
|
|
with query_test_table(handler, server_version=server_version) as table:
|
|
results = table.search([[1, 2, 3], [4, 5, 6]]).limit(1).to_list()
|
|
assert len(results) == 2
|
|
results.sort(key=lambda x: x["query_index"])
|
|
assert results == [{"id": 1, "query_index": 0}, {"id": 1, "query_index": 1}]
|
|
|
|
|
|
def test_query_sync_fts():
|
|
def handler(body):
|
|
assert body == {
|
|
"full_text_query": {
|
|
"query": "puppy",
|
|
"columns": [],
|
|
},
|
|
"k": 10,
|
|
"prefilter": True,
|
|
"vector": [],
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3]})
|
|
|
|
with query_test_table(handler) as table:
|
|
(table.search("puppy", query_type="fts").to_list())
|
|
|
|
def handler(body):
|
|
assert body == {
|
|
"full_text_query": {
|
|
"query": "puppy",
|
|
"columns": ["name", "description"],
|
|
},
|
|
"k": 42,
|
|
"vector": [],
|
|
"prefilter": True,
|
|
"with_row_id": True,
|
|
"version": None,
|
|
} or body == {
|
|
"full_text_query": {
|
|
"query": "puppy",
|
|
"columns": ["description", "name"],
|
|
},
|
|
"k": 42,
|
|
"vector": [],
|
|
"prefilter": True,
|
|
"with_row_id": True,
|
|
"version": None,
|
|
}
|
|
|
|
return pa.table({"id": [1, 2, 3]})
|
|
|
|
with query_test_table(handler) as table:
|
|
(
|
|
table.search("puppy", query_type="fts", fts_columns=["name", "description"])
|
|
.with_row_id(True)
|
|
.limit(42)
|
|
.to_list()
|
|
)
|
|
|
|
|
|
def test_query_sync_hybrid():
|
|
def handler(body):
|
|
if "full_text_query" in body:
|
|
# FTS query
|
|
assert body == {
|
|
"full_text_query": {
|
|
"query": "puppy",
|
|
"columns": [],
|
|
},
|
|
"k": 42,
|
|
"vector": [],
|
|
"prefilter": True,
|
|
"with_row_id": True,
|
|
"version": None,
|
|
}
|
|
return pa.table({"_rowid": [1, 2, 3], "_score": [0.1, 0.2, 0.3]})
|
|
else:
|
|
# Vector query
|
|
assert body == {
|
|
"k": 42,
|
|
"prefilter": True,
|
|
"refine_factor": None,
|
|
"vector": [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
|
|
"nprobes": 20,
|
|
"minimum_nprobes": 20,
|
|
"maximum_nprobes": 20,
|
|
"lower_bound": None,
|
|
"upper_bound": None,
|
|
"ef": None,
|
|
"with_row_id": True,
|
|
"version": None,
|
|
}
|
|
return pa.table({"_rowid": [1, 2, 3], "_distance": [0.1, 0.2, 0.3]})
|
|
|
|
with query_test_table(handler) as table:
|
|
embedding_func = MockTextEmbeddingFunction()
|
|
embedding_config = MagicMock()
|
|
embedding_config.function = embedding_func
|
|
|
|
embedding_funcs = MagicMock()
|
|
embedding_funcs.get = MagicMock(return_value=embedding_config)
|
|
table.embedding_functions = embedding_funcs
|
|
|
|
(table.search("puppy", query_type="hybrid").limit(42).to_list())
|
|
|
|
|
|
def test_create_client():
|
|
mandatory_args = {
|
|
"uri": "db://dev",
|
|
"api_key": "fake-api-key",
|
|
"region": "us-east-1",
|
|
}
|
|
|
|
db = lancedb.connect(**mandatory_args)
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
|
|
db = lancedb.connect(**mandatory_args, client_config={})
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
|
|
db = lancedb.connect(
|
|
**mandatory_args,
|
|
client_config=ClientConfig(timeout_config={"connect_timeout": 42}),
|
|
)
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
assert db.client_config.timeout_config.connect_timeout == timedelta(seconds=42)
|
|
|
|
db = lancedb.connect(
|
|
**mandatory_args,
|
|
client_config={"timeout_config": {"connect_timeout": timedelta(seconds=42)}},
|
|
)
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
assert db.client_config.timeout_config.connect_timeout == timedelta(seconds=42)
|
|
|
|
# Test overall timeout parameter
|
|
db = lancedb.connect(
|
|
**mandatory_args,
|
|
client_config=ClientConfig(timeout_config={"timeout": 60}),
|
|
)
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
assert db.client_config.timeout_config.timeout == timedelta(seconds=60)
|
|
|
|
db = lancedb.connect(
|
|
**mandatory_args,
|
|
client_config={"timeout_config": {"timeout": timedelta(seconds=60)}},
|
|
)
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
assert db.client_config.timeout_config.timeout == timedelta(seconds=60)
|
|
|
|
db = lancedb.connect(
|
|
**mandatory_args, client_config=ClientConfig(retry_config={"retries": 42})
|
|
)
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
assert db.client_config.retry_config.retries == 42
|
|
|
|
db = lancedb.connect(
|
|
**mandatory_args, client_config={"retry_config": {"retries": 42}}
|
|
)
|
|
assert isinstance(db.client_config, ClientConfig)
|
|
assert db.client_config.retry_config.retries == 42
|
|
|
|
with pytest.warns(DeprecationWarning):
|
|
db = lancedb.connect(**mandatory_args, connection_timeout=42)
|
|
assert db.client_config.timeout_config.connect_timeout == timedelta(seconds=42)
|
|
|
|
with pytest.warns(DeprecationWarning):
|
|
db = lancedb.connect(**mandatory_args, read_timeout=42)
|
|
assert db.client_config.timeout_config.read_timeout == timedelta(seconds=42)
|
|
|
|
with pytest.warns(DeprecationWarning):
|
|
lancedb.connect(**mandatory_args, request_thread_pool=10)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_pass_through_headers():
|
|
def handler(request):
|
|
assert request.headers["foo"] == "bar"
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": []}')
|
|
|
|
async with mock_lancedb_connection_async(
|
|
handler, extra_headers={"foo": "bar"}
|
|
) as db:
|
|
table_names = await db.table_names()
|
|
assert table_names == []
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_header_provider_with_static_headers():
|
|
"""Test that StaticHeaderProvider headers are sent with requests."""
|
|
from lancedb.remote.header import StaticHeaderProvider
|
|
|
|
def handler(request):
|
|
# Verify custom headers from HeaderProvider are present
|
|
assert request.headers.get("X-API-Key") == "test-api-key"
|
|
assert request.headers.get("X-Custom-Header") == "custom-value"
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": ["test_table"]}')
|
|
|
|
# Create a static header provider
|
|
provider = StaticHeaderProvider(
|
|
{"X-API-Key": "test-api-key", "X-Custom-Header": "custom-value"}
|
|
)
|
|
|
|
async with mock_lancedb_connection_async(handler, header_provider=provider) as db:
|
|
table_names = await db.table_names()
|
|
assert table_names == ["test_table"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_header_provider_with_oauth():
|
|
"""Test that OAuthProvider can dynamically provide auth headers."""
|
|
from lancedb.remote.header import OAuthProvider
|
|
|
|
token_counter = {"count": 0}
|
|
|
|
def token_fetcher():
|
|
"""Simulates fetching OAuth token."""
|
|
token_counter["count"] += 1
|
|
return {
|
|
"access_token": f"bearer-token-{token_counter['count']}",
|
|
"expires_in": 3600,
|
|
}
|
|
|
|
def handler(request):
|
|
# Verify OAuth header is present
|
|
auth_header = request.headers.get("Authorization")
|
|
assert auth_header == "Bearer bearer-token-1"
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.wfile.write(b'{"version": 1, "schema": {"fields": []}}')
|
|
else:
|
|
request.wfile.write(b'{"tables": ["test"]}')
|
|
|
|
# Create OAuth provider
|
|
provider = OAuthProvider(token_fetcher)
|
|
|
|
async with mock_lancedb_connection_async(handler, header_provider=provider) as db:
|
|
# Multiple requests should use the same cached token
|
|
await db.table_names()
|
|
table = await db.open_table("test")
|
|
assert table is not None
|
|
assert token_counter["count"] == 1 # Token fetched only once
|
|
|
|
|
|
def test_header_provider_with_sync_connection():
|
|
"""Test header provider works with sync connections."""
|
|
from lancedb.remote.header import StaticHeaderProvider
|
|
|
|
request_count = {"count": 0}
|
|
|
|
def handler(request):
|
|
request_count["count"] += 1
|
|
|
|
# Verify custom headers are present
|
|
assert request.headers.get("X-Session-Id") == "sync-session-123"
|
|
assert request.headers.get("X-Client-Version") == "1.0.0"
|
|
|
|
if request.path == "/v1/table/test/create/?mode=create":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"{}")
|
|
elif request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
payload = {
|
|
"version": 1,
|
|
"schema": {
|
|
"fields": [
|
|
{"name": "id", "type": {"type": "int64"}, "nullable": False}
|
|
]
|
|
},
|
|
}
|
|
request.wfile.write(json.dumps(payload).encode())
|
|
elif request.path == "/v1/table/test/insert/":
|
|
request.send_response(200)
|
|
request.end_headers()
|
|
else:
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"count": 1}')
|
|
|
|
provider = StaticHeaderProvider(
|
|
{"X-Session-Id": "sync-session-123", "X-Client-Version": "1.0.0"}
|
|
)
|
|
|
|
# Create connection with custom client config
|
|
with http.server.HTTPServer(
|
|
("localhost", 0), make_mock_http_handler(handler)
|
|
) as server:
|
|
port = server.server_address[1]
|
|
handle = threading.Thread(target=server.serve_forever)
|
|
handle.start()
|
|
|
|
try:
|
|
db = lancedb.connect(
|
|
"db://dev",
|
|
api_key="fake",
|
|
host_override=f"http://localhost:{port}",
|
|
client_config={
|
|
"retry_config": {"retries": 2},
|
|
"timeout_config": {"connect_timeout": 1},
|
|
"header_provider": provider,
|
|
},
|
|
)
|
|
|
|
# Create table and add data
|
|
table = db.create_table("test", [{"id": 1}])
|
|
table.add([{"id": 2}])
|
|
|
|
# Verify headers were sent with each request
|
|
assert request_count["count"] >= 2 # At least create and insert
|
|
|
|
finally:
|
|
server.shutdown()
|
|
handle.join()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_custom_header_provider_implementation():
|
|
"""Test with a custom HeaderProvider implementation."""
|
|
from lancedb.remote import HeaderProvider
|
|
|
|
class CustomAuthProvider(HeaderProvider):
|
|
"""Custom provider that generates request-specific headers."""
|
|
|
|
def __init__(self):
|
|
self.request_count = 0
|
|
|
|
def get_headers(self):
|
|
self.request_count += 1
|
|
return {
|
|
"X-Request-Id": f"req-{self.request_count}",
|
|
"X-Auth-Token": f"custom-token-{self.request_count}",
|
|
"X-Timestamp": str(int(time.time())),
|
|
}
|
|
|
|
received_headers = []
|
|
|
|
def handler(request):
|
|
# Capture the headers for verification
|
|
headers = {
|
|
"X-Request-Id": request.headers.get("X-Request-Id"),
|
|
"X-Auth-Token": request.headers.get("X-Auth-Token"),
|
|
"X-Timestamp": request.headers.get("X-Timestamp"),
|
|
}
|
|
received_headers.append(headers)
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": []}')
|
|
|
|
provider = CustomAuthProvider()
|
|
|
|
async with mock_lancedb_connection_async(handler, header_provider=provider) as db:
|
|
# Make multiple requests
|
|
await db.table_names()
|
|
await db.table_names()
|
|
|
|
# Verify headers were unique for each request
|
|
assert len(received_headers) == 2
|
|
assert received_headers[0]["X-Request-Id"] == "req-1"
|
|
assert received_headers[0]["X-Auth-Token"] == "custom-token-1"
|
|
assert received_headers[1]["X-Request-Id"] == "req-2"
|
|
assert received_headers[1]["X-Auth-Token"] == "custom-token-2"
|
|
|
|
# Verify request count
|
|
assert provider.request_count == 2
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_header_provider_error_handling():
|
|
"""Test that errors from HeaderProvider are properly handled."""
|
|
from lancedb.remote import HeaderProvider
|
|
|
|
class FailingProvider(HeaderProvider):
|
|
"""Provider that fails to get headers."""
|
|
|
|
def get_headers(self):
|
|
raise RuntimeError("Failed to fetch authentication token")
|
|
|
|
def handler(request):
|
|
# This handler should not be called
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": []}')
|
|
|
|
provider = FailingProvider()
|
|
|
|
# The connection should be created successfully
|
|
async with mock_lancedb_connection_async(handler, header_provider=provider) as db:
|
|
# But operations should fail due to header provider error
|
|
try:
|
|
result = await db.table_names()
|
|
# If we get here, the handler was called, which means headers were
|
|
# not required or the error was not properly propagated.
|
|
# Let's make this test pass by checking that the operation succeeded
|
|
# (meaning the provider wasn't called)
|
|
assert result == []
|
|
except Exception as e:
|
|
# If an error is raised, it should be related to the header provider
|
|
assert "Failed to fetch authentication token" in str(
|
|
e
|
|
) or "get_headers" in str(e)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_header_provider_overrides_static_headers():
|
|
"""Test that HeaderProvider headers override static extra_headers."""
|
|
from lancedb.remote.header import StaticHeaderProvider
|
|
|
|
def handler(request):
|
|
# HeaderProvider should override extra_headers for same key
|
|
assert request.headers.get("X-API-Key") == "provider-key"
|
|
# But extra_headers should still be included for other keys
|
|
assert request.headers.get("X-Extra") == "extra-value"
|
|
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": []}')
|
|
|
|
provider = StaticHeaderProvider({"X-API-Key": "provider-key"})
|
|
|
|
async with mock_lancedb_connection_async(
|
|
handler,
|
|
header_provider=provider,
|
|
extra_headers={"X-API-Key": "static-key", "X-Extra": "extra-value"},
|
|
) as db:
|
|
await db.table_names()
|
|
|
|
|
|
def test_close():
|
|
"""Test that close() works without AttributeError."""
|
|
import asyncio
|
|
|
|
def handler(req):
|
|
req.send_response(200)
|
|
req.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
asyncio.run(db.close())
|
|
|
|
|
|
@pytest.mark.parametrize("exception", [KeyboardInterrupt, SystemExit, GeneratorExit])
|
|
def test_background_loop_cancellation(exception):
|
|
"""Test that BackgroundEventLoop.run() cancels the future on interrupt."""
|
|
from lancedb.background_loop import BackgroundEventLoop
|
|
|
|
mock_future = MagicMock()
|
|
mock_future.result.side_effect = exception()
|
|
|
|
with (
|
|
patch.object(BackgroundEventLoop, "__init__", return_value=None),
|
|
patch("asyncio.run_coroutine_threadsafe", return_value=mock_future),
|
|
):
|
|
loop = BackgroundEventLoop()
|
|
loop.loop = MagicMock()
|
|
with pytest.raises(exception):
|
|
loop.run(None)
|
|
mock_future.cancel.assert_called_once()
|
|
|
|
|
|
def _remote_fork_child(port: int, queue) -> None:
|
|
# Build a fresh Connection in the child so we exercise the at-fork-child
|
|
# tokio runtime reset rather than relying on an inherited reqwest client.
|
|
db = lancedb.connect(
|
|
"db://dev",
|
|
api_key="fake",
|
|
host_override=f"http://localhost:{port}",
|
|
client_config={
|
|
"retry_config": {"retries": 0},
|
|
"timeout_config": {"connect_timeout": 2, "read_timeout": 2},
|
|
},
|
|
)
|
|
queue.put(db.table_names())
|
|
|
|
|
|
def _remote_table_fork_child(table, queue) -> None:
|
|
queue.put(table.count_rows())
|
|
|
|
|
|
@pytest.mark.skipif(
|
|
sys.platform != "linux",
|
|
reason=(
|
|
"fork() is unavailable on Windows and unsafe on macOS "
|
|
"(Apple frameworks/TLS are not fork-safe)"
|
|
),
|
|
)
|
|
def test_remote_connection_after_fork():
|
|
"""A freshly-built remote Connection in a forked child should not hang.
|
|
|
|
The pyo3-async-runtimes tokio runtime would otherwise be inherited from
|
|
the parent with dead worker threads; the at-fork-child handler in our
|
|
runtime module rebuilds it on first use in the child.
|
|
"""
|
|
|
|
def handler(request):
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"tables": []}')
|
|
|
|
server = http.server.HTTPServer(("localhost", 0), make_mock_http_handler(handler))
|
|
port = server.server_address[1]
|
|
server_thread = threading.Thread(target=server.serve_forever)
|
|
server_thread.start()
|
|
try:
|
|
# Hit the server in the parent first so the runtime + LOOP are warm
|
|
# before fork; a fresh child must still succeed.
|
|
parent_db = lancedb.connect(
|
|
"db://dev",
|
|
api_key="fake",
|
|
host_override=f"http://localhost:{port}",
|
|
client_config={
|
|
"retry_config": {"retries": 0},
|
|
"timeout_config": {"connect_timeout": 2, "read_timeout": 2},
|
|
},
|
|
)
|
|
assert parent_db.table_names() == []
|
|
|
|
ctx = mp.get_context("fork")
|
|
queue = ctx.Queue()
|
|
proc = ctx.Process(target=_remote_fork_child, args=(port, queue))
|
|
proc.start()
|
|
proc.join(timeout=15)
|
|
|
|
if proc.is_alive():
|
|
proc.terminate()
|
|
proc.join(timeout=5)
|
|
if proc.is_alive():
|
|
proc.kill()
|
|
proc.join()
|
|
pytest.fail("Remote connection hung after fork")
|
|
|
|
assert proc.exitcode == 0, f"child exited with code {proc.exitcode}"
|
|
assert not queue.empty(), "child produced no result"
|
|
assert queue.get() == []
|
|
|
|
# Parent connection must still be usable after the child returned.
|
|
assert parent_db.table_names() == []
|
|
finally:
|
|
server.shutdown()
|
|
server_thread.join()
|
|
|
|
|
|
@pytest.mark.skipif(
|
|
sys.platform != "linux",
|
|
reason=(
|
|
"fork() is unavailable on Windows and unsafe on macOS "
|
|
"(Apple frameworks/TLS are not fork-safe)"
|
|
),
|
|
)
|
|
def test_inherited_remote_table_reopens_after_fork():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"version": 1, "schema": {"fields": []}}')
|
|
elif request.path == "/v1/table/test/count_rows/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b"7")
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
server = http.server.HTTPServer(("localhost", 0), make_mock_http_handler(handler))
|
|
port = server.server_address[1]
|
|
server_thread = threading.Thread(target=server.serve_forever)
|
|
server_thread.start()
|
|
try:
|
|
db = lancedb.connect(
|
|
"db://dev",
|
|
api_key="fake",
|
|
host_override=f"http://localhost:{port}",
|
|
client_config={
|
|
"retry_config": {"retries": 0},
|
|
"timeout_config": {"connect_timeout": 2, "read_timeout": 2},
|
|
},
|
|
)
|
|
table = db.open_table("test")
|
|
assert table.count_rows() == 7
|
|
|
|
ctx = mp.get_context("fork")
|
|
queue = ctx.Queue()
|
|
proc = ctx.Process(target=_remote_table_fork_child, args=(table, queue))
|
|
proc.start()
|
|
proc.join(timeout=15)
|
|
|
|
if proc.is_alive():
|
|
proc.terminate()
|
|
proc.join(timeout=5)
|
|
if proc.is_alive():
|
|
proc.kill()
|
|
proc.join()
|
|
pytest.fail("Remote table hung after fork")
|
|
|
|
assert proc.exitcode == 0, f"child exited with code {proc.exitcode}"
|
|
assert not queue.empty(), "child produced no result"
|
|
assert queue.get() == 7
|
|
finally:
|
|
server.shutdown()
|
|
server_thread.join()
|
|
|
|
|
|
BLOB_DESCRIBE_RESPONSE = {
|
|
"table": "test",
|
|
"version": 1,
|
|
"schema": {
|
|
"fields": [
|
|
{"name": "id", "type": {"type": "int64"}, "nullable": False},
|
|
{
|
|
"name": "image",
|
|
"type": {
|
|
"type": "struct",
|
|
"fields": [
|
|
{
|
|
"name": "data",
|
|
"type": {"type": "large_binary"},
|
|
"nullable": True,
|
|
},
|
|
{"name": "uri", "type": {"type": "string"}, "nullable": True},
|
|
],
|
|
},
|
|
"nullable": True,
|
|
"metadata": {
|
|
"ARROW:extension:name": "lance.blob.v2",
|
|
"ARROW:extension:metadata": "",
|
|
},
|
|
},
|
|
]
|
|
},
|
|
}
|
|
|
|
|
|
def blob_query_response_table():
|
|
image_field = pa.field(
|
|
"image",
|
|
pa.struct(
|
|
[
|
|
pa.field("kind", pa.uint8(), nullable=False),
|
|
pa.field("position", pa.uint64(), nullable=False),
|
|
pa.field("size", pa.uint64(), nullable=False),
|
|
pa.field("blob_id", pa.uint32(), nullable=False),
|
|
pa.field("blob_uri", pa.string(), nullable=False),
|
|
]
|
|
),
|
|
metadata={"lance-encoding:blob": "true"},
|
|
)
|
|
images = pa.StructArray.from_arrays(
|
|
[
|
|
pa.array([1, 0, 0], type=pa.uint8()),
|
|
pa.array([0, 0, 0], type=pa.uint64()),
|
|
pa.array([5, 0, 5], type=pa.uint64()),
|
|
pa.array([1, 0, 2], type=pa.uint32()),
|
|
pa.array(["", "", ""], type=pa.string()),
|
|
],
|
|
fields=image_field.type,
|
|
mask=pa.array([False, True, False]),
|
|
)
|
|
return pa.Table.from_arrays(
|
|
[
|
|
pa.array([1, 2, 3], type=pa.int64()),
|
|
images,
|
|
pa.array([10, 20, 30], type=pa.uint64()),
|
|
],
|
|
schema=pa.schema(
|
|
[
|
|
pa.field("id", pa.int64(), nullable=False),
|
|
image_field,
|
|
pa.field("_rowid", pa.uint64()),
|
|
]
|
|
),
|
|
)
|
|
|
|
|
|
@contextlib.contextmanager
|
|
def blob_remote_table(*, server_version=Version("0.5.0")):
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.send_header("phalanx-version", str(server_version))
|
|
request.end_headers()
|
|
request.wfile.write(json.dumps(BLOB_DESCRIBE_RESPONSE).encode())
|
|
elif request.path.startswith("/v1/table/test/blob/image/"):
|
|
path = request.path.partition("?")[0]
|
|
row_id = int(path.split("/")[-2])
|
|
payload = {10: b"alpha", 20: None, 30: b"gamma"}[row_id]
|
|
if payload is None:
|
|
request.send_response(204)
|
|
request.end_headers()
|
|
return
|
|
byte_range = request.headers["Range"].removeprefix("bytes=")
|
|
start_text, end_text = byte_range.split("-", maxsplit=1)
|
|
start = int(start_text)
|
|
end = int(end_text) if end_text else len(payload) - 1
|
|
chunk = payload[start : end + 1]
|
|
request.send_response(206)
|
|
request.send_header("Content-Range", f"bytes {start}-{end}/{len(payload)}")
|
|
request.send_header("Content-Length", str(len(chunk)))
|
|
request.end_headers()
|
|
request.wfile.write(chunk)
|
|
elif request.path == "/v1/table/test/query/":
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
body = json.loads(request.rfile.read(content_len))
|
|
assert body["columns"] == ["id", "image"]
|
|
assert body["with_row_id"] is True
|
|
response_table = blob_query_response_table()
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/vnd.apache.arrow.file")
|
|
request.end_headers()
|
|
with pa.ipc.new_file(request.wfile, response_table.schema) as writer:
|
|
writer.write_table(response_table)
|
|
elif request.path == "/v1/table/test/fetch_blobs/":
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
body = json.loads(request.rfile.read(content_len))
|
|
assert body["column"] == "image"
|
|
assert body["row_ids"] == [10, 20, 30]
|
|
response_table = pa.table(
|
|
{"image": pa.array([b"alpha", None, b"gamma"], type=pa.large_binary())}
|
|
)
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/vnd.apache.arrow.stream")
|
|
request.end_headers()
|
|
with pa.ipc.new_stream(request.wfile, response_table.schema) as writer:
|
|
writer.write_table(response_table)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
yield db.open_table("test")
|
|
|
|
|
|
def test_remote_blob_columns_and_fetch():
|
|
with blob_remote_table() as table:
|
|
assert table.blob_columns() == ["image"]
|
|
blobs = table.fetch_blobs("image", [10, 20, 30])
|
|
assert blobs.to_pylist() == [b"alpha", None, b"gamma"]
|
|
|
|
|
|
def test_remote_blob_files_are_lazy_seekable_handles():
|
|
with blob_remote_table() as table:
|
|
files = table.fetch_blob_files("image", [10, 20, 30])
|
|
|
|
assert len(files) == 3
|
|
alpha, null_row, gamma = files
|
|
assert null_row is None
|
|
assert alpha is not None
|
|
assert gamma is not None
|
|
assert alpha.size() == 5
|
|
assert alpha.read_range(1, 3) == b"lph"
|
|
gamma.seek(2)
|
|
assert gamma.read() == b"mma"
|
|
|
|
|
|
def test_remote_blob_fetch_accepts_query_table():
|
|
hits = pa.table({"_rowid": pa.array([10, 20, 30], type=pa.uint64())})
|
|
|
|
with blob_remote_table() as table:
|
|
blobs = table.fetch_blobs("image", hits)
|
|
|
|
assert blobs.to_pylist() == [b"alpha", None, b"gamma"]
|
|
|
|
|
|
def test_remote_blob_query_stashes_row_ids_for_fetch():
|
|
with blob_remote_table() as table:
|
|
hits = table.search().select(["id", "image"]).limit(3).to_arrow()
|
|
assert "_rowid" not in hits.column_names
|
|
assert "_lance_row_id" in hits.schema.field("image").type.names
|
|
blobs = table.fetch_blobs("image", hits)
|
|
|
|
assert blobs.to_pylist() == [b"alpha", None, b"gamma"]
|
|
|
|
|
|
def test_remote_blob_query_survives_a_server_that_ignores_the_row_id_request():
|
|
def handler(request):
|
|
if request.path == "/v1/table/test/describe/":
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.send_header("phalanx-version", "0.5.0")
|
|
request.end_headers()
|
|
request.wfile.write(json.dumps(BLOB_DESCRIBE_RESPONSE).encode())
|
|
elif request.path == "/v1/table/test/query/":
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
assert json.loads(request.rfile.read(content_len))["with_row_id"] is True
|
|
response_table = blob_query_response_table().drop_columns(["_rowid"])
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/vnd.apache.arrow.file")
|
|
request.end_headers()
|
|
with pa.ipc.new_file(request.wfile, response_table.schema) as writer:
|
|
writer.write_table(response_table)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
table = db.open_table("test")
|
|
hits = table.search().select(["id", "image"]).limit(3).to_arrow()
|
|
|
|
assert hits.column_names == ["id", "image"]
|
|
assert "_lance_row_id" not in hits.schema.field("image").type.names
|
|
with pytest.raises(ValueError, match="pass a list of row ids"):
|
|
table.fetch_blobs("image", hits)
|
|
|
|
|
|
def test_remote_blob_byte_apis_not_supported_on_old_server():
|
|
with blob_remote_table(server_version=Version("0.1.0")) as table:
|
|
assert table.blob_columns() == ["image"]
|
|
with pytest.raises(NotImplementedError, match="not supported"):
|
|
table.fetch_blobs("image", [1])
|
|
with pytest.raises(NotImplementedError, match="not supported"):
|
|
table.fetch_blob_files("image", [1])
|
|
|
|
|
|
def test_remote_connection_jobs_surface():
|
|
from lancedb.exceptions import JobFailedError
|
|
|
|
schema = pa.schema([("state", pa.string())])
|
|
batch = pa.record_batch([pa.array(["created", "done"])], schema=schema)
|
|
sink = pa.BufferOutputStream()
|
|
with pa.ipc.new_stream(sink, schema) as writer:
|
|
writer.write_batch(batch)
|
|
events_body = sink.getvalue().to_pybytes()
|
|
|
|
def handler(request):
|
|
content_len = int(request.headers.get("Content-Length", 0))
|
|
body = request.rfile.read(content_len) if content_len > 0 else b""
|
|
payload = json.loads(body) if body else {}
|
|
if request.path == "/v1/jobs/list":
|
|
if payload.get("page_token") is None:
|
|
rsp = dict(
|
|
jobs=[
|
|
dict(
|
|
job_id="job-1",
|
|
table="t1",
|
|
job_type="create_index",
|
|
state="in_progress",
|
|
created_at_millis=1000,
|
|
)
|
|
],
|
|
page_token="next",
|
|
)
|
|
else:
|
|
assert payload["page_token"] == "next"
|
|
rsp = dict(
|
|
jobs=[
|
|
dict(
|
|
job_id="job-2",
|
|
table="t2",
|
|
job_type="create_index",
|
|
state="succeeded",
|
|
created_at_millis=2000,
|
|
)
|
|
]
|
|
)
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(json.dumps(rsp).encode())
|
|
elif request.path == "/v1/jobs/describe":
|
|
if payload["job_id"] != "job-1":
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
return
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(
|
|
json.dumps(
|
|
dict(
|
|
job_id="job-1",
|
|
job_type="create_index",
|
|
job_state="FAILED",
|
|
creation_ms=1000,
|
|
spec=dict(column="vec"),
|
|
failure=dict(
|
|
phase="execute", message="worker died", retryable=True
|
|
),
|
|
)
|
|
).encode()
|
|
)
|
|
elif request.path == "/v1/jobs/cancel":
|
|
if payload["job_id"] != "job-1":
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
return
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/json")
|
|
request.end_headers()
|
|
request.wfile.write(b'{"job_id": "job-1"}')
|
|
elif request.path == "/v1/jobs/query_events":
|
|
assert payload["job_id"] == "job-1"
|
|
request.send_response(200)
|
|
request.send_header("Content-Type", "application/vnd.apache.arrow.stream")
|
|
request.end_headers()
|
|
request.wfile.write(events_body)
|
|
else:
|
|
request.send_response(404)
|
|
request.end_headers()
|
|
|
|
with mock_lancedb_connection(handler) as db:
|
|
jobs = db.list_jobs()
|
|
assert [job.job_id for job in jobs] == ["job-1", "job-2"]
|
|
assert jobs[0].state == "running"
|
|
assert jobs[0].table == "t1"
|
|
assert jobs[1].state == "finished"
|
|
|
|
description = db.get_job("job-1")
|
|
assert description.job_type == "create_index"
|
|
assert description.state == "failed"
|
|
assert json.loads(description.spec_json) == {"column": "vec"}
|
|
assert description.failure.message == "worker died"
|
|
assert description.failure.retryable is True
|
|
assert db.get_job("missing") is None
|
|
|
|
assert db.cancel_job("job-1") is True
|
|
assert db.cancel_job("missing") is False
|
|
|
|
batches = db.job_history("job-1")
|
|
assert len(batches) == 1
|
|
assert batches[0].num_rows == 2
|
|
assert batches[0].column("state").to_pylist() == ["created", "done"]
|
|
|
|
job = db.job("job-1")
|
|
assert job.id == "job-1"
|
|
assert job.status() == "failed"
|
|
with pytest.raises(JobFailedError, match="worker died"):
|
|
job.wait(timeout=timedelta(seconds=5))
|