Merge remote-tracking branch 'refs/remotes/origin/main' into gatekeeper/fix-2820-1

# Conflicts:
#	rust/lancedb/src/remote/table.rs
This commit is contained in:
Gatefixer
2026-08-26 16:13:07 +00:00
20 changed files with 869 additions and 60 deletions
+2
View File
@@ -159,6 +159,8 @@ and combined with [BooleanQuery][lancedb.query.BooleanQuery].
::: lancedb.query.FullTextOperator
::: lancedb.query.DocumentGranularity
::: lancedb.query.Occur
## Embeddings
+26 -3
View File
@@ -222,6 +222,7 @@ class PythonEnvironmentSpec(_RemoteValue):
kind: str
packages: tuple[str, ...] = ()
channels: tuple[str, ...] = ()
path: Optional[str] = None
modules: tuple[str, ...] = ()
image: Optional[str] = None
@@ -909,13 +910,25 @@ class UdfDefinition:
pip: tuple[str, ...],
env: Mapping[str, str],
python_version: Optional[str],
conda: tuple[str, ...] = (),
conda_channels: tuple[str, ...] = (),
):
function_name = name or function.__name__
if not _FUNCTION_NAME.fullmatch(function_name):
raise ValueError(f"invalid Function name: {function_name!r}")
packages = tuple(sorted(set(pip)))
if pip and conda:
raise ValueError("a Function environment is pip or conda, not both")
if conda_channels and not conda:
raise ValueError("conda_channels requires conda packages")
packages = tuple(sorted(set(conda if conda else pip)))
if any(not package or package != package.strip() for package in packages):
raise ValueError("pip requirements must be non-empty and trimmed")
raise ValueError("package requirements must be non-empty and trimmed")
if conda:
environment_spec = PythonEnvironmentSpec(
kind="conda", packages=packages, channels=tuple(conda_channels)
)
else:
environment_spec = PythonEnvironmentSpec(kind="pip", packages=packages)
environment = dict(env)
if any(
not isinstance(key, str) or not isinstance(value, str)
@@ -929,7 +942,7 @@ class UdfDefinition:
kind="python",
python_version=python_version
or f"{sys.version_info.major}.{sys.version_info.minor}",
environment=PythonEnvironmentSpec(kind="pip", packages=packages),
environment=environment_spec,
env=environment,
)
self._function = function
@@ -976,6 +989,8 @@ def udf(
pip: tuple[str, ...] | list[str] = (),
env: Optional[Mapping[str, str]] = None,
python_version: Optional[str] = None,
conda: tuple[str, ...] | list[str] = (),
conda_channels: tuple[str, ...] | list[str] = (),
) -> Callable[[Callable[..., Any]], UdfDefinition]: ...
@@ -988,6 +1003,8 @@ def udf(
pip: tuple[str, ...] | list[str] = (),
env: Optional[Mapping[str, str]] = None,
python_version: Optional[str] = None,
conda: tuple[str, ...] | list[str] = (),
conda_channels: tuple[str, ...] | list[str] = (),
):
"""Prepare a scalar Python callable for remote Function registration.
@@ -1010,6 +1027,10 @@ def udf(
provided together with ``input_schema``.
pip : sequence of str, optional
Pip requirements for the remote environment.
conda : sequence of str, optional
Conda packages for the remote environment, instead of ``pip``.
conda_channels : sequence of str, optional
Conda channels in priority order; requires ``conda``.
env : mapping of str to str, optional
Environment variables included in the Function definition.
python_version : str, optional
@@ -1049,6 +1070,8 @@ def udf(
pip=tuple(pip),
env={} if env is None else env,
python_version=python_version,
conda=tuple(conda),
conda_channels=tuple(conda_channels),
)
if function is None:
+12
View File
@@ -7,6 +7,7 @@ from typing import List, Literal, Optional
from ._lancedb import (
IndexConfig,
)
from .query import DocumentGranularity
from .types import BaseTokenizerType
lang_mapping = {
@@ -121,6 +122,11 @@ class FTS:
>>> config = FTS(block_size=256)
Create an index that treats each deepest-list element as one document:
>>> from lancedb.query import DocumentGranularity
>>> config = FTS(document_granularity=DocumentGranularity.LIST_ELEMENT)
Attributes
----------
with_position : bool, default False
@@ -172,6 +178,11 @@ class FTS:
roughly half of the available CPU cores. The effective value is
limited by the available compute capacity. This build-only setting is
not persisted with the index and does not apply to remote tables.
document_granularity : DocumentGranularity, default ROW
``ROW`` treats the selected text in one table row as one document.
``LIST_ELEMENT`` treats each element of the deepest list on the indexed
field path as one document and returns its physical coordinates in
``_doc_index`` for matching queries.
Notes
-----
@@ -196,6 +207,7 @@ class FTS:
custom_stop_words: Optional[List[str]] = None
memory_limit: Optional[int] = None
num_workers: Optional[int] = None
document_granularity: DocumentGranularity = DocumentGranularity.ROW
@dataclass
+23
View File
@@ -376,6 +376,13 @@ class FullTextOperator(str, Enum):
OR = "OR"
class DocumentGranularity(str, Enum):
"""The unit treated as one full-text-search document."""
ROW = "row"
LIST_ELEMENT = "list_element"
class Occur(str, Enum):
SHOULD = "SHOULD"
MUST = "MUST"
@@ -479,6 +486,10 @@ class MatchQuery(FullTextQuery):
prefix_length : int, optional
The number of beginning characters being unchanged for fuzzy matching.
This is useful to achieve prefix matching.
document_granularity : DocumentGranularity, optional
Explicitly select row or deepest-list-element documents. If omitted,
the indexed granularity is inferred. When both granularities are indexed
for the field, this must be specified. With no index, row granularity is used.
"""
query: str
@@ -488,6 +499,9 @@ class MatchQuery(FullTextQuery):
max_expansions: int = pydantic.Field(50, kw_only=True)
operator: FullTextOperator = pydantic.Field(FullTextOperator.OR, kw_only=True)
prefix_length: int = pydantic.Field(0, kw_only=True)
document_granularity: Optional[DocumentGranularity] = pydantic.Field(
None, kw_only=True
)
def query_type(self) -> FullTextQueryType:
return FullTextQueryType.MATCH
@@ -504,11 +518,20 @@ class PhraseQuery(FullTextQuery):
The query string to match against.
column : str
The name of the column to match against.
slop : int, default 0
The maximum number of intervening positions permitted in the phrase.
document_granularity : DocumentGranularity, optional
Explicitly select row or deepest-list-element documents. If omitted,
the indexed granularity is inferred. When both granularities are indexed
for the field, this must be specified. With no index, row granularity is used.
"""
query: str
column: str
slop: int = pydantic.Field(0, kw_only=True)
document_granularity: Optional[DocumentGranularity] = pydantic.Field(
None, kw_only=True
)
def query_type(self) -> FullTextQueryType:
return FullTextQueryType.MATCH_PHRASE
+3
View File
@@ -61,6 +61,7 @@ from lancedb.table import _normalize_progress
from ..query import (
AnalyzePlanDistributedMetrics,
DocumentGranularity,
LanceQueryBuilder,
LanceTakeQueryBuilder,
LanceVectorQueryBuilder,
@@ -349,6 +350,7 @@ class RemoteTable(Table):
ngram_max_length: int = 3,
prefix_only: bool = False,
block_size: int = 128,
document_granularity: DocumentGranularity = DocumentGranularity.ROW,
name: Optional[str] = None,
):
"""Create a full-text search index on a column.
@@ -371,6 +373,7 @@ class RemoteTable(Table):
ngram_max_length=ngram_max_length,
prefix_only=prefix_only,
block_size=block_size,
document_granularity=document_granularity,
)
LOOP.run(
self._table.create_index(
+13 -1
View File
@@ -85,6 +85,7 @@ from .query import (
AsyncQuery,
AsyncTakeQuery,
AsyncVectorQuery,
DocumentGranularity,
FullTextQuery,
LanceEmptyQueryBuilder,
LanceFtsQueryBuilder,
@@ -1168,6 +1169,7 @@ class Table(ABC):
ngram_max_length: int = 3,
prefix_only: bool = False,
block_size: int = 128,
document_granularity: DocumentGranularity = DocumentGranularity.ROW,
wait_timeout: Optional[timedelta] = None,
name: Optional[str] = None,
):
@@ -1246,6 +1248,11 @@ class Table(ABC):
The number of documents per compressed posting block. Must be 128
or 256. A value of 256 uses the experimental FTS V3 format and
may introduce breaking changes.
document_granularity: DocumentGranularity, default ROW
``ROW`` treats the selected text in one table row as one document.
``LIST_ELEMENT`` treats each element of the deepest list on the field
path as one document and returns its physical coordinates in
``_doc_index`` for matching queries.
wait_timeout: timedelta, optional
The timeout to wait if indexing is asynchronous.
name: str, optional
@@ -3273,6 +3280,7 @@ class LanceTable(Table):
ngram_max_length: int = 3,
prefix_only: bool = False,
block_size: int = 128,
document_granularity: DocumentGranularity = DocumentGranularity.ROW,
name: Optional[str] = None,
):
"""Create a full-text search index on a column.
@@ -3324,7 +3332,11 @@ class LanceTable(Table):
tokenizer_configs = self.infer_tokenizer_configs(tokenizer_name)
tokenizer_configs["custom_stop_words"] = custom_stop_words
config = FTS(block_size=block_size, **tokenizer_configs)
config = FTS(
block_size=block_size,
document_granularity=document_granularity,
**tokenizer_configs,
)
try:
LOOP.run(
@@ -69,6 +69,26 @@ def _run_packaged(definition, *args):
return namespace[definition.registration_request.artifact.entrypoint](*args)
def test_udf_conda_environment():
@udf(conda=["scipy", "numpy"], conda_channels=["conda-forge", "defaults"])
def halve(value: float) -> float:
return value / 2
request = json.loads(halve.registration_request.to_canonical_json())
assert request["runtime"]["environment"] == {
"kind": "conda",
"packages": ["numpy", "scipy"],
"channels": ["conda-forge", "defaults"],
}
pip_request = json.loads(normalize_score.registration_request.to_canonical_json())
assert "channels" not in pip_request["runtime"]["environment"]
with pytest.raises(ValueError, match="not both"):
udf(name="both", pip=["numpy"], conda=["numpy"])(lambda value: value)
with pytest.raises(ValueError, match="requires conda"):
udf(name="channels", conda_channels=["conda-forge"])(lambda value: value)
def test_udf_packages_attribute_access_and_body_imports():
@udf
def word_norm(body: str) -> float:
+77
View File
@@ -25,6 +25,7 @@ from lancedb.db import DBConnection
from lancedb.index import FTS
from lancedb.query import (
BoostQuery,
DocumentGranularity,
MatchQuery,
MultiMatchQuery,
PhraseQuery,
@@ -245,6 +246,55 @@ def test_create_inverted_index_rejects_invalid_block_size(table):
table.create_index("text", config=FTS(block_size=129))
def test_list_element_document_granularity(tmp_path):
docs_type = pa.list_(pa.struct([pa.field("content", pa.string())]))
docs = pa.array(
[
[
{"content": "alpha beta"},
None,
{"content": ""},
{"content": "the and"},
{"content": "alpha beta"},
]
],
type=docs_type,
)
table = ldb.connect(tmp_path).create_table(
"list_element_docs", pa.table({"id": [0], "docs": docs})
)
row_table = ldb.connect(tmp_path).create_table(
"row_docs", pa.table({"id": [0], "docs": docs})
)
row_table.create_index("docs.content", config=FTS())
row_result = row_table.search(MatchQuery("alpha", "docs.content")).to_arrow()
assert row_result.num_rows == 1
assert "_doc_index" not in row_result.column_names
granularity = DocumentGranularity.LIST_ELEMENT
table.create_index(
"docs.content",
config=FTS(with_position=True, document_granularity=granularity),
)
assert table.list_indices()[0].columns == ["docs.content"]
def coordinates(query):
result = table.search(query).limit(10).to_arrow()
doc_index_type = result.schema.field("_doc_index").type
assert pa.types.is_list(doc_index_type)
assert doc_index_type.value_type == pa.uint32()
return sorted(result["_doc_index"].to_pylist())
assert coordinates(
MatchQuery("alpha", "docs.content", document_granularity=granularity)
) == [[0], [4]]
assert coordinates(
PhraseQuery("alpha beta", "docs.content", document_granularity=granularity)
) == [[0], [4]]
assert coordinates(MatchQuery("alpha", "docs.content")) == [[0], [4]]
assert FTS().document_granularity is DocumentGranularity.ROW
def test_create_inverted_index_respects_build_memory_limit(table):
with pytest.raises(ValueError, match="exceeds worker memory limit"):
table.create_index(
@@ -1089,6 +1139,20 @@ def test_fts_query_to_json():
)
assert json_str == expected
# Test MatchQuery with list-element document granularity
match_query = MatchQuery(
"hello world",
"text",
document_granularity=DocumentGranularity.LIST_ELEMENT,
)
json_str = match_query.to_json()
expected = (
'{"match":{"column":"text","terms":"hello world","boost":1.0,'
'"fuzziness":0,"max_expansions":50,"operator":"Or","prefix_length":0,'
'"document_granularity":"list_element"}}'
)
assert json_str == expected
# Test MatchQuery with options
match_query = MatchQuery("puppy", "text", fuzziness=2, boost=1.5, prefix_length=3)
json_str = match_query.to_json()
@@ -1098,6 +1162,19 @@ def test_fts_query_to_json():
)
assert json_str == expected
# Test PhraseQuery with list-element document granularity
phrase_query = PhraseQuery(
"quick brown fox",
"title",
document_granularity=DocumentGranularity.LIST_ELEMENT,
)
json_str = phrase_query.to_json()
expected = (
'{"phrase":{"column":"title","terms":"quick brown fox","slop":0,'
'"document_granularity":"list_element"}}'
)
assert json_str == expected
# Test PhraseQuery
phrase_query = PhraseQuery("quick brown fox", "title")
json_str = phrase_query.to_json()
+43
View File
@@ -1643,6 +1643,49 @@ def test_query_sync_fts():
)
def test_query_sync_fts_document_granularity():
from lancedb.query import DocumentGranularity, MatchQuery
def handler(body):
assert body == {
"full_text_query": {
"query": {
"match": {
"column": "docs.content",
"terms": "alpha",
"boost": 1.0,
"fuzziness": 0,
"max_expansions": 50,
"operator": "Or",
"prefix_length": 0,
"document_granularity": "list_element",
}
}
},
"k": 10,
"prefilter": True,
"vector": [],
"version": None,
}
return pa.table(
{
"id": [1, 1],
"_doc_index": pa.array([[0], [4]], type=pa.list_(pa.uint32())),
}
)
with query_test_table(handler, server_version=Version("0.6.0")) as table:
result = table.search(
MatchQuery(
"alpha",
"docs.content",
document_granularity=DocumentGranularity.LIST_ELEMENT,
)
).to_arrow()
assert result["_doc_index"].to_pylist() == [[0], [4]]
def test_query_sync_hybrid():
def handler(body):
if "full_text_query" in body:
+20
View File
@@ -4,6 +4,7 @@
import asyncio
import copy
from concurrent.futures import ThreadPoolExecutor
from datetime import timedelta
import threading
@@ -86,6 +87,25 @@ def test_s3_lifecycle(s3_bucket: str):
asyncio.run(test())
@pytest.mark.s3_test
def test_concurrent_open_table(s3_bucket: str):
uri = f"s3://{s3_bucket}/test_concurrent_open_table"
db = lancedb.connect(uri, storage_options=copy.copy(CONFIG))
db.create_table("test", pa.table({"x": [1, 2, 3]}))
num_workers = 32
barrier = threading.Barrier(num_workers)
def open_and_count(_):
barrier.wait()
return db.open_table("test").count_rows()
with ThreadPoolExecutor(max_workers=num_workers) as pool:
row_counts = list(pool.map(open_and_count, range(num_workers)))
assert row_counts == [3] * num_workers
@pytest.fixture()
def kms_key():
kms = get_boto3_client("kms", endpoint_url=CONFIG["aws_endpoint"])
+8 -2
View File
@@ -8,7 +8,7 @@ use lancedb::index::vector::{
};
use lancedb::index::{
Index as LanceDbIndex,
scalar::{BTreeIndexBuilder, FmIndexBuilder, FtsIndexBuilder},
scalar::{BTreeIndexBuilder, DocumentGranularity, FmIndexBuilder, FtsIndexBuilder},
};
use pyo3::IntoPyObject;
use pyo3::types::PyStringMethods;
@@ -60,7 +60,11 @@ pub fn extract_index_params(source: &Option<Bound<'_, PyAny>>) -> PyResult<Lance
.ngram_min_length(params.ngram_min_length)
.ngram_max_length(params.ngram_max_length)
.ngram_prefix_only(params.prefix_only)
.custom_stop_words(params.custom_stop_words);
.custom_stop_words(params.custom_stop_words)
.document_granularity(
DocumentGranularity::try_from(params.document_granularity.as_str())
.map_err(|err| PyValueError::new_err(err.to_string()))?,
);
if let Some(memory_limit) = params.memory_limit {
inner_opts = inner_opts.memory_limit_mb(memory_limit);
}
@@ -221,6 +225,7 @@ struct FtsParams {
block_size: usize,
memory_limit: Option<u64>,
num_workers: Option<usize>,
document_granularity: String,
}
#[derive(FromPyObject)]
@@ -481,6 +486,7 @@ mod tests {
block_size = 128
memory_limit = 2048
num_workers = 7
document_granularity = 'row'
config = FTS()",
None,
+45 -12
View File
@@ -16,8 +16,8 @@ use arrow::pyarrow::FromPyArrow;
use arrow::pyarrow::IntoPyArrow;
use arrow::pyarrow::ToPyArrow;
use lancedb::index::scalar::{
BooleanQuery, BoostQuery, FtsQuery, FullTextSearchQuery, MatchQuery, MultiMatchQuery, Occur,
Operator, PhraseQuery,
BooleanQuery, BoostQuery, DocumentGranularity, FtsQuery, FullTextSearchQuery, MatchQuery,
MultiMatchQuery, Occur, Operator, PhraseQuery,
};
use lancedb::query::AnalyzePlanDistributedMetrics;
use lancedb::query::QueryBase;
@@ -76,8 +76,16 @@ impl<'a, 'py> FromPyObject<'a, 'py> for PyLanceDB<FtsQuery> {
let max_expansions = ob.getattr("max_expansions")?.extract()?;
let operator = ob.getattr("operator")?.extract::<String>()?;
let prefix_length = ob.getattr("prefix_length")?.extract()?;
let document_granularity = ob
.getattr("document_granularity")?
.extract::<Option<String>>()?
.map(|value| {
DocumentGranularity::try_from(value.as_str())
.map_err(|err| PyValueError::new_err(err.to_string()))
})
.transpose()?;
Ok(Self(
let mut query =
MatchQuery::new(query)
.with_column(Some(column))
.with_boost(boost)
@@ -86,21 +94,32 @@ impl<'a, 'py> FromPyObject<'a, 'py> for PyLanceDB<FtsQuery> {
.with_operator(Operator::try_from(operator.as_str()).map_err(|e| {
PyValueError::new_err(format!("Invalid operator: {}", e))
})?)
.with_prefix_length(prefix_length)
.into(),
))
.with_prefix_length(prefix_length);
if let Some(document_granularity) = document_granularity {
query = query.with_document_granularity(document_granularity);
}
Ok(Self(query.into()))
}
"PhraseQuery" => {
let query = ob.getattr("query")?.extract()?;
let column = ob.getattr("column")?.extract()?;
let slop = ob.getattr("slop")?.extract()?;
let document_granularity = ob
.getattr("document_granularity")?
.extract::<Option<String>>()?
.map(|value| {
DocumentGranularity::try_from(value.as_str())
.map_err(|err| PyValueError::new_err(err.to_string()))
})
.transpose()?;
Ok(Self(
PhraseQuery::new(query)
.with_column(Some(column))
.with_slop(slop)
.into(),
))
let mut query = PhraseQuery::new(query)
.with_column(Some(column))
.with_slop(slop);
if let Some(document_granularity) = document_granularity {
query = query.with_document_granularity(document_granularity);
}
Ok(Self(query.into()))
}
"BoostQuery" => {
let positive: Self = ob.getattr("positive")?.extract()?;
@@ -167,6 +186,13 @@ impl<'py> IntoPyObject<'py> for PyLanceDB<FtsQuery> {
kwargs.set_item("max_expansions", query.max_expansions)?;
kwargs.set_item::<_, &str>("operator", query.operator.into())?;
kwargs.set_item("prefix_length", query.prefix_length)?;
if let Some(document_granularity) = query.document_granularity {
let value = match document_granularity {
DocumentGranularity::Row => "row",
DocumentGranularity::ListElement => "list_element",
};
kwargs.set_item("document_granularity", value)?;
}
namespace
.getattr(intern!(py, "MatchQuery"))?
.call((query.terms, query.column.unwrap()), Some(&kwargs))
@@ -174,6 +200,13 @@ impl<'py> IntoPyObject<'py> for PyLanceDB<FtsQuery> {
FtsQuery::Phrase(query) => {
let kwargs = PyDict::new(py);
kwargs.set_item("slop", query.slop)?;
if let Some(document_granularity) = query.document_granularity {
let value = match document_granularity {
DocumentGranularity::Row => "row",
DocumentGranularity::ListElement => "list_element",
};
kwargs.set_item("document_granularity", value)?;
}
namespace
.getattr(intern!(py, "PhraseQuery"))?
.call((query.terms, query.column.unwrap()), Some(&kwargs))
+54 -1
View File
@@ -1476,7 +1476,7 @@ mod tests {
use crate::table::{AnyQuery, WriteOptions};
use arrow_array::{Int32Array, RecordBatch, StringArray};
use arrow_schema::{DataType, Field, Schema, SchemaRef};
use futures::{TryStreamExt, stream::once};
use futures::{TryStreamExt, future::try_join_all, stream::once};
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
@@ -1614,6 +1614,59 @@ mod tests {
);
}
#[tokio::test]
async fn test_concurrent_open_table_reuses_connection_object_store() {
let tempdir = tempdir().unwrap();
let uri = tempdir.path().to_str().unwrap();
let session = Arc::new(lance::session::Session::default());
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.clone()),
};
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 = session.store_registry().stats();
let opened_tables = try_join_all((0..32).map(|_| {
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();
let after = session.store_registry().stats();
assert_eq!(opened_tables.len(), 32);
assert_eq!(after.misses, before.misses);
assert_eq!(after.active_stores, before.active_stores);
assert!(after.hits >= before.hits + 32);
}
#[tokio::test]
async fn test_listing_database_root_ops_do_not_create_manifest() {
let tempdir = tempdir().unwrap();
+26
View File
@@ -186,6 +186,9 @@ pub struct PythonEnvironmentSpec {
pub kind: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub packages: Vec<String>,
/// Conda channels in priority order; conda environments only.
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub channels: Vec<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub path: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
@@ -583,3 +586,26 @@ impl RefreshColumnResult {
}
impl_json!(RefreshColumnResult);
#[cfg(test)]
mod conda_environment_tests {
use super::PythonEnvironmentSpec;
#[test]
fn conda_channels_round_trip_and_pip_stays_bare() {
let conda: PythonEnvironmentSpec = serde_json::from_str(
r#"{"kind":"conda","packages":["numpy"],"channels":["conda-forge"]}"#,
)
.unwrap();
assert_eq!(conda.channels, ["conda-forge"]);
assert!(
serde_json::to_string(&conda)
.unwrap()
.contains(r#""channels":["conda-forge"]"#)
);
let pip: PythonEnvironmentSpec =
serde_json::from_str(r#"{"kind":"pip","packages":["numpy"]}"#).unwrap();
assert!(!serde_json::to_string(&pip).unwrap().contains("channels"));
}
}
+1
View File
@@ -63,4 +63,5 @@ pub struct FmIndexBuilder {}
pub use lance_index::scalar::FullTextSearchQuery;
pub use lance_index::scalar::InvertedIndexParams as FtsIndexBuilder;
pub use lance_index::scalar::InvertedIndexParams;
pub use lance_index::scalar::inverted::DocumentGranularity;
pub use lance_index::scalar::inverted::query::*;
+4
View File
@@ -87,6 +87,10 @@ impl ServerVersion {
pub fn support_blobs(&self) -> bool {
self.0 >= semver::Version::new(0, 5, 0)
}
pub fn support_fts_document_granularity(&self) -> bool {
self.0 >= semver::Version::new(0, 6, 0)
}
}
pub const OPT_REMOTE_PREFIX: &str = "remote_database_";
+248 -35
View File
@@ -14,6 +14,7 @@ use crate::data::scannable::{PeekedScannable, Scannable, estimate_write_partitio
use crate::expr::expr_to_sql_string;
use crate::index::Index;
use crate::index::IndexStatistics;
use crate::index::scalar::FtsQuery;
use crate::index::waiter::wait_for_index;
use crate::job::Job;
use crate::query::{QueryFilter, QueryRequest, Select, VectorQueryRequest};
@@ -39,8 +40,8 @@ use crate::table::{
use crate::table::{AnyQuery, Filter, Predicate, PreprocessingOutput, TableStatistics};
use crate::utils::background_cache::BackgroundCache;
use crate::utils::{
MaxBatchLengthStream, TimeoutStream, resolve_arrow_field_path, supported_btree_data_type,
supported_vector_data_type,
MaxBatchLengthStream, TimeoutStream, resolve_arrow_field_path, resolve_arrow_fts_field_path,
supported_btree_data_type, supported_vector_data_type,
};
use crate::{DistanceType, Error};
use crate::{
@@ -89,6 +90,32 @@ const SCHEMA_CACHE_TTL: Duration = Duration::from_secs(30);
const SCHEMA_CACHE_REFRESH_WINDOW: Duration = Duration::from_secs(5);
const SCHEMA_SELECTOR_CHANGED: &str = "table selector changed while fetching schema";
fn fts_query_requires_document_granularity_support(query: &FtsQuery) -> bool {
match query {
FtsQuery::Match(query) => query
.document_granularity
.is_some_and(|granularity| granularity.is_list_element()),
FtsQuery::Phrase(query) => query
.document_granularity
.is_some_and(|granularity| granularity.is_list_element()),
FtsQuery::Boost(query) => {
fts_query_requires_document_granularity_support(&query.positive)
|| fts_query_requires_document_granularity_support(&query.negative)
}
FtsQuery::MultiMatch(query) => query.match_queries.iter().any(|query| {
query
.document_granularity
.is_some_and(|granularity| granularity.is_list_element())
}),
FtsQuery::Boolean(query) => query
.must
.iter()
.chain(&query.should)
.chain(&query.must_not)
.any(fts_query_requires_document_granularity_support),
}
}
/// Per-table state driving the freshness headers (`x-lancedb-min-version`,
/// `x-lancedb-min-timestamp`, and `x-lancedb-min-read-version`) sent on table
/// requests.
@@ -481,8 +508,21 @@ impl<S: HttpSend> RemoteTable<S> {
});
}
};
if matches!(
&index.index,
Index::FTS(params) if params.get_document_granularity().is_list_element()
) && !self.server_version.support_fts_document_granularity()
{
return Err(Error::NotSupported {
message: "FTS document granularity requires remote server version 0.6.0 or later"
.into(),
});
}
let schema = self.schema().await?;
let (canonical_column, field) = resolve_arrow_field_path(&schema, &column)?;
let (canonical_column, field) = match &index.index {
Index::FTS(_) => resolve_arrow_fts_field_path(&schema, &column)?,
_ => resolve_arrow_field_path(&schema, &column)?,
};
let mut body = serde_json::json!({
"column": canonical_column
});
@@ -518,7 +558,13 @@ impl<S: HttpSend> RemoteTable<S> {
Index::Bitmap(p) => ("BITMAP", Some(to_json(p)?)),
Index::LabelList(p) => ("LABEL_LIST", Some(to_json(p)?)),
Index::Fm(p) => ("FM", Some(to_json(p)?)),
Index::FTS(p) => ("FTS", Some(to_json(p)?)),
Index::FTS(p) => {
let mut params = to_json(p)?;
if p.get_document_granularity().is_list_element() {
params["document_granularity"] = "list_element".into();
}
("FTS", Some(params))
}
Index::Auto => {
if supported_vector_data_type(field.data_type()) {
body[METRIC_TYPE_KEY] =
@@ -1000,6 +1046,18 @@ impl<S: HttpSend> RemoteTable<S> {
});
}
let requires_document_granularity_support =
fts_query_requires_document_granularity_support(&full_text_search.query);
if requires_document_granularity_support
&& !self.server_version.support_fts_document_granularity()
{
return Err(Error::NotSupported {
message:
"FTS document granularity requires remote server version 0.6.0 or later"
.into(),
});
}
if self.server_version.support_structural_fts() {
body["full_text_query"] = serde_json::json!({
"query": full_text_search.query.clone(),
@@ -3678,7 +3736,7 @@ mod tests {
use arrow_schema::{DataType, Field, Schema};
use chrono::{DateTime, Utc};
use futures::{StreamExt, TryFutureExt, TryStreamExt, future::BoxFuture};
use lance_index::scalar::inverted::query::MatchQuery;
use lance_index::scalar::inverted::{DocumentGranularity, query::MatchQuery};
use lance_index::scalar::{FullTextSearchQuery, InvertedIndexParams};
use reqwest::Body;
use rstest::rstest;
@@ -3955,6 +4013,15 @@ mod tests {
DataType::Struct(vec![Field::new("text", DataType::Utf8, false)].into()),
false,
),
Field::new(
"docs",
DataType::List(Arc::new(Field::new(
"item",
DataType::Struct(vec![Field::new("content", DataType::Utf8, true)].into()),
true,
))),
true,
),
Field::new(
"meta-data",
DataType::Struct(vec![Field::new("user-id", DataType::Int32, false)].into()),
@@ -5691,7 +5758,7 @@ mod tests {
#[tokio::test]
async fn test_query_structured_fts() {
let table =
Table::new_with_handler_version("my_table", semver::Version::new(0, 3, 0), |request| {
Table::new_with_handler_version("my_table", semver::Version::new(0, 6, 0), |request| {
assert_eq!(request.method(), "POST");
assert_eq!(request.url().path(), "/v1/table/my_table/query/");
assert_eq!(
@@ -5712,6 +5779,7 @@ mod tests {
"max_expansions": 50,
"operator": "Or",
"prefix_length": 0,
"document_granularity": "list_element",
},
}
},
@@ -5741,6 +5809,7 @@ mod tests {
.full_text_search(FullTextSearchQuery::new_query(
MatchQuery::new("hello world".to_owned())
.with_column(Some("payload.text".to_owned()))
.with_document_granularity(DocumentGranularity::ListElement)
.into(),
))
.with_row_id()
@@ -5750,6 +5819,76 @@ mod tests {
.unwrap();
}
#[tokio::test]
async fn test_query_row_document_granularity_uses_structured_fts() {
let table =
Table::new_with_handler_version("my_table", semver::Version::new(0, 3, 0), |request| {
let body = request.body().unwrap().as_bytes().unwrap();
let body: serde_json::Value = serde_json::from_slice(body).unwrap();
assert_eq!(
body["full_text_query"]["query"]["match"]["document_granularity"],
"row"
);
let data = RecordBatch::try_new(
Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, false)])),
vec![Arc::new(Int32Array::from(vec![1]))],
)
.unwrap();
http::Response::builder()
.status(200)
.header(CONTENT_TYPE, ARROW_FILE_CONTENT_TYPE)
.body(write_ipc_file(&data))
.unwrap()
});
table
.query()
.full_text_search(FullTextSearchQuery::new_query(
MatchQuery::new("hello world".to_owned())
.with_column(Some("payload.text".to_owned()))
.with_document_granularity(DocumentGranularity::Row)
.into(),
))
.execute()
.await
.unwrap();
}
#[rstest]
#[case(DEFAULT_SERVER_VERSION.clone())]
#[case(semver::Version::new(0, 3, 0))]
#[case(semver::Version::new(0, 5, 0))]
#[tokio::test]
async fn test_query_document_granularity_requires_server_support(
#[case] version: semver::Version,
) {
let table =
Table::new_with_handler_version("my_table", version, |_| -> http::Response<String> {
panic!("unsupported remote query must fail before sending a request")
});
let result = table
.query()
.full_text_search(FullTextSearchQuery::new_query(
MatchQuery::new("hello world".to_owned())
.with_column(Some("payload.text".to_owned()))
.with_document_granularity(DocumentGranularity::ListElement)
.into(),
))
.execute()
.await;
let err = match result {
Ok(_) => panic!("legacy remote query unexpectedly succeeded"),
Err(err) => err,
};
assert!(
err.to_string()
.contains("document granularity requires remote server version 0.6.0 or later")
);
}
#[rstest]
#[case(DEFAULT_SERVER_VERSION.clone())]
#[case(semver::Version::new(0, 2, 0))]
@@ -6028,40 +6167,56 @@ mod tests {
"CAT".to_string(),
]))),
),
(
"FTS",
{
let mut body = serde_json::to_value(InvertedIndexParams::default()).unwrap();
body["document_granularity"] = "list_element".into();
body
},
Index::FTS(
InvertedIndexParams::default()
.document_granularity(DocumentGranularity::ListElement),
),
),
];
for (index_type, expected_body, index) in cases {
let table = Table::new_with_handler("my_table", move |request| {
assert_eq!(request.method(), "POST");
match request.url().path() {
"/v1/table/my_table/describe/" => {
let schema = Schema::new(vec![Field::new("a", DataType::Int32, false)]);
http::Response::builder()
.status(200)
.body(describe_response(&schema))
.unwrap()
}
"/v1/table/my_table/create_index/" => {
assert_eq!(
request.headers().get("Content-Type").unwrap(),
JSON_CONTENT_TYPE
);
let body = request.body().unwrap().as_bytes().unwrap();
let body: serde_json::Value = serde_json::from_slice(body).unwrap();
let mut expected_body = expected_body.clone();
expected_body["column"] = "a".into();
expected_body[INDEX_TYPE_KEY] = index_type.into();
let table = Table::new_with_handler_version(
"my_table",
semver::Version::new(0, 6, 0),
move |request| {
assert_eq!(request.method(), "POST");
match request.url().path() {
"/v1/table/my_table/describe/" => {
let schema = Schema::new(vec![Field::new("a", DataType::Int32, false)]);
http::Response::builder()
.status(200)
.body(describe_response(&schema))
.unwrap()
}
"/v1/table/my_table/create_index/" => {
assert_eq!(
request.headers().get("Content-Type").unwrap(),
JSON_CONTENT_TYPE
);
let body = request.body().unwrap().as_bytes().unwrap();
let body: serde_json::Value = serde_json::from_slice(body).unwrap();
let mut expected_body = expected_body.clone();
expected_body["column"] = "a".into();
expected_body[INDEX_TYPE_KEY] = index_type.into();
assert_eq!(body, expected_body);
assert_eq!(body, expected_body);
http::Response::builder()
.status(200)
.body("{}".to_string())
.unwrap()
http::Response::builder()
.status(200)
.body("{}".to_string())
.unwrap()
}
path => panic!("Unexpected path: {}", path),
}
path => panic!("Unexpected path: {}", path),
}
});
},
);
table.create_index(&["a"], index).execute().await.unwrap();
}
@@ -6320,6 +6475,19 @@ mod tests {
body["index_type"] = "FTS".into();
body
},
{
let mut body = serde_json::to_value(InvertedIndexParams::default()).unwrap();
body["column"] = "docs.content".into();
body["index_type"] = "FTS".into();
body
},
{
let mut body = serde_json::to_value(InvertedIndexParams::default()).unwrap();
body["column"] = "docs.content".into();
body["index_type"] = "FTS".into();
body["document_granularity"] = "list_element".into();
body
},
json!({
"column": "`meta-data`.`user-id`",
"index_type": "BTREE",
@@ -6330,7 +6498,7 @@ mod tests {
}),
]);
let request_idx = Arc::new(AtomicUsize::new(0));
let table = Table::new_with_handler("my_table", {
let table = Table::new_with_handler_version("my_table", semver::Version::new(0, 6, 0), {
let schema = schema.clone();
let expected_requests = expected_requests.clone();
let request_idx = request_idx.clone();
@@ -6395,6 +6563,22 @@ mod tests {
.execute()
.await
.unwrap();
table
.create_index(&["Docs.Content"], Index::FTS(Default::default()))
.execute()
.await
.unwrap();
table
.create_index(
&["Docs.Content"],
Index::FTS(
InvertedIndexParams::default()
.document_granularity(DocumentGranularity::ListElement),
),
)
.execute()
.await
.unwrap();
table
.create_index(&["`META-DATA`.`USER-ID`"], Index::BTree(Default::default()))
.execute()
@@ -6409,6 +6593,35 @@ mod tests {
assert_eq!(request_idx.load(Ordering::SeqCst), expected_requests.len());
}
#[tokio::test]
async fn test_create_list_element_fts_requires_server_support() {
let table = Table::new_with_handler_version(
"my_table",
semver::Version::new(0, 5, 0),
|_| -> http::Response<String> {
panic!("unsupported index creation must fail before sending a request")
},
);
let result = table
.create_index(
&["docs.content"],
Index::FTS(
InvertedIndexParams::default()
.document_granularity(DocumentGranularity::ListElement),
),
)
.execute()
.await;
assert!(
result
.unwrap_err()
.to_string()
.contains("document granularity requires remote server version 0.6.0 or later")
);
}
#[tokio::test]
async fn test_list_indices() {
let schema = Schema::new(vec![
+11 -2
View File
@@ -59,7 +59,9 @@ use crate::index::{IndexConfig, IndexStatisticsImpl, IndexType};
use crate::job::Job;
use crate::query::{IntoQueryVector, Query, QueryExecutionOptions, TakeQuery, VectorQuery};
use crate::table::datafusion::insert::InsertExec;
use crate::utils::{PatchReadParam, PatchWriteParam, resolve_arrow_field_path};
use crate::utils::{
PatchReadParam, PatchWriteParam, public_fts_field_path_by_id, resolve_arrow_field_path,
};
use self::dataset::DatasetConsistencyWrapper;
use self::merge::MergeInsertBuilder;
@@ -3567,7 +3569,14 @@ impl BaseTable for NativeTable {
let field_ids = idx_desc.field_ids();
let mut columns = Vec::with_capacity(field_ids.len());
for field_id in field_ids {
let field_path = match dataset.schema().field_path(*field_id as i32) {
let field_path = match if index_type == crate::index::IndexType::FTS {
public_fts_field_path_by_id(dataset.schema(), *field_id as i32)
} else {
dataset
.schema()
.field_path(*field_id as i32)
.map_err(Into::into)
} {
Ok(field_path) => field_path,
Err(e) => {
log::warn!(
+50 -4
View File
@@ -28,8 +28,9 @@ pub(super) type PreparedIndex = (String, Box<dyn lance::index::IndexParams>, Ind
use crate::index::Index;
use crate::index::vector::{VectorIndex, suggested_num_sub_vectors};
use crate::utils::{
supported_bitmap_data_type, supported_btree_data_type, supported_fm_data_type,
supported_fts_data_type, supported_label_list_data_type, supported_vector_data_type,
resolve_lance_fts_field_path, supported_bitmap_data_type, supported_btree_data_type,
supported_fm_data_type, supported_fts_data_type, supported_label_list_data_type,
supported_vector_data_type,
};
use super::NativeTable;
@@ -122,7 +123,20 @@ impl NativeTable {
}
self.dataset.ensure_mutable()?;
let dataset = self.dataset.get().await?;
let (column, field) = Self::resolve_index_field(dataset.schema(), &opts.columns[0])?;
let (column, field) = if let Index::FTS(params) = &opts.index {
let resolved = resolve_lance_fts_field_path(dataset.schema(), &opts.columns[0])?;
if params.get_document_granularity().is_list_element() && resolved.list_depth == 0 {
return Err(Error::InvalidInput {
message: format!(
"FTS field path '{}' has no List layer and cannot use ListElement document granularity",
resolved.canonical_path
),
});
}
(resolved.canonical_path, resolved.field)
} else {
Self::resolve_index_field(dataset.schema(), &opts.columns[0])?
};
let params = self.make_index_params(&field, opts.index.clone()).await?;
let index_type = self.get_index_type_for_field(&field, &opts.index);
Ok((column, params, index_type))
@@ -436,7 +450,7 @@ mod tests {
use crate::connection::ConnectBuilder;
use crate::index::Index;
use crate::index::scalar::{
BTreeIndexBuilder, BitmapIndexBuilder, FmIndexBuilder, FtsIndexBuilder,
BTreeIndexBuilder, BitmapIndexBuilder, DocumentGranularity, FmIndexBuilder, FtsIndexBuilder,
};
use crate::index::vector::{
IvfHnswFlatIndexBuilder, IvfHnswPqIndexBuilder, IvfHnswSqIndexBuilder,
@@ -553,6 +567,38 @@ mod tests {
job.cancel().await.unwrap();
}
#[tokio::test]
async fn test_execute_async_validates_fts_input_before_starting_job() {
let conn = connect("memory://").execute().await.unwrap();
let batch =
record_batch!(("id", Int32, [1, 2]), ("text", Utf8, ["alpha", "beta"])).unwrap();
let table = conn.create_table("t", batch).execute().await.unwrap();
let missing = table
.create_index(&["missing"], Index::FTS(FtsIndexBuilder::default()))
.execute_async()
.await;
assert!(missing.is_err());
let invalid_type = table
.create_index(&["id"], Index::FTS(FtsIndexBuilder::default()))
.execute_async()
.await;
assert!(invalid_type.is_err());
let invalid_granularity = table
.create_index(
&["text"],
Index::FTS(
FtsIndexBuilder::default()
.document_granularity(DocumentGranularity::ListElement),
),
)
.execute_async()
.await;
assert!(invalid_granularity.is_err());
}
/// Concurrent waiters, and a wait issued after the job settled, all
/// succeed once the build does.
#[tokio::test]
+183
View File
@@ -225,6 +225,159 @@ pub(crate) fn resolve_arrow_field_path(schema: &Schema, column: &str) -> Result<
Ok((canonical_path, Field::from(*field)))
}
pub(crate) struct ResolvedFtsField {
pub canonical_path: String,
pub field: Field,
pub list_depth: usize,
}
/// Canonicalize a public FTS field path while keeping Arrow list item names hidden.
pub(crate) fn resolve_lance_fts_field_path(
schema: &lance_core::datatypes::Schema,
column: &str,
) -> Result<ResolvedFtsField> {
let names =
lance_core::datatypes::parse_field_path(column).map_err(|e| Error::InvalidInput {
message: format!("Invalid field path `{}`: {}", column, e),
})?;
let (root_name, remaining_names) = names.split_first().ok_or_else(|| Error::InvalidInput {
message: "FTS field path cannot be empty".to_string(),
})?;
let mut field = schema
.fields
.iter()
.find(|field| field.name == *root_name)
.or_else(|| {
schema
.fields
.iter()
.find(|field| field.name.eq_ignore_ascii_case(root_name))
})
.ok_or_else(|| fts_field_not_found(schema, column))?;
let mut canonical_names = vec![field.name.clone()];
let mut list_depth = 0;
for name in remaining_names {
while matches!(
field.data_type(),
DataType::List(_) | DataType::LargeList(_)
) {
list_depth += 1;
field = field.children.first().ok_or_else(|| Error::Schema {
message: format!(
"FTS field path `{}` has a list without an item field",
column
),
})?;
}
if !matches!(field.data_type(), DataType::Struct(_)) {
return Err(fts_field_not_found(schema, column));
}
field = field
.children
.iter()
.find(|field| field.name == *name)
.or_else(|| {
field
.children
.iter()
.find(|field| field.name.eq_ignore_ascii_case(name))
})
.ok_or_else(|| fts_field_not_found(schema, column))?;
canonical_names.push(field.name.clone());
}
let mut terminal = field;
while matches!(
terminal.data_type(),
DataType::List(_) | DataType::LargeList(_)
) {
list_depth += 1;
terminal = terminal.children.first().ok_or_else(|| Error::Schema {
message: format!(
"FTS field path `{}` has a list without an item field",
column
),
})?;
}
let canonical_path = lance_core::datatypes::format_field_path(
&canonical_names
.iter()
.map(String::as_str)
.collect::<Vec<_>>(),
);
Ok(ResolvedFtsField {
canonical_path,
field: Field::from(field),
list_depth,
})
}
fn fts_field_not_found(schema: &lance_core::datatypes::Schema, column: &str) -> Error {
Error::Schema {
message: format!(
"Field path `{}` not found in schema. Available field paths: {}",
column,
schema.field_paths().join(", ")
),
}
}
fn find_public_fts_field_path_by_id(
field: &lance_core::datatypes::Field,
field_id: i32,
path: &mut Vec<String>,
) -> bool {
if field.id == field_id {
return true;
}
match field.data_type() {
DataType::List(_) | DataType::LargeList(_) => field
.children
.first()
.is_some_and(|child| find_public_fts_field_path_by_id(child, field_id, path)),
DataType::Struct(_) => field.children.iter().any(|child| {
path.push(child.name.clone());
let found = find_public_fts_field_path_by_id(child, field_id, path);
if !found {
path.pop();
}
found
}),
_ => false,
}
}
pub(crate) fn public_fts_field_path_by_id(
schema: &lance_core::datatypes::Schema,
field_id: i32,
) -> Result<String> {
for root in &schema.fields {
let mut path = vec![root.name.clone()];
if find_public_fts_field_path_by_id(root, field_id, &mut path) {
return Ok(lance_core::datatypes::format_field_path(
&path.iter().map(String::as_str).collect::<Vec<_>>(),
));
}
}
Err(Error::Schema {
message: format!("Field id `{}` not found in schema", field_id),
})
}
pub(crate) fn resolve_arrow_fts_field_path(
schema: &Schema,
column: &str,
) -> Result<(String, Field)> {
let lance_schema =
lance_core::datatypes::Schema::try_from(schema).map_err(|e| Error::Schema {
message: format!("Invalid schema: {}", e),
})?;
let resolved = resolve_lance_fts_field_path(&lance_schema, column)?;
Ok((resolved.canonical_path, resolved.field))
}
pub fn supported_btree_data_type(dtype: &DataType) -> bool {
dtype.is_integer()
|| dtype.is_floating()
@@ -480,6 +633,36 @@ mod tests {
use super::*;
#[test]
fn test_public_fts_field_path_prefers_exact_case() {
let text_list = || {
DataType::List(Arc::new(Field::new(
"item",
DataType::Struct(vec![Field::new("content", DataType::Utf8, true)].into()),
true,
)))
};
let schema = Schema::new(vec![
Field::new("Docs", text_list(), true),
Field::new("docs", text_list(), true),
]);
let (path, _) = resolve_arrow_fts_field_path(&schema, "docs.content").unwrap();
assert_eq!(path, "docs.content");
let lance_schema = lance_core::datatypes::Schema::try_from(&schema).unwrap();
let field_id = lance_schema
.resolve_case_insensitive("docs.item.content")
.unwrap()
.last()
.unwrap()
.id;
assert_eq!(
public_fts_field_path_by_id(&lance_schema, field_id).unwrap(),
"docs.content"
);
}
#[test]
fn test_guess_default_column() {
let schema_no_vector = Schema::new(vec![