From 4bbb368eeeb1b0f49f0be9a422369c7fae15ebcb Mon Sep 17 00:00:00 2001 From: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Mon, 24 Aug 2026 19:37:11 +0000 Subject: [PATCH] fix: preserve streaming take conversions --- python/python/lancedb/_lancedb.pyi | 1 + python/python/lancedb/query.py | 6 + python/python/lancedb/table.py | 20 +- python/python/tests/test_query.py | 7 + python/src/query.rs | 3 + rust/lancedb/src/query.rs | 304 +++++++++++++++++++++++------ rust/lancedb/src/remote/table.rs | 101 +++++++++- rust/lancedb/src/table/query.rs | 9 +- 8 files changed, 382 insertions(+), 69 deletions(-) diff --git a/python/python/lancedb/_lancedb.pyi b/python/python/lancedb/_lancedb.pyi index 1e314ede8..c9571f03a 100644 --- a/python/python/lancedb/_lancedb.pyi +++ b/python/python/lancedb/_lancedb.pyi @@ -604,6 +604,7 @@ class FullTextQuery: class PyQueryRequest: limit: Optional[int] offset: Optional[int] + take_offsets: Optional[List[int]] filter: Optional[Union[str, bytes]] full_text_search: Optional[FullTextQuery] select: Optional[Union[str, List[str]]] diff --git a/python/python/lancedb/query.py b/python/python/lancedb/query.py index bd5e2cea2..51b4dad4e 100644 --- a/python/python/lancedb/query.py +++ b/python/python/lancedb/query.py @@ -109,6 +109,7 @@ def _query_is_plain_scan(query: Query) -> bool: return ( query.vector is None and query.full_text_query is None + and query.take_offsets is None and not query.postfilter and not query.order_by ) @@ -775,6 +776,10 @@ class Query(pydantic.BaseModel): # offset to start fetching results from offset: Optional[int] = None + # Dataset offsets whose duplicate occurrences must be restored after lookup. + # This is populated when a take query is converted to this serializable form. + take_offsets: Optional[List[int]] = None + # if true, will only search the indexed data fast_search: Optional[bool] = None @@ -796,6 +801,7 @@ class Query(pydantic.BaseModel): query = cls() query.limit = req.limit query.offset = req.offset + query.take_offsets = req.take_offsets query.filter = req.filter query.full_text_query = req.full_text_search query.columns = req.select diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index 3ae4300cb..fd0c721a5 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -3873,6 +3873,7 @@ class LanceTable(Table): ) and not self._route_pushdown_to_rust and self.current_branch() is None + and query.take_offsets is None ): from lancedb.namespace import _execute_server_side_query @@ -5759,7 +5760,23 @@ class AsyncTable: def _sync_query_to_async( self, query: Query - ) -> AsyncHybridQuery | AsyncFTSQuery | AsyncVectorQuery | AsyncQuery: + ) -> ( + AsyncHybridQuery + | AsyncFTSQuery + | AsyncVectorQuery + | AsyncQuery + | AsyncTakeQuery + ): + if query.take_offsets is not None: + take_query = self.take_offsets(query.take_offsets) + if query.columns: + take_query = take_query.select(query.columns) + if query.use_lsm is not None: + take_query = take_query.use_lsm(query.use_lsm) + if query.with_row_id: + take_query = take_query.with_row_id() + return take_query + async_query = self.query() if query.limit is not None: async_query = async_query.limit(query.limit) @@ -5824,6 +5841,7 @@ class AsyncTable: self._namespace_client, self._pushdown_operations ) and not self._route_pushdown_to_rust + and query.take_offsets is None ): from lancedb.namespace import _execute_server_side_query diff --git a/python/python/tests/test_query.py b/python/python/tests/test_query.py index a475e6956..701b75308 100644 --- a/python/python/tests/test_query.py +++ b/python/python/tests/test_query.py @@ -1899,6 +1899,13 @@ def test_take_queries(tmp_path): 17, ] + # Converting a take builder to its serializable query representation must + # retain occurrence metadata and execute with the same multiplicity. + query = table.take_offsets([5, 2, 5, 17]).select(["idx"]).to_query_object() + assert query.take_offsets == [5, 2, 5, 17] + converted = table._execute_query(query).read_all() + assert sorted(converted["idx"].to_pylist()) == [2, 5, 5, 17] + # Take by row id assert list( sorted(table.take_row_ids([5, 2, 17]).to_pandas()["idx"].to_list()) diff --git a/python/src/query.rs b/python/src/query.rs index affbdf4fd..1020c822d 100644 --- a/python/src/query.rs +++ b/python/src/query.rs @@ -289,6 +289,7 @@ impl<'py> IntoPyObject<'py> for PyQueryVectors { pub struct PyQueryRequest { pub limit: Option, pub offset: Option, + pub take_offsets: Option>, pub filter: Option, pub full_text_search: Option>, pub select: PySelect, @@ -318,6 +319,7 @@ impl From for PyQueryRequest { AnyQuery::Query(query_request) => Self { limit: query_request.limit, offset: query_request.offset, + take_offsets: query_request.take_offsets, filter: query_request.filter.map(PyQueryFilter), full_text_search: query_request .full_text_search @@ -345,6 +347,7 @@ impl From for PyQueryRequest { AnyQuery::VectorQuery(vector_query) => Self { limit: vector_query.base.limit, offset: vector_query.base.offset, + take_offsets: vector_query.base.take_offsets, filter: vector_query.base.filter.map(PyQueryFilter), full_text_search: None, select: PySelect(vector_query.base.select), diff --git a/rust/lancedb/src/query.rs b/rust/lancedb/src/query.rs index 9afb35356..2d1159eff 100644 --- a/rust/lancedb/src/query.rs +++ b/rust/lancedb/src/query.rs @@ -2,6 +2,7 @@ // SPDX-FileCopyrightText: Copyright The LanceDB Authors use std::collections::{HashMap, HashSet}; +use std::pin::Pin; use std::sync::Arc; use std::{future::Future, time::Duration}; @@ -18,13 +19,13 @@ use datafusion_execution::TaskContext; use datafusion_expr::{Expr, col, lit}; use datafusion_physical_expr::{EquivalenceProperties, Partitioning}; use datafusion_physical_plan::{ - DisplayAs, DisplayFormatType, ExecutionPlan, PlanProperties, + DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, PlanProperties, coalesce_partitions::CoalescePartitionsExec, execution_plan::{Boundedness, EmissionType}, limit::GlobalLimitExec, stream::RecordBatchStreamAdapter, }; -use futures::{FutureExt, TryFutureExt, TryStreamExt, stream, try_join}; +use futures::{FutureExt, StreamExt, TryFutureExt, TryStreamExt, stream, try_join}; use half::f16; /// Re-export Lance ColumnOrdering type for use in query ordering pub use lance::dataset::scanner::ColumnOrdering; @@ -837,6 +838,14 @@ pub struct QueryRequest { /// Offset of the query. pub offset: Option, + /// Dataset offsets whose occurrence multiplicity must be restored after + /// executing the physical lookup represented by this request. + /// + /// This is client-side execution metadata used when a [`TakeQuery`] is + /// converted into a request. It is not sent to remote services. + #[doc(hidden)] + pub take_offsets: Option>, + /// Apply filter to the returned rows. pub filter: Option, @@ -905,6 +914,7 @@ impl Default for QueryRequest { Self { limit: None, offset: None, + take_offsets: None, filter: None, filter_error: None, full_text_search: None, @@ -1530,9 +1540,18 @@ impl HasQuery for VectorQuery { } } -fn restore_take_batch( +fn take_occurrences(offsets: &[u64]) -> HashMap { + let mut occurrences = HashMap::with_capacity(offsets.len()); + for offset in offsets { + *occurrences.entry(*offset).or_insert(0) += 1; + } + occurrences +} + +fn restore_take_batch_with_occurrences( batch: RecordBatch, offsets: &[u64], + occurrences: &HashMap, ordering_column: &str, drop_ordering_column: bool, preserve_order: bool, @@ -1585,15 +1604,11 @@ fn restore_take_batch( .filter_map(|offset| ordering.get(offset).copied()), ); } else { - let mut occurrences = HashMap::with_capacity(offsets.len()); - for offset in offsets { - *occurrences.entry(*offset).or_insert(0) += 1; - } // Public take queries do not guarantee output order. Preserve the lookup's // existing order and only restore the multiplicity of each matching row. for (index, offset) in actual_offsets.iter().enumerate() { - if let Some(count) = occurrences.remove(offset) { - desired_order.extend(std::iter::repeat_n(index as u64, count)); + if let Some(count) = occurrences.get(offset) { + desired_order.extend(std::iter::repeat_n(index as u64, *count)); } } } @@ -1616,16 +1631,36 @@ fn restore_take_batch( Ok(ordered_batch) } +#[cfg(test)] +fn restore_take_batch( + batch: RecordBatch, + offsets: &[u64], + ordering_column: &str, + drop_ordering_column: bool, + preserve_order: bool, +) -> Result { + restore_take_batch_with_occurrences( + batch, + offsets, + &take_occurrences(offsets), + ordering_column, + drop_ordering_column, + preserve_order, + ) +} + /// Restores the logical offset occurrence sequence above the physical lookup plan. /// -/// The lookup plan returns each matching row at most once. This operator collects -/// those rows, expands duplicates, and emits one partition. It preserves lookup order -/// unless the caller explicitly requests offset order. Pagination must remain above -/// this operator so it applies to occurrences. +/// The lookup plan returns each matching row at most once. For ordinary unordered +/// takes this operator expands each input batch incrementally and preserves the +/// lookup's partitioning. The explicitly ordered reader path collects one coalesced +/// input before restoring requested order. Pagination must remain above this operator +/// so it applies to occurrences. #[derive(Debug)] struct TakeRestoreExec { input: Arc, offsets: Vec, + occurrences: Arc>, ordering_column: String, drop_ordering_column: bool, preserve_order: bool, @@ -1648,15 +1683,26 @@ impl TakeRestoreExec { } else { input.schema() }; + let partition_count = if preserve_order { + 1 + } else { + input.output_partitioning().partition_count() + }; + let emission_type = if preserve_order { + EmissionType::Final + } else { + EmissionType::Incremental + }; let properties = Arc::new(PlanProperties::new( EquivalenceProperties::new(schema.clone()), - Partitioning::UnknownPartitioning(1), - EmissionType::Final, + Partitioning::UnknownPartitioning(partition_count), + emission_type, Boundedness::Bounded, )); Ok(Self { input, + occurrences: Arc::new(take_occurrences(&offsets)), offsets, ordering_column, drop_ordering_column, @@ -1695,7 +1741,7 @@ impl ExecutionPlan for TakeRestoreExec { } fn maintains_input_order(&self) -> Vec { - vec![false] + vec![!self.preserve_order] } fn benefits_from_input_partitioning(&self) -> Vec { @@ -1729,35 +1775,55 @@ impl ExecutionPlan for TakeRestoreExec { partition: usize, context: Arc, ) -> DataFusionResult { - if partition != 0 { + let partition_count = self.input.output_partitioning().partition_count(); + if partition >= partition_count || (self.preserve_order && partition != 0) { return Err(DataFusionError::Internal(format!( - "TakeRestoreExec only supports partition 0, got {partition}" + "TakeRestoreExec cannot execute partition {partition}; input has {partition_count} partitions" ))); } - let input = self.input.execute(0, context)?; - let input_schema = input.schema(); + let input = self.input.execute(partition, context)?; let output_schema = self.schema.clone(); let offsets = self.offsets.clone(); + let occurrences = self.occurrences.clone(); let ordering_column = self.ordering_column.clone(); let drop_ordering_column = self.drop_ordering_column; let preserve_order = self.preserve_order; - let stream = stream::once(async move { - let batches = input.try_collect::>().await?; - let batch = if batches.is_empty() { - RecordBatch::new_empty(input_schema.clone()) + let stream: Pin> + Send>> = + if preserve_order { + let input_schema = input.schema(); + Box::pin(stream::once(async move { + let batches = input.try_collect::>().await?; + let batch = if batches.is_empty() { + RecordBatch::new_empty(input_schema.clone()) + } else { + concat_batches(&input_schema, &batches)? + }; + restore_take_batch_with_occurrences( + batch, + &offsets, + &occurrences, + &ordering_column, + drop_ordering_column, + true, + ) + .map_err(|error| DataFusionError::External(Box::new(error))) + })) } else { - concat_batches(&input_schema, &batches)? + Box::pin(input.map(move |batch| { + batch.and_then(|batch| { + restore_take_batch_with_occurrences( + batch, + &offsets, + &occurrences, + &ordering_column, + drop_ordering_column, + false, + ) + .map_err(|error| DataFusionError::External(Box::new(error))) + }) + })) }; - restore_take_batch( - batch, - &offsets, - &ordering_column, - drop_ordering_column, - preserve_order, - ) - .map_err(|error| DataFusionError::External(Box::new(error))) - }); Ok(Box::pin(RecordBatchStreamAdapter::new( output_schema, @@ -1808,6 +1874,7 @@ impl TakeQuery { filter: Some(QueryFilter::Datafusion( col("_rowoffset").in_list(in_list, false), )), + take_offsets: Some(offsets.clone()), ..Default::default() }, offsets: Some(offsets), @@ -1840,15 +1907,20 @@ impl TakeQuery { self } - async fn request_with_row_offset(&self) -> Result<(QueryRequest, String, bool)> { + async fn request_with_row_offset( + parent: &dyn BaseTable, + request: &QueryRequest, + ) -> Result<(QueryRequest, String, bool)> { const ROW_OFFSET: &str = "_rowoffset"; const INTERNAL_ROW_OFFSET: &str = "__lancedb_take_row_offset"; - let mut request = self.request.clone(); + let mut request = request.clone(); + // The physical lookup must not recursively restore occurrences. The + // wrapper above this request owns that logical operation. + request.take_offsets = None; let (ordering_column, drop_ordering_column) = match &mut request.select { Select::All => { - let mut columns = self - .parent + let mut columns = parent .schema() .await? .fields() @@ -1889,10 +1961,11 @@ impl TakeQuery { } async fn prepare_offsets_lookup( - &self, + parent: &dyn BaseTable, + request: &QueryRequest, ) -> Result<(QueryRequest, String, bool, usize, Option)> { let (mut request, ordering_column, drop_ordering_column) = - self.request_with_row_offset().await?; + Self::request_with_row_offset(parent, request).await?; // The lookup operates on distinct physical rows. Pagination is a logical // operation over occurrences and must be applied only after restoration. let output_offset = request.offset.take().unwrap_or_default(); @@ -1916,7 +1989,11 @@ impl TakeQuery { output_limit: Option, preserve_order: bool, ) -> Result> { - let lookup = Arc::new(CoalescePartitionsExec::new(lookup)); + let lookup = if preserve_order { + Arc::new(CoalescePartitionsExec::new(lookup)) as Arc + } else { + lookup + }; let restored: Arc = Arc::new(TakeRestoreExec::try_new( lookup, offsets.to_vec(), @@ -1941,6 +2018,7 @@ impl TakeQuery { occurrence_count: usize, output_offset: usize, output_limit: Option, + preserve_order: bool, ) -> String { fn indent(plan: &str, spaces: usize) -> String { let indentation = " ".repeat(spaces); @@ -1950,10 +2028,17 @@ impl TakeQuery { .join("\n") } - let restored = format!( - "TakeRestoreExec: occurrences={occurrence_count}\n CoalescePartitionsExec\n{}", - indent(lookup, 4) - ); + let restored = if preserve_order { + format!( + "TakeRestoreExec: occurrences={occurrence_count}\n CoalescePartitionsExec\n{}", + indent(lookup, 4) + ) + } else { + format!( + "TakeRestoreExec: occurrences={occurrence_count}\n{}", + indent(lookup, 2) + ) + }; if output_offset > 0 || output_limit.is_some() { let fetch = output_limit @@ -1973,24 +2058,14 @@ impl TakeQuery { offsets: &[u64], options: QueryExecutionOptions, ) -> Result> { - let (request, ordering_column, drop_ordering_column, output_offset, output_limit) = - self.prepare_offsets_lookup().await?; - let query = AnyQuery::Query(request); - let lookup = self - .parent - .clone() - .create_plan(&query, options.without_output_batch_length_limit()) - .await?; - - Self::wrap_offsets_plan( - lookup, + create_take_offsets_plan( + self.parent.as_ref(), + &self.request, offsets, - ordering_column, - drop_ordering_column, - output_offset, - output_limit, + options, self.preserve_order, ) + .await } /// Convert the `TakeQuery` into a `QueryRequest`. @@ -2037,6 +2112,63 @@ impl TakeQuery { } } +pub(crate) async fn create_take_offsets_plan( + parent: &dyn BaseTable, + request: &QueryRequest, + offsets: &[u64], + options: QueryExecutionOptions, + preserve_order: bool, +) -> Result> { + let (request, ordering_column, drop_ordering_column, output_offset, output_limit) = + TakeQuery::prepare_offsets_lookup(parent, request).await?; + let lookup_options = if preserve_order { + options.without_output_batch_length_limit() + } else { + options + }; + let lookup = parent + .create_plan(&AnyQuery::Query(request), lookup_options) + .await?; + + TakeQuery::wrap_offsets_plan( + lookup, + offsets, + ordering_column, + drop_ordering_column, + output_offset, + output_limit, + preserve_order, + ) +} + +pub(crate) async fn explain_take_offsets_plan( + parent: &dyn BaseTable, + request: &QueryRequest, + offsets: &[u64], + verbose: bool, +) -> Result { + let (request, _, _, output_offset, output_limit) = + TakeQuery::prepare_offsets_lookup(parent, request).await?; + let lookup = parent + .explain_plan(&AnyQuery::Query(request), verbose) + .await?; + Ok(TakeQuery::wrap_offsets_explanation( + &lookup, + offsets.len(), + output_offset, + output_limit, + false, + )) +} + +pub(crate) async fn prepare_take_offsets_request( + parent: &dyn BaseTable, + request: &QueryRequest, +) -> Result { + let (request, _, _, _, _) = TakeQuery::prepare_offsets_lookup(parent, request).await?; + Ok(request) +} + impl HasQuery for TakeQuery { fn mut_query(&mut self) -> &mut QueryRequest { &mut self.request @@ -2078,7 +2210,7 @@ impl ExecutableQuery for TakeQuery { async fn explain_plan(&self, verbose: bool) -> Result { if let Some(offsets) = &self.offsets { let (request, _, _, output_offset, output_limit) = - self.prepare_offsets_lookup().await?; + Self::prepare_offsets_lookup(self.parent.as_ref(), &self.request).await?; // Ask the backend to explain only the distinct-row lookup. This keeps // remote explanation non-executing while still showing the client-side // operators that create_plan and execution place above that lookup. @@ -2091,6 +2223,7 @@ impl ExecutableQuery for TakeQuery { offsets.len(), output_offset, output_limit, + self.preserve_order, )); } @@ -2101,7 +2234,8 @@ impl ExecutableQuery for TakeQuery { async fn analyze_plan_with_options(&self, options: QueryExecutionOptions) -> Result { if self.offsets.is_some() { if self.parent.analyze_plan_is_remote() { - let (request, _, _, _, _) = self.prepare_offsets_lookup().await?; + let (request, _, _, _, _) = + Self::prepare_offsets_lookup(self.parent.as_ref(), &self.request).await?; // Remote analysis is owned by the service. The current wire // request represents only the distinct-row lookup, so return // the service report unchanged instead of fabricating metrics @@ -2135,6 +2269,7 @@ mod tests { types::Float32Type, }; use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema}; + use datafusion_physical_plan::display::DisplayableExecutionPlan; use futures::{StreamExt, TryStreamExt}; use lance_testing::datagen::{BatchGenerator, IncrementingInt32, RandomVector}; use rand::seq::IndexedRandom; @@ -3175,6 +3310,47 @@ mod tests { assert_eq!(ids, vec![1, 5, 5, 17]); } + #[tokio::test] + async fn test_take_offsets_plan_is_incremental() { + let tmp_dir = tempdir().unwrap(); + let table = make_test_table(&tmp_dir).await; + + let plan = table + .take_offsets(vec![5, 1, 17]) + .create_plan(QueryExecutionOptions { + max_batch_length: 1, + ..Default::default() + }) + .await + .unwrap(); + + assert_eq!(plan.properties().emission_type, EmissionType::Incremental); + let displayed = DisplayableExecutionPlan::new(plan.as_ref()) + .indent(false) + .to_string(); + assert!(displayed.contains("TakeRestoreExec")); + assert!(!displayed.contains("CoalescePartitionsExec")); + } + + #[tokio::test] + async fn test_take_into_request_preserves_duplicate_multiplicity() { + let tmp_dir = tempdir().unwrap(); + let table = make_test_table(&tmp_dir).await; + let request = table.take_offsets(vec![5, 5]).into_request(); + assert_eq!(request.take_offsets, Some(vec![5, 5])); + + let batches = table + .base_table() + .query(&AnyQuery::Query(request), QueryExecutionOptions::default()) + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 2); + } + #[test] fn test_restore_take_batch_only_reorders_when_requested() { let batch = RecordBatch::try_from_iter([ @@ -3303,12 +3479,12 @@ mod tests { let explained = take.explain_plan(false).await.unwrap(); assert!(explained.contains("GlobalLimitExec")); assert!(explained.contains("TakeRestoreExec")); - assert!(explained.contains("CoalescePartitionsExec")); + assert!(!explained.contains("CoalescePartitionsExec")); let analyzed = take.analyze_plan().await.unwrap(); assert!(analyzed.contains("GlobalLimitExec")); assert!(analyzed.contains("TakeRestoreExec")); - assert!(analyzed.contains("CoalescePartitionsExec")); + assert!(!analyzed.contains("CoalescePartitionsExec")); } #[tokio::test] diff --git a/rust/lancedb/src/remote/table.rs b/rust/lancedb/src/remote/table.rs index dc064c48d..ba8f0aaae 100644 --- a/rust/lancedb/src/remote/table.rs +++ b/rust/lancedb/src/remote/table.rs @@ -39,7 +39,8 @@ use crate::table::{ use crate::table::{AnyQuery, Filter, Predicate, PreprocessingOutput, TableStatistics}; use crate::utils::background_cache::BackgroundCache; use crate::utils::{ - resolve_arrow_field_path, supported_btree_data_type, supported_vector_data_type, + MaxBatchLengthStream, TimeoutStream, resolve_arrow_field_path, supported_btree_data_type, + supported_vector_data_type, }; use crate::{DistanceType, Error}; use crate::{ @@ -2209,6 +2210,13 @@ impl BaseTable for RemoteTable { query: &AnyQuery, options: QueryExecutionOptions, ) -> Result> { + if let AnyQuery::Query(request) = query + && let Some(offsets) = &request.take_offsets + { + return crate::query::create_take_offsets_plan(self, request, offsets, options, false) + .await; + } + let streams = self.execute_query(query, &options).await?; if streams.len() == 1 { let stream = streams.into_iter().next().unwrap(); @@ -2227,6 +2235,27 @@ impl BaseTable for RemoteTable { query: &AnyQuery, options: QueryExecutionOptions, ) -> Result { + if let AnyQuery::Query(request) = query + && let Some(offsets) = &request.take_offsets + { + let plan = crate::query::create_take_offsets_plan( + self, + request, + offsets, + options.clone(), + false, + ) + .await?; + let inner = execute_plan(plan, Default::default())?; + let inner = MaxBatchLengthStream::new_boxed(inner, options.max_batch_length as usize); + let inner = if let Some(timeout) = options.timeout { + TimeoutStream::new_boxed(inner, timeout) + } else { + inner + }; + return Ok(DatasetRecordBatchStream::new(inner)); + } + let streams = self.execute_query(query, &options).await?; if streams.len() == 1 { @@ -2264,6 +2293,12 @@ impl BaseTable for RemoteTable { } async fn explain_plan(&self, query: &AnyQuery, verbose: bool) -> Result { + if let AnyQuery::Query(request) = query + && let Some(offsets) = &request.take_offsets + { + return crate::query::explain_take_offsets_plan(self, request, offsets, verbose).await; + } + let base_request = self.post_read(&format!("/v1/table/{}/explain_plan/", self.identifier)); let query_bodies = self.prepare_query_bodies(query).await?; @@ -2311,6 +2346,17 @@ impl BaseTable for RemoteTable { query: &AnyQuery, options: QueryExecutionOptions, ) -> Result { + let prepared_query = if let AnyQuery::Query(request) = query + && request.take_offsets.is_some() + { + Some(AnyQuery::Query( + crate::query::prepare_take_offsets_request(self, request).await?, + )) + } else { + None + }; + let query = prepared_query.as_ref().unwrap_or(query); + let mut request = self.post_read(&format!("/v1/table/{}/analyze_plan/", self.identifier)); if options.analyze_plan_distributed_metrics != AnalyzePlanDistributedMetrics::Aggregate { @@ -3263,7 +3309,7 @@ mod tests { use arrow_array::{BinaryArray, Int32Array, RecordBatch, RecordBatchIterator, record_batch}; use arrow_schema::{DataType, Field, Schema}; use chrono::{DateTime, Utc}; - use futures::{StreamExt, TryFutureExt, future::BoxFuture}; + use futures::{StreamExt, TryFutureExt, TryStreamExt, future::BoxFuture}; use lance_index::scalar::inverted::query::MatchQuery; use lance_index::scalar::{FullTextSearchQuery, InvertedIndexParams}; use reqwest::Body; @@ -5069,10 +5115,59 @@ mod tests { assert!(explained.contains("GlobalLimitExec")); assert!(explained.contains("TakeRestoreExec")); - assert!(explained.contains("CoalescePartitionsExec")); + assert!(!explained.contains("CoalescePartitionsExec")); assert!(explained.contains("RemoteLookupExec")); } + #[tokio::test] + async fn test_converted_take_request_restores_remote_occurrences() { + let table = Table::new_with_handler("my_table", |request| { + assert_eq!(request.method(), "POST"); + assert_eq!(request.url().path(), "/v1/table/my_table/query/"); + + let body: serde_json::Value = + serde_json::from_slice(request.body().unwrap().as_bytes().unwrap()).unwrap(); + assert_eq!(body["columns"], json!(["id", "_rowoffset"])); + + let data = RecordBatch::try_new( + Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int32, false), + Field::new("_rowoffset", DataType::UInt64, false), + ])), + vec![ + Arc::new(Int32Array::from(vec![5])), + Arc::new(arrow_array::UInt64Array::from(vec![5])), + ], + ) + .unwrap(); + http::Response::builder() + .status(200) + .header(CONTENT_TYPE, ARROW_FILE_CONTENT_TYPE) + .body(write_ipc_file(&data)) + .unwrap() + }); + + let request = table + .take_offsets(vec![5, 5]) + .select(crate::query::Select::columns(&["id"])) + .into_request(); + let batches = table + .base_table() + .query(&AnyQuery::Query(request), QueryExecutionOptions::default()) + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + + assert_eq!(batches.iter().map(RecordBatch::num_rows).sum::(), 2); + assert!( + batches + .iter() + .all(|batch| batch.schema().fields().len() == 1) + ); + } + #[tokio::test] async fn test_take_offsets_analyze_plan_delegates_to_remote() { let table = Table::new_with_handler("my_table", |request| { diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 9feb9d5ab..3151f8b56 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -90,7 +90,7 @@ fn requires_local_namespace_execution(query: &AnyQuery) -> bool { // pushing these down would silently ignore the user's setting. For use_lsm that // is worse than a tuning miss: MemWAL read routing lives only in `create_plan`, // so a pushed-down query would return stale base-only data with no error. - if query.base().use_lsm.is_some() { + if query.base().use_lsm.is_some() || query.base().take_offsets.is_some() { return true; } matches!( @@ -133,6 +133,13 @@ pub async fn create_plan( query: &AnyQuery, options: QueryExecutionOptions, ) -> Result> { + if let AnyQuery::Query(request) = query + && let Some(offsets) = &request.take_offsets + { + return crate::query::create_take_offsets_plan(table, request, offsets, options, false) + .await; + } + let query = match query { AnyQuery::VectorQuery(query) => query.clone(), AnyQuery::Query(query) => VectorQueryRequest::from_plain_query(query.clone()),