use std::collections::hash_map::Entry;
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use arrow_array::builder::{FixedSizeListBuilder, Float32Builder};
use arrow_array::cast::AsArray;
use arrow_array::types::Float32Type;
use arrow_array::{Array, FixedSizeListArray, RecordBatch};
use arrow_schema::{DataType, SchemaRef};
use datafusion::common::ScalarValue;
use datafusion::logical_expr::Operator;
use datafusion::physical_plan::limit::GlobalLimitExec;
use datafusion::physical_plan::{ExecutionPlan, SendableRecordBatchStream};
use datafusion::prelude::{Expr, SessionContext};
use futures::TryStreamExt;
use lance_core::{Error, Result, is_system_column};
use lance_datafusion::expr::safe_coerce_scalar;
use lance_index::scalar::FullTextSearchQuery;
use lance_linalg::distance::DistanceType;
use uuid::Uuid;
use super::collector::{InMemoryMemTableRef, InMemoryMemTables, LsmDataSourceCollector};
use super::data_source::{FreshTierWatermark, ShardSnapshot};
use super::planner::LsmScanPlanner;
use super::point_lookup::LsmPointLookupPlanner;
use super::projection::validate_projection_names;
use super::sstable_cache::{DatasetCache, SsTableWarmer};
use crate::dataset::Dataset;
use crate::dataset::mem_wal::util::derived_store_params;
use crate::session::Session;
use lance_io::object_store::ObjectStoreParams;
#[derive(Debug, Clone)]
struct LsmVectorQuery {
column: String,
key: Arc<dyn Array>,
k: usize,
nprobes: usize,
refine: bool,
metric_type: Option<DistanceType>,
}
fn extract_pk_point_keys(
filter: &Expr,
pk_col: &str,
pk_type: &DataType,
) -> Option<Vec<ScalarValue>> {
match filter {
Expr::BinaryExpr(b) if matches!(b.op, Operator::Eq) => {
match (b.left.as_ref(), b.right.as_ref()) {
(Expr::Column(c), Expr::Literal(lit, _))
| (Expr::Literal(lit, _), Expr::Column(c))
if c.name == pk_col =>
{
Some(vec![lit.clone()])
}
_ => None,
}
}
Expr::InList(in_list) if !in_list.negated => {
let Expr::Column(c) = in_list.expr.as_ref() else {
return None;
};
if c.name != pk_col {
return None;
}
let mut vals = Vec::with_capacity(in_list.list.len());
let mut seen = HashSet::with_capacity(in_list.list.len());
for e in &in_list.list {
let Expr::Literal(lit, _) = e else {
return None; };
let identity = safe_coerce_scalar(lit, pk_type)?;
if seen.insert(identity) {
vals.push(lit.clone());
}
}
(!vals.is_empty()).then_some(vals)
}
_ => None,
}
}
enum BaseSource {
Table(Arc<Dataset>),
PathOnly(String),
}
fn vector_dim(schema: &arrow_schema::Schema, column: &str) -> Result<i32> {
let field = schema
.field_with_name(column)
.map_err(|_| Error::invalid_input(format!("vector column '{}' not found", column)))?;
match field.data_type() {
DataType::FixedSizeList(_, dim) => Ok(*dim),
other => Err(Error::invalid_input(format!(
"column '{}' is not a fixed-size-list vector (got {:?})",
column, other
))),
}
}
fn key_to_fsl(key: &dyn Array, dim: i32) -> Result<FixedSizeListArray> {
if let Some(fsl) = key.as_any().downcast_ref::<FixedSizeListArray>() {
if fsl.len() != 1 {
return Err(Error::invalid_input(format!(
"LSM vector search supports a single query vector, got {} rows",
fsl.len()
)));
}
if fsl.value_length() != dim {
return Err(Error::invalid_input(format!(
"query vector dimension {} does not match column dimension {}",
fsl.value_length(),
dim
)));
}
return Ok(fsl.clone());
}
let values = key
.as_primitive_opt::<Float32Type>()
.ok_or_else(|| Error::invalid_input("query vector must be Float32".to_string()))?;
if values.len() != dim as usize {
return Err(Error::invalid_input(format!(
"query vector dimension {} does not match column dimension {}",
values.len(),
dim
)));
}
let mut builder =
FixedSizeListBuilder::with_capacity(Float32Builder::with_capacity(dim as usize), dim, 1);
builder.values().append_slice(values.values());
builder.append(true);
Ok(builder.finish())
}
pub struct LsmScanner {
base: BaseSource,
schema: SchemaRef,
shard_snapshots: Vec<ShardSnapshot>,
in_memory_memtables: HashMap<Uuid, InMemoryMemTables>,
projection: Option<Vec<String>>,
filter: Option<Expr>,
limit: Option<usize>,
offset: Option<usize>,
nearest: Option<LsmVectorQuery>,
full_text_query: Option<FullTextSearchQuery>,
with_row_address: bool,
with_memtable_gen: bool,
pk_columns: Vec<String>,
session: Option<Arc<Session>>,
store_params: Option<ObjectStoreParams>,
sstable_cache: Option<Arc<dyn DatasetCache>>,
warmer: Option<Arc<dyn SsTableWarmer>>,
overfetch_factor: Option<f64>,
}
impl LsmScanner {
pub fn new(
base_table: Arc<Dataset>,
shard_snapshots: Vec<ShardSnapshot>,
pk_columns: Vec<String>,
) -> Self {
let lance_schema = base_table.schema();
let arrow_schema: arrow_schema::Schema = lance_schema.into();
let session = Some(base_table.session());
let store_params = base_table.store_params().map(derived_store_params);
Self {
base: BaseSource::Table(base_table),
schema: Arc::new(arrow_schema),
shard_snapshots,
in_memory_memtables: HashMap::new(),
projection: None,
filter: None,
limit: None,
offset: None,
nearest: None,
full_text_query: None,
with_row_address: false,
with_memtable_gen: false,
pk_columns,
session,
store_params,
sstable_cache: None,
warmer: None,
overfetch_factor: None,
}
}
pub fn without_base_table(
schema: SchemaRef,
base_path: impl Into<String>,
shard_snapshots: Vec<ShardSnapshot>,
pk_columns: Vec<String>,
) -> Self {
Self {
base: BaseSource::PathOnly(base_path.into()),
schema,
shard_snapshots,
in_memory_memtables: HashMap::new(),
projection: None,
filter: None,
limit: None,
offset: None,
nearest: None,
full_text_query: None,
with_row_address: false,
with_memtable_gen: false,
pk_columns,
session: None,
store_params: None,
sstable_cache: None,
warmer: None,
overfetch_factor: None,
}
}
pub fn with_active_memtable(mut self, shard_id: Uuid, memtable: InMemoryMemTableRef) -> Self {
match self.in_memory_memtables.entry(shard_id) {
Entry::Occupied(mut e) => e.get_mut().active = memtable,
Entry::Vacant(e) => {
e.insert(InMemoryMemTables {
active: memtable,
frozen: Vec::new(),
});
}
}
self
}
pub fn with_in_memory_memtables(
mut self,
shard_id: Uuid,
memtables: InMemoryMemTables,
) -> Self {
self.in_memory_memtables.insert(shard_id, memtables);
self
}
pub fn with_session(mut self, session: Arc<Session>) -> Self {
self.session = Some(session);
self
}
pub fn with_store_params(mut self, store_params: ObjectStoreParams) -> Self {
self.store_params = Some(derived_store_params(&store_params));
self
}
pub fn with_sstable_cache(mut self, cache: Arc<dyn DatasetCache>) -> Self {
self.sstable_cache = Some(cache);
self
}
pub fn with_warmer(mut self, warmer: Arc<dyn SsTableWarmer>) -> Self {
self.warmer = Some(warmer);
self
}
pub fn with_overfetch_factor(mut self, factor: f64) -> Self {
self.overfetch_factor = Some(factor);
self
}
pub fn project<T: AsRef<str>>(mut self, columns: &[T]) -> Result<Self> {
self.projection = Some(columns.iter().map(|s| s.as_ref().to_string()).collect());
Ok(self)
}
pub fn filter(mut self, filter_expr: &str) -> Result<Self> {
let expr = super::parse_filter_expr(self.schema.as_ref(), filter_expr)?;
self.filter = Some(expr);
Ok(self)
}
pub fn filter_expr(mut self, expr: Expr) -> Self {
self.filter = Some(expr);
self
}
pub fn limit(mut self, limit: Option<i64>, offset: Option<i64>) -> Result<Self> {
if let Some(value) = limit
&& value < 0
{
return Err(Error::invalid_input(
"limit must be non-negative".to_string(),
));
}
if let Some(value) = offset
&& value < 0
{
return Err(Error::invalid_input(
"offset must be non-negative".to_string(),
));
}
self.limit = limit.map(|value| value as usize);
self.offset = offset.map(|value| value as usize);
Ok(self)
}
pub fn nearest(mut self, column: &str, key: &dyn Array, k: usize) -> Result<Self> {
if k == 0 {
return Err(Error::invalid_input("k must be positive".to_string()));
}
if key.is_empty() {
return Err(Error::invalid_input(
"query vector must have non-zero length".to_string(),
));
}
self.nearest = Some(LsmVectorQuery {
column: column.to_string(),
key: key.slice(0, key.len()),
k,
nprobes: 1,
refine: false,
metric_type: None,
});
Ok(self)
}
pub fn nprobes(mut self, nprobes: usize) -> Self {
if let Some(q) = self.nearest.as_mut() {
q.nprobes = nprobes;
}
self
}
pub fn refine(mut self, refine_factor: u32) -> Self {
if let Some(q) = self.nearest.as_mut() {
q.refine = refine_factor > 0;
}
self
}
pub fn distance_metric(mut self, metric: DistanceType) -> Self {
if let Some(q) = self.nearest.as_mut() {
q.metric_type = Some(metric);
}
self
}
pub fn with_row_address(mut self) -> Self {
self.with_row_address = true;
self
}
pub fn with_memtable_gen(mut self) -> Self {
self.with_memtable_gen = true;
self
}
pub fn schema(&self) -> SchemaRef {
self.schema.clone()
}
fn projection_has_system_columns(&self) -> bool {
self.projection
.as_ref()
.map(|p| p.iter().any(|c| is_system_column(c)))
.unwrap_or(false)
}
pub async fn create_plan(&self) -> Result<Arc<dyn ExecutionPlan>> {
if self.nearest.is_some() && self.full_text_query.is_some() {
return Err(Error::invalid_input(
"LSM scanner does not support combined vector and full-text search".to_string(),
));
}
if self.nearest.is_some() {
return self.plan_vector().await;
}
if self.full_text_query.is_some() {
return self.plan_fts().await;
}
self.plan_scan().await
}
fn apply_limit_offset(&self, plan: Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
let skip = self.offset.unwrap_or(0);
if skip == 0 && self.limit.is_none() {
return plan;
}
Arc::new(GlobalLimitExec::new(plan, skip, self.limit))
}
async fn plan_vector(&self) -> Result<Arc<dyn ExecutionPlan>> {
let nearest = self
.nearest
.as_ref()
.expect("plan_vector requires a nearest query");
let base_schema = self.schema();
let dim = vector_dim(base_schema.as_ref(), &nearest.column)?;
let query_fsl = key_to_fsl(nearest.key.as_ref(), dim)?;
let distance_type = nearest.metric_type.unwrap_or(DistanceType::L2);
let collector = self.build_collector();
let mut planner = super::LsmVectorSearchPlanner::new(
collector,
self.pk_columns.clone(),
base_schema,
nearest.column.clone(),
distance_type,
)
.with_filter(self.filter.clone());
if let BaseSource::Table(dataset) = &self.base {
planner = planner.with_dataset(dataset.clone());
}
if let Some(session) = &self.session {
planner = planner.with_session(session.clone());
}
if let Some(store_params) = &self.store_params {
planner = planner.with_store_params(store_params.clone());
}
if let Some(cache) = &self.sstable_cache {
planner = planner.with_sstable_cache(cache.clone());
}
if let Some(warmer) = &self.warmer {
planner = planner.with_warmer(warmer.clone());
}
let per_source_k = nearest.k.saturating_add(self.offset.unwrap_or(0));
let overfetch_factor = self.overfetch_factor.unwrap_or(1.0);
let plan = planner
.plan_search(
&query_fsl,
per_source_k,
nearest.nprobes,
self.projection.as_deref(),
nearest.refine,
overfetch_factor,
)
.await?;
Ok(self.apply_limit_offset(plan))
}
async fn plan_fts(&self) -> Result<Arc<dyn ExecutionPlan>> {
let query = self
.full_text_query
.as_ref()
.expect("plan_fts requires a full-text query");
let columns: Vec<String> = query.columns().into_iter().collect();
if columns.len() > 1 {
return Err(Error::invalid_input(
"LSM full-text search supports a single column".to_string(),
));
}
let column = columns.into_iter().next().ok_or_else(|| {
Error::invalid_input(
"full_text_search requires a column; set it with `FullTextSearchQuery::with_column`"
.to_string(),
)
})?;
let base_schema = self.schema();
let query_limit = query
.limit
.map(|limit| {
if limit < 0 {
Err(Error::invalid_input(
"full-text search limit must be non-negative".to_string(),
))
} else {
Ok(limit as usize)
}
})
.transpose()?;
let source_limit = match query_limit {
Some(limit) => Some(limit),
None => self
.limit
.map(|limit| limit.saturating_add(self.offset.unwrap_or(0))),
};
let collector = self.build_collector();
let mut planner =
super::LsmFtsSearchPlanner::new(collector, self.pk_columns.clone(), base_schema)
.with_filter(self.filter.clone());
if let Some(session) = &self.session {
planner = planner.with_session(session.clone());
}
if let Some(store_params) = &self.store_params {
planner = planner.with_store_params(store_params.clone());
}
if let Some(cache) = &self.sstable_cache {
planner = planner.with_sstable_cache(cache.clone());
}
if let Some(warmer) = &self.warmer {
planner = planner.with_warmer(warmer.clone());
}
if let Some(factor) = self.overfetch_factor {
planner = planner.with_overfetch_factor(factor);
}
let plan = planner
.plan_search(
&column,
query.clone(),
source_limit,
self.projection.as_deref(),
)
.await?;
Ok(self.apply_limit_offset(plan))
}
async fn plan_scan(&self) -> Result<Arc<dyn ExecutionPlan>> {
let collector = self.build_collector();
let base_schema = self.schema();
validate_projection_names(self.projection.as_deref(), &base_schema, &[])?;
if self.pk_columns.len() == 1
&& self.offset.is_none()
&& !self.with_memtable_gen
&& !self.with_row_address
&& !self.projection_has_system_columns()
&& let Some(filter) = &self.filter
&& let Ok(pk_field) = base_schema.field_with_name(&self.pk_columns[0])
&& let Some(keys) =
extract_pk_point_keys(filter, &self.pk_columns[0], pk_field.data_type())
{
let mut planner =
LsmPointLookupPlanner::new(collector, self.pk_columns.clone(), base_schema);
if let Some(session) = &self.session {
planner = planner.with_session(session.clone());
}
if let Some(store_params) = &self.store_params {
planner = planner.with_store_params(store_params.clone());
}
if let Some(cache) = &self.sstable_cache {
planner = planner.with_sstable_cache(cache.clone());
}
if let Some(warmer) = &self.warmer {
planner = planner.with_warmer(warmer.clone());
}
let plan = planner
.plan_point_lookup(&keys, self.projection.as_deref())
.await?;
return Ok(match self.limit {
Some(n) => Arc::new(GlobalLimitExec::new(plan, 0, Some(n))),
None => plan,
});
}
let mut planner = LsmScanPlanner::new(collector, self.pk_columns.clone(), base_schema);
if let Some(session) = &self.session {
planner = planner.with_session(session.clone());
}
if let Some(store_params) = &self.store_params {
planner = planner.with_store_params(store_params.clone());
}
if let Some(cache) = &self.sstable_cache {
planner = planner.with_sstable_cache(cache.clone());
}
if let Some(warmer) = &self.warmer {
planner = planner.with_warmer(warmer.clone());
}
planner
.plan_scan(
self.projection.as_deref(),
self.filter.as_ref(),
self.limit,
self.offset,
self.with_memtable_gen,
self.with_row_address,
)
.await
}
pub fn full_text_search(mut self, query: FullTextSearchQuery) -> Result<Self> {
self.full_text_query = Some(query);
Ok(self)
}
pub async fn try_into_stream(&self) -> Result<SendableRecordBatchStream> {
let plan = self.create_plan().await?;
let ctx = SessionContext::new();
let task_ctx = ctx.task_ctx();
plan.execute(0, task_ctx)
.map_err(|e| Error::io(format!("Failed to execute plan: {}", e)))
}
pub async fn try_into_batch(&self) -> Result<RecordBatch> {
let stream = self.try_into_stream().await?;
let output_schema = stream.schema();
let batches: Vec<RecordBatch> = stream
.try_collect()
.await
.map_err(|e| Error::io(format!("Failed to collect batches: {}", e)))?;
if batches.is_empty() {
return Ok(RecordBatch::new_empty(output_schema));
}
let schema = batches[0].schema();
arrow_select::concat::concat_batches(&schema, &batches)
.map_err(|e| Error::io(format!("Failed to concatenate batches: {}", e)))
}
pub async fn count_rows(&self) -> Result<u64> {
let stream = self.try_into_stream().await?;
let batches: Vec<RecordBatch> = stream
.try_collect()
.await
.map_err(|e| Error::io(format!("Failed to count rows: {}", e)))?;
Ok(batches.iter().map(|b| b.num_rows() as u64).sum())
}
pub async fn contains_pks(&self, pks: &RecordBatch) -> Result<Vec<bool>> {
self.contains_pks_at(pks, None).await
}
pub async fn contains_pks_at(
&self,
pks: &RecordBatch,
watermarks: Option<&HashMap<Uuid, FreshTierWatermark>>,
) -> Result<Vec<bool>> {
let sources = self.build_collector().collect()?;
let memberships = super::block_list::fresh_tier_block_list(
&sources,
self.session.as_ref(),
self.store_params.as_ref(),
self.sstable_cache.as_ref(),
watermarks,
)
.await?;
let pk_indices = super::exec::resolve_pk_indices(pks, &self.pk_columns)
.map_err(|e| Error::invalid_input(e.to_string()))?;
let keys: Vec<ScalarValue> = (0..pks.num_rows())
.map(|row| {
let values: Vec<ScalarValue> = pk_indices
.iter()
.map(|&col| ScalarValue::try_from_array(pks.column(col), row))
.collect::<std::result::Result<_, _>>()
.map_err(|e| Error::invalid_input(e.to_string()))?;
super::block_list::on_disk_pk_key(&values)
})
.collect::<Result<_>>()?;
let mut contained = vec![false; keys.len()];
let mut live: Vec<usize> = (0..keys.len()).collect();
for membership in &memberships {
if live.is_empty() {
break;
}
let live_keys: Vec<ScalarValue> = live.iter().map(|&i| keys[i].clone()).collect();
let mask = membership.contains_keys(&live_keys).await?;
let mut next_live = Vec::with_capacity(live.len());
for (pos, &row) in live.iter().enumerate() {
if mask[pos] {
contained[row] = true;
} else {
next_live.push(row);
}
}
live = next_live;
}
Ok(contained)
}
fn build_collector(&self) -> LsmDataSourceCollector {
let mut collector = match &self.base {
BaseSource::Table(dataset) => {
LsmDataSourceCollector::new(dataset.clone(), self.shard_snapshots.clone())
}
BaseSource::PathOnly(path) => LsmDataSourceCollector::without_base_table(
path.clone(),
self.shard_snapshots.clone(),
),
};
for (shard_id, mems) in &self.in_memory_memtables {
collector = collector.with_in_memory_memtables(*shard_id, mems.clone());
}
collector
}
}
impl std::fmt::Debug for LsmScanner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let (label, value) = match &self.base {
BaseSource::Table(dataset) => ("base_table", dataset.uri().to_string()),
BaseSource::PathOnly(path) => ("base_path", path.clone()),
};
f.debug_struct("LsmScanner")
.field(label, &value)
.field("num_shards", &self.shard_snapshots.len())
.field(
"num_in_memory_memtables",
&self
.in_memory_memtables
.values()
.map(|m| 1 + m.frozen.len())
.sum::<usize>(),
)
.field("projection", &self.projection)
.field("limit", &self.limit)
.field("offset", &self.offset)
.field("pk_columns", &self.pk_columns)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use arrow_array::{Int32Array, ListArray, StringArray, StructArray, UInt32Array};
use arrow_buffer::{OffsetBuffer, ScalarBuffer};
use arrow_schema::{Field, Fields};
use lance_index::scalar::inverted::{DOC_INDEX_COL, DocumentGranularity, InvertedIndexParams};
use crate::dataset::mem_wal::write::{BatchStore, IndexStore};
#[test]
fn test_lsm_scanner_builder() {
let pk_columns = ["id".to_string()];
let shard_snapshots: Vec<ShardSnapshot> = vec![];
assert_eq!(pk_columns.len(), 1);
assert!(shard_snapshots.is_empty());
}
#[test]
fn point_lookup_extraction_requires_normalizable_literals() {
use datafusion::prelude::{col, lit};
let filters = [
col("id").in_list(vec![lit(1i32), lit("not an integer")], false),
col("id").in_list(vec![lit(1i32), col("other")], false),
];
for filter in filters {
assert!(
extract_pk_point_keys(&filter, "id", &DataType::Int32).is_none(),
"unsupported point-lookup filter must use the scan path: {filter}"
);
}
}
#[test]
fn test_shard_snapshot_construction() {
use super::super::data_source::ShardSnapshot;
let shard_id = Uuid::new_v4();
let snapshot = ShardSnapshot::new(shard_id)
.with_spec_id(1)
.with_current_generation(5)
.with_sstable(1, "path/gen_1".to_string())
.with_sstable(2, "path/gen_2".to_string());
assert_eq!(snapshot.shard_id, shard_id);
assert_eq!(snapshot.spec_id, 1);
assert_eq!(snapshot.current_generation, 5);
assert_eq!(snapshot.sstables.len(), 2);
}
#[test]
fn test_in_memory_memtable_ref() {
use crate::dataset::mem_wal::write::{BatchStore, IndexStore};
let batch_store = Arc::new(BatchStore::with_capacity(100));
let index_store = Arc::new(IndexStore::new());
let schema = Arc::new(arrow_schema::Schema::empty());
let memtable_ref = InMemoryMemTableRef {
batch_store,
index_store,
schema,
generation: 10,
};
assert_eq!(memtable_ref.generation, 10);
}
fn pk_schema() -> SchemaRef {
use arrow_schema::{DataType, Field, Schema};
Arc::new(Schema::new(vec![Field::new("id", DataType::Int32, false)]))
}
fn id_pk_batch(ids: &[i32]) -> RecordBatch {
use arrow_array::Int32Array;
RecordBatch::try_new(pk_schema(), vec![Arc::new(Int32Array::from(ids.to_vec()))]).unwrap()
}
fn mk_pk_memtable(ids: &[i32], generation: u64) -> InMemoryMemTableRef {
use crate::dataset::mem_wal::write::{BatchStore, IndexStore};
let store = BatchStore::with_capacity(8);
let mut index = IndexStore::new();
index.enable_pk_index(&[("id".to_string(), 0)]);
let b = id_pk_batch(ids);
let (bp, off, _) = store.append(b.clone()).unwrap();
index.insert_with_batch_position(&b, off, Some(bp)).unwrap();
InMemoryMemTableRef {
batch_store: Arc::new(store),
index_store: Arc::new(index),
schema: pk_schema(),
generation,
}
}
#[tokio::test]
async fn overfetch_factor_only_applies_to_searches() {
let scan = LsmScanner::without_base_table(
pk_schema(),
"memory://scan",
vec![],
vec!["id".to_string()],
)
.with_overfetch_factor(0.5)
.try_into_batch()
.await
.unwrap();
assert_eq!(scan.num_rows(), 0);
let vector_schema = pk_schema_with(arrow_schema::Field::new(
"vector",
DataType::FixedSizeList(
Arc::new(arrow_schema::Field::new("item", DataType::Float32, true)),
4,
),
false,
));
let vector_search = LsmScanner::without_base_table(
vector_schema,
"memory://vector",
vec![],
vec!["id".to_string()],
)
.nearest(
"vector",
&arrow_array::Float32Array::from(vec![0.0f32, 0.0, 0.0, 0.0]),
1,
)
.unwrap()
.with_overfetch_factor(0.5);
let Err(err) = vector_search.try_into_stream().await else {
panic!("invalid overfetch factor should fail vector search planning");
};
assert!(
err.to_string().contains("overfetch_factor"),
"unexpected error for invalid overfetch factor: {err}"
);
let fts_search = LsmScanner::without_base_table(
pk_schema_with(arrow_schema::Field::new("text", DataType::Utf8, true)),
"memory://fts",
vec![],
vec!["id".to_string()],
)
.full_text_search(
FullTextSearchQuery::new("lance".to_string())
.with_column("text".to_string())
.unwrap(),
)
.unwrap()
.with_overfetch_factor(0.5);
let Err(err) = fts_search.try_into_stream().await else {
panic!("invalid overfetch factor should fail full-text search planning");
};
assert!(
err.to_string().contains("overfetch_factor"),
"unexpected error for invalid overfetch factor: {err}"
);
}
#[tokio::test]
async fn unknown_scan_projection_column_is_rejected() {
let scanner = LsmScanner::without_base_table(
pk_schema(),
"memory://t",
vec![],
vec!["id".to_string()],
)
.project(&["missing"])
.unwrap();
let Err(err) = scanner.try_into_stream().await else {
panic!("unknown projection column should fail planning");
};
assert!(
err.to_string().contains("missing"),
"unexpected missing-column projection error: {err}"
);
}
#[tokio::test]
async fn point_lookup_fast_route_rejects_missing_projection_column() {
let shard = Uuid::new_v4();
let scanner = LsmScanner::without_base_table(
pk_schema(),
"memory://t",
vec![],
vec!["id".to_string()],
)
.with_in_memory_memtables(
shard,
InMemoryMemTables {
active: mk_pk_memtable(&[1, 2], 2),
frozen: vec![],
},
)
.project(&["missing"])
.unwrap()
.filter("id = 1")
.unwrap();
let Err(err) = scanner.try_into_stream().await else {
panic!("unknown projection column should fail on point-lookup fast route");
};
assert!(
err.to_string().contains("missing"),
"unexpected missing-column point lookup error: {err}"
);
}
#[tokio::test]
async fn full_text_search_missing_column_is_rejected() {
use arrow_schema::{DataType, Field};
let scanner = LsmScanner::without_base_table(
pk_schema_with(Field::new("text", DataType::Utf8, true)),
"memory://t",
vec![],
vec!["id".to_string()],
)
.full_text_search(
FullTextSearchQuery::new("lance".to_string())
.with_column("missing".to_string())
.unwrap(),
)
.unwrap();
let Err(err) = scanner.try_into_stream().await else {
panic!("unknown FTS column should fail planning");
};
assert!(
err.to_string().contains("missing"),
"unexpected missing FTS column error: {err}"
);
}
#[tokio::test]
async fn contains_pks_reports_fresh_tier_membership() {
let shard = Uuid::new_v4();
let scanner = LsmScanner::without_base_table(
pk_schema(),
"memory://t",
vec![],
vec!["id".to_string()],
)
.with_in_memory_memtables(
shard,
InMemoryMemTables {
active: mk_pk_memtable(&[1, 2], 2),
frozen: vec![mk_pk_memtable(&[3], 1)],
},
);
let result = scanner
.contains_pks(&id_pk_batch(&[1, 4, 3]))
.await
.unwrap();
assert_eq!(result, vec![true, false, true]);
}
#[tokio::test]
async fn contains_pks_at_batched_probe_respects_watermark() {
use crate::dataset::mem_wal::scanner::data_source::FreshTierWatermark;
let shard = Uuid::new_v4();
let scanner = LsmScanner::without_base_table(
pk_schema(),
"memory://t",
vec![],
vec!["id".to_string()],
)
.with_in_memory_memtables(
shard,
InMemoryMemTables {
active: mk_pk_memtable(&[1, 2], 2),
frozen: vec![mk_pk_memtable(&[3, 4], 1)],
},
);
let probe = id_pk_batch(&[4, 1, 9, 3, 2, 1]);
let live = scanner.contains_pks_at(&probe, None).await.unwrap();
assert_eq!(live, vec![true, true, false, true, true, true]);
let watermarks: HashMap<Uuid, FreshTierWatermark> = [(
shard,
FreshTierWatermark {
active_generation: 1,
active_batch_count: u64::MAX,
},
)]
.into_iter()
.collect();
let bounded = scanner
.contains_pks_at(&probe, Some(&watermarks))
.await
.unwrap();
assert_eq!(bounded, vec![true, false, false, true, false, false]);
}
fn pk_schema_with(extra: arrow_schema::Field) -> SchemaRef {
use arrow_schema::{DataType, Field, Schema};
let mut id_meta = HashMap::new();
id_meta.insert(
"lance-schema:unenforced-primary-key".to_string(),
"true".to_string(),
);
Arc::new(Schema::new(vec![
Field::new("id", DataType::Int32, false).with_metadata(id_meta),
extra,
]))
}
#[tokio::test]
async fn nearest_dispatches_through_facade() {
use crate::dataset::mem_wal::write::{BatchStore, IndexStore};
use crate::dataset::{Dataset, WriteParams};
use arrow_array::builder::{FixedSizeListBuilder, Float32Builder};
use arrow_array::{Float32Array, Int32Array, RecordBatchIterator};
use arrow_schema::{DataType, Field};
let schema = pk_schema_with(Field::new(
"vector",
DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float32, true)), 4),
false,
));
let make_batch = |ids: &[i32]| -> RecordBatch {
let mut vb = FixedSizeListBuilder::new(Float32Builder::new(), 4);
for id in ids {
let base = *id as f32 * 0.1;
for d in 0..4 {
vb.values().append_value(base + d as f32 * 0.1);
}
vb.append(true);
}
RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(ids.to_vec())),
Arc::new(vb.finish()),
],
)
.unwrap()
};
let tmp = tempfile::tempdir().unwrap();
let uri = format!("{}/base", tmp.path().to_str().unwrap());
let reader = RecordBatchIterator::new(vec![Ok(make_batch(&[100, 200]))], schema.clone());
let base = Arc::new(
Dataset::write(reader, &uri, Some(WriteParams::default()))
.await
.unwrap(),
);
let store = Arc::new(BatchStore::with_capacity(16));
let mut index = IndexStore::new();
index.enable_pk_index(&[("id".to_string(), 0)]);
index.add_hnsw(
"vec_hnsw".to_string(),
1,
"vector".to_string(),
DistanceType::L2,
64,
8,
);
let batch = make_batch(&[1, 2, 3]);
store.append(batch.clone()).unwrap();
index
.insert_with_batch_position(&batch, 0, Some(0))
.unwrap();
let scanner = LsmScanner::new(base, vec![], vec!["id".to_string()])
.with_in_memory_memtables(
Uuid::new_v4(),
InMemoryMemTables {
active: InMemoryMemTableRef {
batch_store: store,
index_store: Arc::new(index),
schema: schema.clone(),
generation: 1,
},
frozen: vec![],
},
)
.nearest(
"vector",
&Float32Array::from(vec![0.1f32, 0.2, 0.3, 0.4]),
3,
)
.unwrap();
let batches: Vec<RecordBatch> = scanner
.try_into_stream()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let ids: Vec<i32> = batches
.iter()
.flat_map(|b| {
b.column_by_name("id")
.unwrap()
.as_any()
.downcast_ref::<Int32Array>()
.unwrap()
.values()
.to_vec()
})
.collect();
assert_eq!(
ids.first().copied(),
Some(1),
"nearest neighbor via the facade should be id=1; got {ids:?}"
);
}
#[tokio::test]
async fn nearest_pagination_skips_offset_and_caps_limit() {
use crate::dataset::{Dataset, WriteParams};
use crate::index::DatasetIndexExt;
use crate::index::vector::VectorIndexParams;
use arrow_array::builder::{FixedSizeListBuilder, Float32Builder};
use arrow_array::{Float32Array, Int32Array, RecordBatchIterator};
use arrow_schema::{DataType, Field};
use lance_index::IndexType;
let schema = pk_schema_with(Field::new(
"vector",
DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float32, true)), 4),
false,
));
let mut vb = FixedSizeListBuilder::new(Float32Builder::new(), 4);
let ids: Vec<i32> = (1..=6).collect();
for id in &ids {
let base = *id as f32 * 0.1;
for d in 0..4 {
vb.values().append_value(base + d as f32 * 0.1);
}
vb.append(true);
}
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(ids.clone())),
Arc::new(vb.finish()),
],
)
.unwrap();
let tmp = tempfile::tempdir().unwrap();
let uri = format!("{}/base", tmp.path().to_str().unwrap());
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema.clone());
let mut base = Dataset::write(reader, &uri, Some(WriteParams::default()))
.await
.unwrap();
let ivf_flat = VectorIndexParams::ivf_flat(1, DistanceType::L2);
base.create_index(&["vector"], IndexType::Vector, None, &ivf_flat, true)
.await
.unwrap();
let base = Arc::new(base);
let scanner = LsmScanner::new(base, vec![], vec!["id".to_string()])
.nearest(
"vector",
&Float32Array::from(vec![0.1f32, 0.2, 0.3, 0.4]),
2,
)
.unwrap()
.limit(Some(2), Some(1))
.unwrap();
let batches: Vec<RecordBatch> = scanner
.try_into_stream()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let mut out: Vec<i32> = batches
.iter()
.flat_map(|b| {
b.column_by_name("id")
.unwrap()
.as_any()
.downcast_ref::<Int32Array>()
.unwrap()
.values()
.to_vec()
})
.collect();
out.sort();
assert_eq!(
out,
vec![2, 3],
"offset=1, limit=2 over k=2 must return the 2nd-3rd nearest (id=2, id=3); got {out:?}"
);
}
#[tokio::test]
async fn full_text_search_pagination_skips_offset_and_caps_limit() {
use crate::dataset::{Dataset, WriteParams};
use crate::index::DatasetIndexExt;
use arrow_array::{Int32Array, RecordBatchIterator, StringArray};
use arrow_schema::{DataType, Field};
use lance_index::IndexType;
use lance_index::scalar::inverted::tokenizer::InvertedIndexParams;
let schema = pk_schema_with(Field::new("text", DataType::Utf8, true));
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(vec![1, 2, 3])),
Arc::new(StringArray::from(vec![
"lance",
"lance filler",
"lance filler filler",
])),
],
)
.unwrap();
let tmp = tempfile::tempdir().unwrap();
let uri = format!("{}/base", tmp.path().to_str().unwrap());
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema.clone());
let mut base = Dataset::write(reader, &uri, Some(WriteParams::default()))
.await
.unwrap();
base.create_index(
&["text"],
IndexType::Inverted,
Some("text_fts".to_string()),
&InvertedIndexParams::default(),
false,
)
.await
.unwrap();
let base = Arc::new(Dataset::open(&uri).await.unwrap());
let query_limited = LsmScanner::new(base.clone(), vec![], vec!["id".to_string()])
.full_text_search(
FullTextSearchQuery::new("lance".to_string())
.with_column("text".to_string())
.unwrap()
.limit(Some(1)),
)
.unwrap();
let batches: Vec<RecordBatch> = query_limited
.try_into_stream()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let out: Vec<i32> = batches
.iter()
.flat_map(|b| {
b.column_by_name("id")
.unwrap()
.as_any()
.downcast_ref::<Int32Array>()
.unwrap()
.values()
.to_vec()
})
.collect();
assert_eq!(
out,
vec![1],
"query-level FTS limit=1 must cap the unpaginated scanner result; got {out:?}"
);
let query_limited_with_offset =
LsmScanner::new(base.clone(), vec![], vec!["id".to_string()])
.full_text_search(
FullTextSearchQuery::new("lance".to_string())
.with_column("text".to_string())
.unwrap()
.limit(Some(2)),
)
.unwrap()
.limit(None, Some(1))
.unwrap();
let batches: Vec<RecordBatch> = query_limited_with_offset
.try_into_stream()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let out: Vec<i32> = batches
.iter()
.flat_map(|b| {
b.column_by_name("id")
.unwrap()
.as_any()
.downcast_ref::<Int32Array>()
.unwrap()
.values()
.to_vec()
})
.collect();
assert_eq!(
out,
vec![2],
"query-level FTS limit=2 plus offset=1 must page within the top 2; got {out:?}"
);
let scanner = LsmScanner::new(base, vec![], vec!["id".to_string()])
.full_text_search(
FullTextSearchQuery::new("lance".to_string())
.with_column("text".to_string())
.unwrap(),
)
.unwrap()
.limit(Some(1), Some(1))
.unwrap();
let batches: Vec<RecordBatch> = scanner
.try_into_stream()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let out: Vec<i32> = batches
.iter()
.flat_map(|b| {
b.column_by_name("id")
.unwrap()
.as_any()
.downcast_ref::<Int32Array>()
.unwrap()
.values()
.to_vec()
})
.collect();
assert_eq!(
out,
vec![2],
"offset=1, limit=1 must return the 2nd-ranked hit (id=2); got {out:?}"
);
}
#[tokio::test]
async fn full_text_search_without_limit_returns_all_matches() {
use crate::dataset::{Dataset, WriteParams};
use crate::index::DatasetIndexExt;
use arrow_array::{Int32Array, RecordBatchIterator, StringArray};
use arrow_schema::{DataType, Field};
use lance_index::IndexType;
use lance_index::scalar::inverted::tokenizer::InvertedIndexParams;
let schema = pk_schema_with(Field::new("text", DataType::Utf8, true));
let ids: Vec<i32> = (0..12).collect();
let texts: Vec<&str> = (0..12).map(|_| "lance").collect();
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(ids)),
Arc::new(StringArray::from(texts)),
],
)
.unwrap();
let tmp = tempfile::tempdir().unwrap();
let uri = format!("{}/base", tmp.path().to_str().unwrap());
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema.clone());
let mut base = Dataset::write(reader, &uri, Some(WriteParams::default()))
.await
.unwrap();
base.create_index(
&["text"],
IndexType::Inverted,
Some("text_fts".to_string()),
&InvertedIndexParams::default(),
false,
)
.await
.unwrap();
let base = Arc::new(Dataset::open(&uri).await.unwrap());
let scanner = LsmScanner::new(base, vec![], vec!["id".to_string()])
.full_text_search(
FullTextSearchQuery::new("lance".to_string())
.with_column("text".to_string())
.unwrap(),
)
.unwrap();
let batches: Vec<RecordBatch> = scanner
.try_into_stream()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let total: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(
total, 12,
"unbounded LSM FTS must not apply the old default top-10 cap"
);
}
#[tokio::test]
async fn combined_vector_and_fts_is_rejected() {
use crate::dataset::{Dataset, WriteParams};
use arrow_array::builder::{FixedSizeListBuilder, Float32Builder};
use arrow_array::{Float32Array, Int32Array, RecordBatchIterator};
use arrow_schema::{DataType, Field};
let schema = pk_schema_with(Field::new(
"vector",
DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float32, true)), 4),
false,
));
let mut vb = FixedSizeListBuilder::new(Float32Builder::new(), 4);
for d in 0..4 {
vb.values().append_value(d as f32);
}
vb.append(true);
let batch = RecordBatch::try_new(
schema.clone(),
vec![Arc::new(Int32Array::from(vec![1])), Arc::new(vb.finish())],
)
.unwrap();
let tmp = tempfile::tempdir().unwrap();
let uri = format!("{}/base", tmp.path().to_str().unwrap());
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema.clone());
let base = Arc::new(
Dataset::write(reader, &uri, Some(WriteParams::default()))
.await
.unwrap(),
);
let scanner = LsmScanner::new(base, vec![], vec!["id".to_string()])
.nearest(
"vector",
&Float32Array::from(vec![0.0f32, 1.0, 2.0, 3.0]),
1,
)
.unwrap()
.full_text_search(
FullTextSearchQuery::new("lance".to_string())
.with_column("vector".to_string())
.unwrap(),
)
.unwrap();
let err = scanner.create_plan().await.unwrap_err();
assert!(
err.to_string()
.contains("combined vector and full-text search"),
"expected combined-search rejection, got: {err}"
);
}
#[tokio::test]
async fn multi_row_query_vector_is_rejected() {
use crate::dataset::{Dataset, WriteParams};
use arrow_array::builder::{FixedSizeListBuilder, Float32Builder};
use arrow_array::{Int32Array, RecordBatchIterator};
use arrow_schema::{DataType, Field};
let schema = pk_schema_with(Field::new(
"vector",
DataType::FixedSizeList(Arc::new(Field::new("item", DataType::Float32, true)), 4),
false,
));
let mut col = FixedSizeListBuilder::new(Float32Builder::new(), 4);
for d in 0..4 {
col.values().append_value(d as f32);
}
col.append(true);
let batch = RecordBatch::try_new(
schema.clone(),
vec![Arc::new(Int32Array::from(vec![1])), Arc::new(col.finish())],
)
.unwrap();
let tmp = tempfile::tempdir().unwrap();
let uri = format!("{}/base", tmp.path().to_str().unwrap());
let reader = RecordBatchIterator::new(vec![Ok(batch)], schema.clone());
let base = Arc::new(
Dataset::write(reader, &uri, Some(WriteParams::default()))
.await
.unwrap(),
);
let mut q = FixedSizeListBuilder::new(Float32Builder::new(), 4);
for _ in 0..2 {
for d in 0..4 {
q.values().append_value(d as f32);
}
q.append(true);
}
let query = q.finish();
let scanner = LsmScanner::new(base, vec![], vec!["id".to_string()])
.nearest("vector", &query, 1)
.unwrap();
let err = scanner.create_plan().await.unwrap_err();
assert!(
err.to_string().contains("single query vector"),
"expected single-query-vector rejection, got: {err}"
);
}
#[tokio::test]
async fn full_text_search_dispatches_through_facade() {
use crate::dataset::mem_wal::write::{BatchStore, IndexStore};
use crate::dataset::{Dataset, WriteParams};
use crate::index::DatasetIndexExt;
use arrow_array::{Int32Array, RecordBatchIterator, StringArray};
use arrow_schema::{DataType, Field};
use lance_index::IndexType;
use lance_index::scalar::inverted::tokenizer::InvertedIndexParams;
let schema = pk_schema_with(Field::new("text", DataType::Utf8, true));
let make_batch = |ids: &[i32], texts: &[&str]| {
RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(ids.to_vec())),
Arc::new(StringArray::from(texts.to_vec())),
],
)
.unwrap()
};
let tmp = tempfile::tempdir().unwrap();
let uri = format!("{}/base", tmp.path().to_str().unwrap());
let reader =
RecordBatchIterator::new(vec![Ok(make_batch(&[1], &["alpha"]))], schema.clone());
let mut base = Dataset::write(reader, &uri, Some(WriteParams::default()))
.await
.unwrap();
base.create_index(
&["text"],
IndexType::Inverted,
Some("text_fts".to_string()),
&InvertedIndexParams::default(),
false,
)
.await
.unwrap();
let base = Arc::new(Dataset::open(&uri).await.unwrap());
let store = Arc::new(BatchStore::with_capacity(16));
let mut index = IndexStore::new();
index.enable_pk_index(&[("id".to_string(), 0)]);
index.add_fts("text_fts".to_string(), 1, "text".to_string());
let batch = make_batch(&[99], &["zebra"]);
store.append(batch.clone()).unwrap();
index
.insert_with_batch_position(&batch, 0, Some(0))
.unwrap();
let scanner = LsmScanner::new(base, vec![], vec!["id".to_string()])
.with_in_memory_memtables(
Uuid::new_v4(),
InMemoryMemTables {
active: InMemoryMemTableRef {
batch_store: store,
index_store: Arc::new(index),
schema: schema.clone(),
generation: 1,
},
frozen: vec![],
},
)
.full_text_search(
FullTextSearchQuery::new("zebra".to_string())
.with_column("text".to_string())
.unwrap(),
)
.unwrap();
let batches: Vec<RecordBatch> = scanner
.try_into_stream()
.await
.unwrap()
.try_collect()
.await
.unwrap();
let rows: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(
rows, 1,
"facade FTS should surface the memtable 'zebra' row"
);
}
#[tokio::test]
async fn full_text_search_supports_nested_list_element_path() {
let doc_fields = Fields::from(vec![Field::new("content", DataType::Utf8, true)]);
let doc_values = StructArray::new(
doc_fields.clone(),
vec![Arc::new(StringArray::from(vec!["alpha", "beta", "alpha"]))],
None,
);
let doc_item = Arc::new(Field::new("item", DataType::Struct(doc_fields), true));
let docs = ListArray::new(
doc_item.clone(),
OffsetBuffer::new(ScalarBuffer::from(vec![0_i32, 1, 3])),
Arc::new(doc_values),
None,
);
let group_fields = Fields::from(vec![Field::new("docs", DataType::List(doc_item), true)]);
let group_values = StructArray::new(group_fields.clone(), vec![Arc::new(docs)], None);
let group_item = Arc::new(Field::new("item", DataType::Struct(group_fields), true));
let groups = ListArray::new(
group_item,
OffsetBuffer::new(ScalarBuffer::from(vec![0_i32, 2])),
Arc::new(group_values),
None,
);
let schema = pk_schema_with(Field::new("groups", groups.data_type().clone(), true));
let batch = RecordBatch::try_new(
schema.clone(),
vec![Arc::new(Int32Array::from(vec![7])), Arc::new(groups)],
)
.unwrap();
let batch_store = Arc::new(BatchStore::with_capacity(16));
let mut indexes = IndexStore::new();
indexes.enable_pk_index(&[("id".to_string(), 0)]);
indexes
.add_fts_with_params(
"content_list_element_fts".to_string(),
1,
"groups.docs.content".to_string(),
InvertedIndexParams::default()
.document_granularity(DocumentGranularity::ListElement),
)
.unwrap();
let (batch_position, row_offset, _) = batch_store.append(batch.clone()).unwrap();
indexes
.insert_with_batch_position(&batch, row_offset, Some(batch_position))
.unwrap();
let scanner = LsmScanner::without_base_table(
schema.clone(),
"memory://nested_fts",
vec![],
vec!["id".to_string()],
)
.with_in_memory_memtables(
Uuid::new_v4(),
InMemoryMemTables {
active: InMemoryMemTableRef {
batch_store,
index_store: Arc::new(indexes),
schema,
generation: 1,
},
frozen: vec![],
},
)
.project(&["id"])
.unwrap()
.full_text_search(
FullTextSearchQuery::new("alpha".to_string())
.with_column("groups.docs.content".to_string())
.unwrap(),
)
.unwrap();
let batch = scanner.try_into_batch().await.unwrap();
let ids = batch["id"].as_any().downcast_ref::<Int32Array>().unwrap();
let coordinates = batch[DOC_INDEX_COL]
.as_any()
.downcast_ref::<ListArray>()
.unwrap();
let mut hits = (0..batch.num_rows())
.map(|row| {
let coordinate = coordinates
.value(row)
.as_any()
.downcast_ref::<UInt32Array>()
.unwrap()
.values()
.to_vec();
(ids.value(row), coordinate)
})
.collect::<Vec<_>>();
hits.sort_unstable();
assert_eq!(hits, vec![(7, vec![0, 0]), (7, vec![1, 1])]);
}
#[tokio::test]
async fn try_into_batch_empty_fts_keeps_score_schema() {
use arrow_schema::{DataType, Field};
let schema = pk_schema_with(Field::new("text", DataType::Utf8, true));
let scanner = LsmScanner::without_base_table(
schema,
"memory://empty",
vec![],
vec!["id".to_string()],
)
.full_text_search(
FullTextSearchQuery::new("missing".to_string())
.with_column("text".to_string())
.unwrap(),
)
.unwrap();
let batch = scanner.try_into_batch().await.unwrap();
assert_eq!(batch.num_rows(), 0);
assert!(
batch.schema().field_with_name("_score").is_ok(),
"empty FTS batch must keep the planned _score column"
);
}
fn mk_indexed_memtable(schema: &SchemaRef, ids: &[i32], names: &[&str]) -> InMemoryMemTableRef {
use crate::dataset::mem_wal::write::{BatchStore, IndexStore};
use arrow_array::{Int32Array, StringArray};
let store = BatchStore::with_capacity(8);
let mut index = IndexStore::new();
index.add_btree("id_idx".to_string(), 0, "id".to_string());
let batch = RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int32Array::from(ids.to_vec())),
Arc::new(StringArray::from(names.to_vec())),
],
)
.unwrap();
let (idx, row_offset, _) = store.append(batch.clone()).unwrap();
index
.insert_with_batch_position(&batch, row_offset, Some(idx))
.unwrap();
InMemoryMemTableRef {
batch_store: Arc::new(store),
index_store: Arc::new(index),
schema: schema.clone(),
generation: 1,
}
}
#[tokio::test]
async fn point_lookup_filter_routes_to_fast_path() {
use arrow_schema::{DataType, Field, Schema};
use datafusion::physical_plan::displayable;
use datafusion::prelude::{SessionContext, col, lit};
let schema = Arc::new(Schema::new(vec![
Field::new("id", DataType::Int32, false),
Field::new("name", DataType::Utf8, true),
]));
let memtable = mk_indexed_memtable(&schema, &[1, 2, 3, 4, 5], &["a", "b", "c", "d", "e"]);
let shard = Uuid::new_v4();
let scanner = || {
LsmScanner::without_base_table(
schema.clone(),
"memory://t",
vec![],
vec!["id".to_string()],
)
.with_in_memory_memtables(
shard,
InMemoryMemTables {
active: memtable.clone(),
frozen: vec![],
},
)
};
let ctx = SessionContext::new();
let count = |plan: Arc<dyn ExecutionPlan>| {
let ctx = ctx.clone();
async move {
let rows: Vec<RecordBatch> = plan
.execute(0, ctx.task_ctx())
.unwrap()
.try_collect()
.await
.unwrap();
rows.iter().map(|b| b.num_rows()).sum::<usize>()
}
};
let collect_ids = |plan: Arc<dyn ExecutionPlan>| {
let ctx = ctx.clone();
async move {
let rows: Vec<RecordBatch> = plan
.execute(0, ctx.task_ctx())
.unwrap()
.try_collect()
.await
.unwrap();
let mut ids = Vec::new();
for batch in rows {
let id_array = batch
.column_by_name("id")
.unwrap()
.as_any()
.downcast_ref::<arrow_array::Int32Array>()
.unwrap();
ids.extend(id_array.values().iter().copied());
}
ids.sort_unstable();
ids
}
};
let plan = scanner()
.filter_expr(col("id").eq(lit(2i32)))
.create_plan()
.await
.unwrap();
let disp = format!("{}", displayable(plan.as_ref()).indent(true));
assert!(disp.contains("OneShotStream"), "pk=lit must route: {disp}");
assert!(
!disp.contains("Union"),
"must not use the union path: {disp}"
);
assert_eq!(count(plan).await, 1);
let plan = scanner()
.filter_expr(col("id").in_list(vec![lit(1i32), lit(3i32)], false))
.create_plan()
.await
.unwrap();
assert!(
format!("{}", displayable(plan.as_ref()).indent(true)).contains("OneShotStream"),
"pk IN (..) must route"
);
assert_eq!(count(plan).await, 2);
let plan = scanner()
.filter_expr(col("id").in_list(vec![lit(1i32), lit(1i32), lit(2i32)], false))
.limit(Some(2), None)
.unwrap()
.create_plan()
.await
.unwrap();
assert_eq!(collect_ids(plan).await, vec![1, 2]);
let plan = scanner()
.filter_expr(col("id").in_list(vec![lit(1i32), lit(1i64), lit(2i32)], false))
.limit(Some(2), None)
.unwrap()
.create_plan()
.await
.unwrap();
assert_eq!(collect_ids(plan).await, vec![1, 2]);
let plan = scanner()
.filter_expr(col("id").in_list(vec![lit(1i64), lit(1i64), lit(3i64)], false))
.create_plan()
.await
.unwrap();
assert_eq!(collect_ids(plan).await, vec![1, 3]);
let null = lit(ScalarValue::Int32(None));
let plan = scanner()
.filter_expr(col("id").in_list(vec![null.clone(), lit(2i32), null], false))
.create_plan()
.await
.unwrap();
assert_eq!(collect_ids(plan).await, vec![2]);
let plan = scanner()
.filter_expr(col("id").gt(lit(2i32)))
.create_plan()
.await
.unwrap();
assert!(
!format!("{}", displayable(plan.as_ref()).indent(true)).contains("OneShotStream"),
"range filter must not route to the point-lookup node"
);
assert_eq!(count(plan).await, 3);
let plan = scanner()
.filter_expr(col("id").in_list(vec![lit(1i32), lit(3i32), lit(5i32)], false))
.limit(None, Some(1))
.unwrap()
.create_plan()
.await
.unwrap();
let disp = format!("{}", displayable(plan.as_ref()).indent(true));
assert!(
!disp.contains("OneShotStream"),
"offset point-lookup filters must use the scan path: {disp}"
);
assert_eq!(collect_ids(plan).await, vec![3, 5]);
}
}