diff --git a/docs/src/python/python.md b/docs/src/python/python.md index 3cbeee6f0..3cb996a15 100644 --- a/docs/src/python/python.md +++ b/docs/src/python/python.md @@ -159,6 +159,8 @@ and combined with [BooleanQuery][lancedb.query.BooleanQuery]. ::: lancedb.query.FullTextOperator +::: lancedb.query.DocumentGranularity + ::: lancedb.query.Occur ## Embeddings diff --git a/python/python/lancedb/functions.py b/python/python/lancedb/functions.py index 237ed73c2..8a19a9d37 100644 --- a/python/python/lancedb/functions.py +++ b/python/python/lancedb/functions.py @@ -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: diff --git a/python/python/lancedb/index.py b/python/python/lancedb/index.py index d2b63baf6..948342887 100644 --- a/python/python/lancedb/index.py +++ b/python/python/lancedb/index.py @@ -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 diff --git a/python/python/lancedb/query.py b/python/python/lancedb/query.py index 51b4dad4e..7e7d313cd 100644 --- a/python/python/lancedb/query.py +++ b/python/python/lancedb/query.py @@ -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 diff --git a/python/python/lancedb/remote/table.py b/python/python/lancedb/remote/table.py index d0bf9f67a..d9139396b 100644 --- a/python/python/lancedb/remote/table.py +++ b/python/python/lancedb/remote/table.py @@ -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( diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index aa3021317..d16cf0128 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -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( diff --git a/python/python/tests/test_first_class_function_slice2.py b/python/python/tests/test_first_class_function_slice2.py index 55257d322..7ce6b6b91 100644 --- a/python/python/tests/test_first_class_function_slice2.py +++ b/python/python/tests/test_first_class_function_slice2.py @@ -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: diff --git a/python/python/tests/test_fts.py b/python/python/tests/test_fts.py index 625198d92..e5129dd9c 100644 --- a/python/python/tests/test_fts.py +++ b/python/python/tests/test_fts.py @@ -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() diff --git a/python/python/tests/test_remote_db.py b/python/python/tests/test_remote_db.py index da6347142..01e2cc4c5 100644 --- a/python/python/tests/test_remote_db.py +++ b/python/python/tests/test_remote_db.py @@ -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: diff --git a/python/python/tests/test_s3.py b/python/python/tests/test_s3.py index 256ccb1d4..70b423adb 100644 --- a/python/python/tests/test_s3.py +++ b/python/python/tests/test_s3.py @@ -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"]) diff --git a/python/src/index.rs b/python/src/index.rs index a5ca63c68..54b15f55e 100644 --- a/python/src/index.rs +++ b/python/src/index.rs @@ -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>) -> PyResult, num_workers: Option, + 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, diff --git a/python/src/query.rs b/python/src/query.rs index 1020c822d..82c05d056 100644 --- a/python/src/query.rs +++ b/python/src/query.rs @@ -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 { let max_expansions = ob.getattr("max_expansions")?.extract()?; let operator = ob.getattr("operator")?.extract::()?; let prefix_length = ob.getattr("prefix_length")?.extract()?; + let document_granularity = ob + .getattr("document_granularity")? + .extract::>()? + .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 { .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::>()? + .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 { 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::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)) diff --git a/rust/lancedb/src/database/listing.rs b/rust/lancedb/src/database/listing.rs index 17dc82756..064b5d28f 100644 --- a/rust/lancedb/src/database/listing.rs +++ b/rust/lancedb/src/database/listing.rs @@ -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, + 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(); diff --git a/rust/lancedb/src/function.rs b/rust/lancedb/src/function.rs index 52c70a4b1..5366d984e 100644 --- a/rust/lancedb/src/function.rs +++ b/rust/lancedb/src/function.rs @@ -186,6 +186,9 @@ pub struct PythonEnvironmentSpec { pub kind: String, #[serde(default, skip_serializing_if = "Vec::is_empty")] pub packages: Vec, + /// Conda channels in priority order; conda environments only. + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub channels: Vec, #[serde(default, skip_serializing_if = "Option::is_none")] pub path: Option, #[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")); + } +} diff --git a/rust/lancedb/src/index/scalar.rs b/rust/lancedb/src/index/scalar.rs index 10d835bb1..dba05b776 100644 --- a/rust/lancedb/src/index/scalar.rs +++ b/rust/lancedb/src/index/scalar.rs @@ -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::*; diff --git a/rust/lancedb/src/remote/db.rs b/rust/lancedb/src/remote/db.rs index 3f4216bc8..da9a4b09b 100644 --- a/rust/lancedb/src/remote/db.rs +++ b/rust/lancedb/src/remote/db.rs @@ -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_"; diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index 34b22d138..721b8b955 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -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 RemoteTable { }); } }; + 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 RemoteTable { 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 RemoteTable { }); } + 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 { + 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 { + 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![ diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index f0c08f0ba..ecb83055b 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -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!( diff --git a/rust/lancedb/src/table/create_index.rs b/rust/lancedb/src/table/create_index.rs index e373522bc..c7d6b5675 100644 --- a/rust/lancedb/src/table/create_index.rs +++ b/rust/lancedb/src/table/create_index.rs @@ -28,8 +28,9 @@ pub(super) type PreparedIndex = (String, Box, 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] diff --git a/rust/lancedb/src/utils/mod.rs b/rust/lancedb/src/utils/mod.rs index 8bd306988..07d1836a1 100644 --- a/rust/lancedb/src/utils/mod.rs +++ b/rust/lancedb/src/utils/mod.rs @@ -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 { + 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::>(), + ); + 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, +) -> 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 { + 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::>(), + )); + } + } + 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![