This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/paimon-rust.git
The following commit(s) were added to refs/heads/main by this push:
new b44ffea feat(datafusion): support lateral vector search joins (#452)
b44ffea is described below
commit b44ffea692e0d36a4847f1da06641c2703d6c3fc
Author: Jingsong Lee <[email protected]>
AuthorDate: Sun Jul 5 08:47:29 2026 +0800
feat(datafusion): support lateral vector search joins (#452)
---
crates/integrations/datafusion/Cargo.toml | 1 +
.../datafusion/src/lateral_vector_search.rs | 698 +++++++++++++++++++++
crates/integrations/datafusion/src/lib.rs | 1 +
crates/integrations/datafusion/src/sql_context.rs | 21 +-
.../integrations/datafusion/src/vector_search.rs | 98 ++-
.../integrations/datafusion/tests/read_tables.rs | 51 ++
6 files changed, 846 insertions(+), 24 deletions(-)
diff --git a/crates/integrations/datafusion/Cargo.toml
b/crates/integrations/datafusion/Cargo.toml
index a0a9bf5..e78d10f 100644
--- a/crates/integrations/datafusion/Cargo.toml
+++ b/crates/integrations/datafusion/Cargo.toml
@@ -33,6 +33,7 @@ mosaic = ["paimon/mosaic"]
vortex = ["paimon/vortex"]
[dependencies]
+arrow-select = { workspace = true }
async-trait = "0.1"
chrono = "0.4"
constant_time_eq = { workspace = true }
diff --git a/crates/integrations/datafusion/src/lateral_vector_search.rs
b/crates/integrations/datafusion/src/lateral_vector_search.rs
new file mode 100644
index 0000000..91a202a
--- /dev/null
+++ b/crates/integrations/datafusion/src/lateral_vector_search.rs
@@ -0,0 +1,698 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements. See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership. The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License. You may obtain a copy of the License at
+//
+// http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied. See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+use std::any::Any;
+use std::cmp::Ordering;
+use std::collections::HashMap;
+use std::fmt;
+use std::hash::{Hash, Hasher};
+use std::sync::Arc;
+
+use async_trait::async_trait;
+use datafusion::arrow::array::{
+ new_empty_array, Array, ArrayRef, FixedSizeListArray, Float32Array,
Int64Array, ListArray,
+ RecordBatch, UInt32Array,
+};
+use datafusion::arrow::datatypes::SchemaRef as ArrowSchemaRef;
+use datafusion::catalog::default_table_source::source_as_provider;
+use datafusion::common::stats::Precision;
+use datafusion::common::tree_node::Transformed;
+use datafusion::common::{
+ internal_err, DFSchemaRef, DataFusionError, JoinType, Result as DFResult,
Statistics,
+};
+use datafusion::datasource::TableProvider;
+use datafusion::execution::context::{QueryPlanner, SessionState};
+use datafusion::execution::{SendableRecordBatchStream, TaskContext};
+use datafusion::logical_expr::{Expr, Extension, LogicalPlan, TableScan,
UserDefinedLogicalNode};
+use datafusion::optimizer::{ApplyOrder, Optimizer, OptimizerConfig,
OptimizerRule};
+use datafusion::physical_expr::PhysicalExpr;
+use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
+use datafusion::physical_plan::stream::RecordBatchStreamAdapter;
+use datafusion::physical_plan::{
+ DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties,
Partitioning,
+ PlanProperties,
+};
+use datafusion::physical_planner::{DefaultPhysicalPlanner, ExtensionPlanner,
PhysicalPlanner};
+use datafusion::prelude::SessionConfig;
+use futures::{StreamExt, TryStreamExt};
+use paimon::spec::ROW_ID_FIELD_NAME;
+use paimon::table::{RowRange, Table};
+use paimon::vector_search::SearchResult;
+
+use crate::error::to_datafusion_error;
+use crate::vector_search::LateralVectorSearchTableProvider;
+
+#[derive(Debug)]
+pub(crate) struct PaimonQueryPlanner;
+
+impl PaimonQueryPlanner {
+ pub(crate) fn new() -> Self {
+ Self
+ }
+}
+
+#[async_trait]
+impl QueryPlanner for PaimonQueryPlanner {
+ async fn create_physical_plan(
+ &self,
+ logical_plan: &LogicalPlan,
+ session_state: &SessionState,
+ ) -> DFResult<Arc<dyn ExecutionPlan>> {
+ let planner =
DefaultPhysicalPlanner::with_extension_planners(vec![Arc::new(
+ LateralVectorSearchExtensionPlanner,
+ )]);
+ planner
+ .create_physical_plan(logical_plan, session_state)
+ .await
+ }
+}
+
+#[derive(Debug)]
+pub(crate) struct RewriteLateralVectorSearch;
+
+impl RewriteLateralVectorSearch {
+ pub(crate) fn new() -> Self {
+ Self
+ }
+}
+
+pub(crate) fn optimizer_rules() -> Vec<Arc<dyn OptimizerRule + Send + Sync>> {
+ let mut rules: Vec<Arc<dyn OptimizerRule + Send + Sync>> =
+ vec![Arc::new(RewriteLateralVectorSearch::new())];
+ rules.extend(Optimizer::default().rules);
+ rules
+}
+
+impl OptimizerRule for RewriteLateralVectorSearch {
+ fn name(&self) -> &str {
+ "rewrite_lateral_vector_search"
+ }
+
+ fn apply_order(&self) -> Option<ApplyOrder> {
+ Some(ApplyOrder::BottomUp)
+ }
+
+ fn rewrite(
+ &self,
+ plan: LogicalPlan,
+ _config: &dyn OptimizerConfig,
+ ) -> DFResult<Transformed<LogicalPlan>> {
+ let LogicalPlan::Join(join) = plan else {
+ return Ok(Transformed::no(plan));
+ };
+
+ if join.join_type != JoinType::Inner || !join.on.is_empty() ||
join.filter.is_some() {
+ return Ok(Transformed::no(LogicalPlan::Join(join)));
+ }
+
+ let Some(spec) = find_lateral_vector_search_provider(&join.right)?
else {
+ return Ok(Transformed::no(LogicalPlan::Join(join)));
+ };
+
+ let node = LateralVectorSearchNode::new(
+ Arc::clone(&join.left),
+ spec.target_table,
+ spec.target_schema,
+ spec.target_column,
+ spec.query_vector_expr,
+ spec.limit,
+ Arc::clone(&join.schema),
+ );
+ Ok(Transformed::yes(LogicalPlan::Extension(Extension {
+ node: Arc::new(node),
+ })))
+ }
+}
+
+fn find_lateral_vector_search_provider(
+ plan: &LogicalPlan,
+) -> DFResult<Option<LateralVectorSearchSpec>> {
+ match plan {
+ LogicalPlan::TableScan(TableScan { source, .. }) => {
+ let provider = source_as_provider(source)?;
+ let Some(provider) = provider
+ .as_any()
+ .downcast_ref::<LateralVectorSearchTableProvider>()
+ else {
+ return Ok(None);
+ };
+ Ok(Some(LateralVectorSearchSpec {
+ target_table: provider.inner().table().clone(),
+ target_schema: provider.inner().schema(),
+ target_column: provider.column_name().to_string(),
+ query_vector_expr: provider.query_vector_expr().clone(),
+ limit: provider.limit(),
+ }))
+ }
+ LogicalPlan::Subquery(subquery) =>
find_lateral_vector_search_provider(&subquery.subquery),
+ LogicalPlan::SubqueryAlias(alias) =>
find_lateral_vector_search_provider(&alias.input),
+ _ => Ok(None),
+ }
+}
+
+struct LateralVectorSearchSpec {
+ target_table: Table,
+ target_schema: ArrowSchemaRef,
+ target_column: String,
+ query_vector_expr: Expr,
+ limit: usize,
+}
+
+#[derive(Debug, Clone)]
+pub(crate) struct LateralVectorSearchNode {
+ input: Arc<LogicalPlan>,
+ target_table: Table,
+ target_schema: ArrowSchemaRef,
+ target_column: String,
+ query_vector_expr: Expr,
+ limit: usize,
+ schema: DFSchemaRef,
+}
+
+impl LateralVectorSearchNode {
+ fn new(
+ input: Arc<LogicalPlan>,
+ target_table: Table,
+ target_schema: ArrowSchemaRef,
+ target_column: String,
+ query_vector_expr: Expr,
+ limit: usize,
+ schema: DFSchemaRef,
+ ) -> Self {
+ Self {
+ input,
+ target_table,
+ target_schema,
+ target_column,
+ query_vector_expr,
+ limit,
+ schema,
+ }
+ }
+
+ fn target_table(&self) -> &Table {
+ &self.target_table
+ }
+
+ fn target_schema(&self) -> &ArrowSchemaRef {
+ &self.target_schema
+ }
+
+ fn target_column(&self) -> &str {
+ &self.target_column
+ }
+
+ fn query_vector_expr(&self) -> &Expr {
+ &self.query_vector_expr
+ }
+
+ fn limit(&self) -> usize {
+ self.limit
+ }
+}
+
+impl UserDefinedLogicalNode for LateralVectorSearchNode {
+ fn as_any(&self) -> &dyn Any {
+ self
+ }
+
+ fn name(&self) -> &str {
+ "LateralVectorSearch"
+ }
+
+ fn inputs(&self) -> Vec<&LogicalPlan> {
+ vec![&self.input]
+ }
+
+ fn schema(&self) -> &DFSchemaRef {
+ &self.schema
+ }
+
+ fn check_invariants(&self, _check:
datafusion::logical_expr::InvariantLevel) -> DFResult<()> {
+ Ok(())
+ }
+
+ fn expressions(&self) -> Vec<Expr> {
+ vec![self.query_vector_expr.clone()]
+ }
+
+ fn fmt_for_explain(&self, f: &mut fmt::Formatter) -> fmt::Result {
+ write!(
+ f,
+ "LateralVectorSearch: column={}, limit={}",
+ self.target_column, self.limit
+ )
+ }
+
+ fn with_exprs_and_inputs(
+ &self,
+ exprs: Vec<Expr>,
+ inputs: Vec<LogicalPlan>,
+ ) -> DFResult<Arc<dyn UserDefinedLogicalNode>> {
+ if exprs.len() != 1 || inputs.len() != 1 {
+ return internal_err!("LateralVectorSearch expects one expression
and one input");
+ }
+ Ok(Arc::new(Self {
+ input: Arc::new(inputs.into_iter().next().unwrap()),
+ target_table: self.target_table.clone(),
+ target_schema: Arc::clone(&self.target_schema),
+ target_column: self.target_column.clone(),
+ query_vector_expr: exprs.into_iter().next().unwrap(),
+ limit: self.limit,
+ schema: Arc::clone(&self.schema),
+ }))
+ }
+
+ fn dyn_hash(&self, mut state: &mut dyn Hasher) {
+ self.name().hash(&mut state);
+ self.input.hash(&mut state);
+ self.target_table.location().hash(&mut state);
+ self.target_column.hash(&mut state);
+ self.query_vector_expr.hash(&mut state);
+ self.limit.hash(&mut state);
+ }
+
+ fn dyn_eq(&self, other: &dyn UserDefinedLogicalNode) -> bool {
+ other.as_any().downcast_ref::<Self>().is_some_and(|other| {
+ self.input == other.input
+ && self.target_table.location() ==
other.target_table.location()
+ && self.target_column == other.target_column
+ && self.query_vector_expr == other.query_vector_expr
+ && self.limit == other.limit
+ })
+ }
+
+ fn dyn_ord(&self, other: &dyn UserDefinedLogicalNode) -> Option<Ordering> {
+ let other = other.as_any().downcast_ref::<Self>()?;
+ if self.dyn_eq(other) {
+ Some(Ordering::Equal)
+ } else {
+ Some(format!("{self:?}").cmp(&format!("{other:?}")))
+ }
+ }
+}
+
+#[derive(Debug)]
+struct LateralVectorSearchExtensionPlanner;
+
+#[async_trait]
+impl ExtensionPlanner for LateralVectorSearchExtensionPlanner {
+ async fn plan_extension(
+ &self,
+ planner: &dyn PhysicalPlanner,
+ node: &dyn UserDefinedLogicalNode,
+ logical_inputs: &[&LogicalPlan],
+ physical_inputs: &[Arc<dyn ExecutionPlan>],
+ session_state: &SessionState,
+ ) -> DFResult<Option<Arc<dyn ExecutionPlan>>> {
+ let Some(node) =
node.as_any().downcast_ref::<LateralVectorSearchNode>() else {
+ return Ok(None);
+ };
+ if logical_inputs.len() != 1 || physical_inputs.len() != 1 {
+ return internal_err!("LateralVectorSearch physical planning
expects one input");
+ }
+
+ let query_vector_expr = planner.create_physical_expr(
+ node.query_vector_expr(),
+ logical_inputs[0].schema(),
+ session_state,
+ )?;
+ Ok(Some(Arc::new(LateralVectorSearchExec::new(
+ Arc::clone(&physical_inputs[0]),
+ node.target_table().clone(),
+ Arc::clone(node.target_schema()),
+ node.target_column().to_string(),
+ query_vector_expr,
+ node.limit(),
+ Arc::new(node.schema().as_arrow().clone()),
+ ))))
+ }
+}
+
+#[derive(Debug, Clone)]
+struct LateralVectorSearchExec {
+ input: Arc<dyn ExecutionPlan>,
+ target_table: Table,
+ target_schema: ArrowSchemaRef,
+ target_column: String,
+ query_vector_expr: Arc<dyn PhysicalExpr>,
+ limit: usize,
+ output_schema: ArrowSchemaRef,
+ plan_properties: Arc<PlanProperties>,
+}
+
+impl LateralVectorSearchExec {
+ fn new(
+ input: Arc<dyn ExecutionPlan>,
+ target_table: Table,
+ target_schema: ArrowSchemaRef,
+ target_column: String,
+ query_vector_expr: Arc<dyn PhysicalExpr>,
+ limit: usize,
+ output_schema: ArrowSchemaRef,
+ ) -> Self {
+ let partition_count = input.output_partitioning().partition_count();
+ let plan_properties = Arc::new(PlanProperties::new(
+
datafusion::physical_expr::EquivalenceProperties::new(output_schema.clone()),
+ Partitioning::UnknownPartitioning(partition_count),
+ EmissionType::Incremental,
+ Boundedness::Bounded,
+ ));
+ Self {
+ input,
+ target_table,
+ target_schema,
+ target_column,
+ query_vector_expr,
+ limit,
+ output_schema,
+ plan_properties,
+ }
+ }
+
+ async fn process_batch(&self, batch: RecordBatch) -> DFResult<RecordBatch>
{
+ if batch.num_rows() == 0 {
+ return empty_batch(self.output_schema.clone());
+ }
+
+ let vector_array = self
+ .query_vector_expr
+ .evaluate(&batch)?
+ .into_array(batch.num_rows())?;
+ let (query_vectors, left_query_rows) =
collect_query_vectors(&vector_array)?;
+ if query_vectors.is_empty() {
+ return empty_batch(self.output_schema.clone());
+ }
+
+ let mut builder = self.target_table.new_batch_vector_search_builder();
+ let results = builder
+ .with_vector_column(&self.target_column)
+ .with_query_vectors(query_vectors)
+ .with_limit(self.limit)
+ .execute()
+ .await
+ .map_err(to_datafusion_error)?;
+
+ let (target_batch, target_row_id_to_index) =
+ read_target_rows(&self.target_table, &self.target_schema,
&results).await?;
+
+ let mut left_indices = Vec::new();
+ let mut right_indices = Vec::new();
+ for (query_index, result) in results.iter().enumerate() {
+ let left_row = left_query_rows[query_index] as u32;
+ for row_id in &result.row_ids {
+ if let Some(&right_row) = target_row_id_to_index.get(row_id) {
+ left_indices.push(left_row);
+ right_indices.push(right_row);
+ }
+ }
+ }
+
+ if left_indices.is_empty() {
+ return empty_batch(self.output_schema.clone());
+ }
+
+ let left_indices = UInt32Array::from(left_indices);
+ let right_indices = UInt32Array::from(right_indices);
+ let mut columns = Vec::with_capacity(batch.num_columns() +
target_batch.num_columns());
+ for column in batch.columns() {
+ columns.push(arrow_select::take::take(
+ column.as_ref(),
+ &left_indices,
+ None,
+ )?);
+ }
+ for column in target_batch.columns() {
+ columns.push(arrow_select::take::take(
+ column.as_ref(),
+ &right_indices,
+ None,
+ )?);
+ }
+
+ RecordBatch::try_new(self.output_schema.clone(),
columns).map_err(DataFusionError::from)
+ }
+}
+
+impl DisplayAs for LateralVectorSearchExec {
+ fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) ->
fmt::Result {
+ write!(
+ f,
+ "LateralVectorSearchExec: column={}, limit={}",
+ self.target_column, self.limit
+ )
+ }
+}
+
+impl ExecutionPlan for LateralVectorSearchExec {
+ fn name(&self) -> &str {
+ "LateralVectorSearchExec"
+ }
+
+ fn as_any(&self) -> &dyn Any {
+ self
+ }
+
+ fn properties(&self) -> &Arc<PlanProperties> {
+ &self.plan_properties
+ }
+
+ fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
+ vec![&self.input]
+ }
+
+ fn with_new_children(
+ self: Arc<Self>,
+ mut children: Vec<Arc<dyn ExecutionPlan>>,
+ ) -> DFResult<Arc<dyn ExecutionPlan>> {
+ if children.len() != 1 {
+ return internal_err!("LateralVectorSearchExec expects one child");
+ }
+ Ok(Arc::new(Self::new(
+ children.remove(0),
+ self.target_table.clone(),
+ Arc::clone(&self.target_schema),
+ self.target_column.clone(),
+ Arc::clone(&self.query_vector_expr),
+ self.limit,
+ Arc::clone(&self.output_schema),
+ )))
+ }
+
+ fn execute(
+ &self,
+ partition: usize,
+ context: Arc<TaskContext>,
+ ) -> DFResult<SendableRecordBatchStream> {
+ let input = self.input.execute(partition, context)?;
+ let exec = self.clone();
+ let stream = input.then(move |batch| {
+ let exec = exec.clone();
+ async move {
+ let batch = batch?;
+ exec.process_batch(batch).await
+ }
+ });
+ Ok(Box::pin(RecordBatchStreamAdapter::new(
+ self.output_schema.clone(),
+ Box::pin(stream),
+ )))
+ }
+
+ fn partition_statistics(&self, _partition: Option<usize>) ->
DFResult<Statistics> {
+ Ok(Statistics {
+ num_rows: Precision::Absent,
+ total_byte_size: Precision::Absent,
+ column_statistics: Statistics::unknown_column(&self.output_schema),
+ })
+ }
+}
+
+fn collect_query_vectors(array: &ArrayRef) -> DFResult<(Vec<Vec<f32>>,
Vec<usize>)> {
+ enum VectorLayout<'a> {
+ List(&'a ListArray),
+ Fixed(&'a FixedSizeListArray),
+ }
+ let layout = if let Some(array) =
array.as_any().downcast_ref::<ListArray>() {
+ VectorLayout::List(array)
+ } else if let Some(array) =
array.as_any().downcast_ref::<FixedSizeListArray>() {
+ VectorLayout::Fixed(array)
+ } else {
+ return Err(DataFusionError::Plan(
+ "lateral vector_search query vector must be List<Float32> or
FixedSizeList<Float32>"
+ .to_string(),
+ ));
+ };
+ let values = match layout {
+ VectorLayout::List(array) => array.values(),
+ VectorLayout::Fixed(array) => array.values(),
+ }
+ .as_any()
+ .downcast_ref::<Float32Array>()
+ .ok_or_else(|| {
+ DataFusionError::Plan(
+ "lateral vector_search query vector elements must be
Float32".to_string(),
+ )
+ })?;
+
+ let row_count = match layout {
+ VectorLayout::List(array) => array.len(),
+ VectorLayout::Fixed(array) => array.len(),
+ };
+ let mut vectors = Vec::new();
+ let mut rows = Vec::new();
+ for row in 0..row_count {
+ let is_null = match layout {
+ VectorLayout::List(array) => array.is_null(row),
+ VectorLayout::Fixed(array) => array.is_null(row),
+ };
+ if is_null {
+ continue;
+ }
+
+ let (start, end) = match layout {
+ VectorLayout::List(array) => {
+ let offsets = array.value_offsets();
+ (offsets[row] as usize, offsets[row + 1] as usize)
+ }
+ VectorLayout::Fixed(array) => {
+ let len = array.value_length() as usize;
+ (row * len, (row + 1) * len)
+ }
+ };
+ let mut vector = Vec::with_capacity(end - start);
+ for value_index in start..end {
+ if values.is_null(value_index) {
+ return Err(DataFusionError::Plan(
+ "lateral vector_search query vector cannot contain null
elements".to_string(),
+ ));
+ }
+ vector.push(values.value(value_index));
+ }
+ vectors.push(vector);
+ rows.push(row);
+ }
+ Ok((vectors, rows))
+}
+
+async fn read_target_rows(
+ table: &Table,
+ target_schema: &ArrowSchemaRef,
+ results: &[SearchResult],
+) -> DFResult<(RecordBatch, HashMap<u64, u32>)> {
+ let mut row_ids = results
+ .iter()
+ .flat_map(|result| result.row_ids.iter().copied())
+ .collect::<Vec<_>>();
+ row_ids.sort_unstable();
+ row_ids.dedup();
+ if row_ids.is_empty() {
+ return Ok((empty_batch(target_schema.clone())?, HashMap::new()));
+ }
+
+ let row_ranges = row_ranges_from_row_ids(&row_ids)?;
+ let mut projection = target_schema
+ .fields()
+ .iter()
+ .map(|field| field.name().to_string())
+ .collect::<Vec<_>>();
+ if !projection.iter().any(|column| column == ROW_ID_FIELD_NAME) {
+ projection.push(ROW_ID_FIELD_NAME.to_string());
+ }
+ let projection_refs =
projection.iter().map(String::as_str).collect::<Vec<_>>();
+
+ let mut read_builder = table.new_read_builder();
+ read_builder
+ .with_projection(&projection_refs)
+ .with_row_ranges(row_ranges);
+ let plan = read_builder
+ .new_scan()
+ .plan()
+ .await
+ .map_err(to_datafusion_error)?;
+ if plan.splits().is_empty() {
+ return Ok((empty_batch(target_schema.clone())?, HashMap::new()));
+ }
+
+ let read = read_builder.new_read().map_err(to_datafusion_error)?;
+ let mut stream =
read.to_arrow(plan.splits()).map_err(to_datafusion_error)?;
+ let mut batches = Vec::new();
+ while let Some(batch) =
stream.try_next().await.map_err(to_datafusion_error)? {
+ batches.push(batch);
+ }
+ if batches.is_empty() {
+ return Ok((empty_batch(target_schema.clone())?, HashMap::new()));
+ }
+
+ let read_schema = batches[0].schema();
+ let batch = arrow_select::concat::concat_batches(&read_schema, &batches)
+ .map_err(DataFusionError::from)?;
+ let row_id_index = batch
+ .schema()
+ .index_of(ROW_ID_FIELD_NAME)
+ .map_err(DataFusionError::from)?;
+ let row_id_array = batch
+ .column(row_id_index)
+ .as_any()
+ .downcast_ref::<Int64Array>()
+ .ok_or_else(|| {
+ DataFusionError::Internal(
+ "_ROW_ID must be Int64 in vector search target
read".to_string(),
+ )
+ })?;
+
+ let mut row_id_to_index = HashMap::new();
+ for row in 0..batch.num_rows() {
+ if row_id_array.is_null(row) {
+ continue;
+ }
+ let row_id = u64::try_from(row_id_array.value(row)).map_err(|_| {
+ DataFusionError::Internal(format!(
+ "negative _ROW_ID {} in vector search target read",
+ row_id_array.value(row)
+ ))
+ })?;
+ row_id_to_index.insert(row_id, row as u32);
+ }
+
+ let target_columns = (0..target_schema.fields().len())
+ .map(|index| Arc::clone(batch.column(index)))
+ .collect::<Vec<_>>();
+ let target_batch = RecordBatch::try_new(target_schema.clone(),
target_columns)
+ .map_err(DataFusionError::from)?;
+ Ok((target_batch, row_id_to_index))
+}
+
+fn row_ranges_from_row_ids(row_ids: &[u64]) -> DFResult<Vec<RowRange>> {
+ let scores = vec![0.0; row_ids.len()];
+ SearchResult::new(row_ids.to_vec(), scores)
+ .to_row_ranges()
+ .map_err(to_datafusion_error)
+}
+
+fn empty_batch(schema: ArrowSchemaRef) -> DFResult<RecordBatch> {
+ let columns = schema
+ .fields()
+ .iter()
+ .map(|field| new_empty_array(field.data_type()))
+ .collect::<Vec<_>>();
+ RecordBatch::try_new(schema, columns).map_err(DataFusionError::from)
+}
+
+pub(crate) fn session_config() -> SessionConfig {
+ SessionConfig::new().with_information_schema(true)
+}
diff --git a/crates/integrations/datafusion/src/lib.rs
b/crates/integrations/datafusion/src/lib.rs
index f11cfce..de6e932 100644
--- a/crates/integrations/datafusion/src/lib.rs
+++ b/crates/integrations/datafusion/src/lib.rs
@@ -43,6 +43,7 @@ mod error;
mod filter_pushdown;
#[cfg(feature = "fulltext")]
mod full_text_search;
+mod lateral_vector_search;
mod merge_into;
mod physical_plan;
mod procedures;
diff --git a/crates/integrations/datafusion/src/sql_context.rs
b/crates/integrations/datafusion/src/sql_context.rs
index 515cead..8ccdc50 100644
--- a/crates/integrations/datafusion/src/sql_context.rs
+++ b/crates/integrations/datafusion/src/sql_context.rs
@@ -46,7 +46,8 @@ use datafusion::arrow::record_batch::RecordBatch;
use datafusion::common::TableReference;
use datafusion::datasource::{MemTable, TableProvider};
use datafusion::error::{DataFusionError, Result as DFResult};
-use datafusion::prelude::{DataFrame, SessionConfig, SessionContext};
+use datafusion::execution::SessionStateBuilder;
+use datafusion::prelude::{DataFrame, SessionContext};
use datafusion::sql::sqlparser::ast::{
AlterTableOperation, ColumnDef, CreateTable, CreateTableOptions,
CreateView, Delete,
Expr as SqlExpr, FromTable, Insert, Merge, ObjectName, ObjectType,
RenameTableNameKind, Reset,
@@ -93,12 +94,18 @@ impl Default for SQLContext {
impl SQLContext {
/// Creates a new empty SQL context.
pub fn new() -> Self {
- let ctx =
-
SessionContext::new_with_config(SessionConfig::new().with_information_schema(true));
- ctx.register_relation_planner(Arc::new(
- crate::relation_planner::PaimonRelationPlanner::new(),
- ))
- .expect("failed to register relation planner");
+ let state = SessionStateBuilder::new()
+ .with_config(crate::lateral_vector_search::session_config())
+ .with_default_features()
+ .with_relation_planners(vec![Arc::new(
+ crate::relation_planner::PaimonRelationPlanner::new(),
+ )])
+
.with_optimizer_rules(crate::lateral_vector_search::optimizer_rules())
+ .with_query_planner(Arc::new(
+ crate::lateral_vector_search::PaimonQueryPlanner::new(),
+ ))
+ .build();
+ let ctx = SessionContext::new_with_state(state);
Self {
ctx,
catalogs: HashMap::new(),
diff --git a/crates/integrations/datafusion/src/vector_search.rs
b/crates/integrations/datafusion/src/vector_search.rs
index 5b53d14..536e379 100644
--- a/crates/integrations/datafusion/src/vector_search.rs
+++ b/crates/integrations/datafusion/src/vector_search.rs
@@ -25,7 +25,7 @@ use datafusion::catalog::Session;
use datafusion::catalog::TableFunctionImpl;
use datafusion::common::project_schema;
use datafusion::datasource::{TableProvider, TableType};
-use datafusion::error::Result as DFResult;
+use datafusion::error::{DataFusionError, Result as DFResult};
use datafusion::logical_expr::{Expr, TableProviderFilterPushDown};
use datafusion::physical_plan::empty::EmptyExec;
use datafusion::physical_plan::ExecutionPlan;
@@ -84,29 +84,14 @@ impl TableFunctionImpl for VectorSearchFunction {
let table_name = extract_string_literal(FUNCTION_NAME, &args[0],
"table_name")?;
let column_name = extract_string_literal(FUNCTION_NAME, &args[1],
"column_name")?;
- let query_vector_json =
- extract_string_literal(FUNCTION_NAME, &args[2],
"query_vector_json")?;
let limit = extract_int_literal(FUNCTION_NAME, &args[3], "limit")?;
if limit <= 0 {
- return Err(datafusion::error::DataFusionError::Plan(
+ return Err(DataFusionError::Plan(
"vector_search: limit must be positive".to_string(),
));
}
- let query_vector: Vec<f32> =
serde_json::from_str(&query_vector_json).map_err(|e| {
- datafusion::error::DataFusionError::Plan(format!(
- "vector_search: query_vector_json must be a JSON array of
floats, got '{}': {}",
- query_vector_json, e
- ))
- })?;
-
- if query_vector.is_empty() {
- return Err(datafusion::error::DataFusionError::Plan(
- "vector_search: query vector cannot be empty".to_string(),
- ));
- }
-
let identifier =
parse_table_identifier(FUNCTION_NAME, &table_name,
&self.default_database)?;
@@ -118,6 +103,32 @@ impl TableFunctionImpl for VectorSearchFunction {
.map_err(to_datafusion_error)?;
let inner = PaimonTableProvider::try_new(table)?;
+ let query_vector_json =
+ match extract_string_literal(FUNCTION_NAME, &args[2],
"query_vector_json") {
+ Ok(value) => value,
+ Err(_) if matches!(args[2], Expr::Column(_)) => {
+ return Ok(Arc::new(LateralVectorSearchTableProvider {
+ inner,
+ column_name,
+ query_vector_expr: args[2].clone(),
+ limit: limit as usize,
+ }));
+ }
+ Err(err) => return Err(err),
+ };
+
+ let query_vector: Vec<f32> =
serde_json::from_str(&query_vector_json).map_err(|e| {
+ DataFusionError::Plan(format!(
+ "vector_search: query_vector_json must be a JSON array of
floats, got '{}': {}",
+ query_vector_json, e
+ ))
+ })?;
+
+ if query_vector.is_empty() {
+ return Err(DataFusionError::Plan(
+ "vector_search: query vector cannot be empty".to_string(),
+ ));
+ }
Ok(Arc::new(VectorSearchTableProvider {
inner,
@@ -128,6 +139,59 @@ impl TableFunctionImpl for VectorSearchFunction {
}
}
+#[derive(Debug)]
+pub(crate) struct LateralVectorSearchTableProvider {
+ inner: PaimonTableProvider,
+ column_name: String,
+ query_vector_expr: Expr,
+ limit: usize,
+}
+
+impl LateralVectorSearchTableProvider {
+ pub(crate) fn inner(&self) -> &PaimonTableProvider {
+ &self.inner
+ }
+
+ pub(crate) fn column_name(&self) -> &str {
+ &self.column_name
+ }
+
+ pub(crate) fn query_vector_expr(&self) -> &Expr {
+ &self.query_vector_expr
+ }
+
+ pub(crate) fn limit(&self) -> usize {
+ self.limit
+ }
+}
+
+#[async_trait]
+impl TableProvider for LateralVectorSearchTableProvider {
+ fn as_any(&self) -> &dyn Any {
+ self
+ }
+
+ fn schema(&self) -> ArrowSchemaRef {
+ self.inner.schema()
+ }
+
+ fn table_type(&self) -> TableType {
+ TableType::Base
+ }
+
+ async fn scan(
+ &self,
+ _state: &dyn Session,
+ _projection: Option<&Vec<usize>>,
+ _filters: &[Expr],
+ _limit: Option<usize>,
+ ) -> DFResult<Arc<dyn ExecutionPlan>> {
+ Err(DataFusionError::Plan(
+ "lateral vector_search must be planned through a lateral
join".to_string(),
+ ))
+ }
+}
+
#[derive(Debug)]
struct VectorSearchTableProvider {
inner: PaimonTableProvider,
diff --git a/crates/integrations/datafusion/tests/read_tables.rs
b/crates/integrations/datafusion/tests/read_tables.rs
index cc4c781..1d72380 100644
--- a/crates/integrations/datafusion/tests/read_tables.rs
+++ b/crates/integrations/datafusion/tests/read_tables.rs
@@ -1343,6 +1343,7 @@ mod vector_search_tests {
DataType as ArrowDataType, Field as ArrowField, Schema as ArrowSchema,
};
use datafusion::arrow::record_batch::RecordBatch;
+ use datafusion::datasource::MemTable;
use paimon::catalog::Identifier;
use paimon::spec::{ArrayType, DataType, FloatType, IntType, Schema};
use paimon::{Catalog, CatalogOptions, FileSystemCatalog, Options};
@@ -1467,6 +1468,27 @@ mod vector_search_tests {
ids
}
+ fn extract_query_result_ids(
+ batches: &[datafusion::arrow::record_batch::RecordBatch],
+ ) -> Vec<(i32, i32)> {
+ let mut rows = Vec::new();
+ for batch in batches {
+ let query_id_array = batch
+ .column_by_name("query_id")
+ .and_then(|c| c.as_any().downcast_ref::<Int32Array>())
+ .expect("Expected Int32Array for query_id");
+ let result_id_array = batch
+ .column_by_name("result_id")
+ .and_then(|c| c.as_any().downcast_ref::<Int32Array>())
+ .expect("Expected Int32Array for result_id");
+ for i in 0..batch.num_rows() {
+ rows.push((query_id_array.value(i), result_id_array.value(i)));
+ }
+ }
+ rows.sort();
+ rows
+ }
+
fn extract_index_rows(
batches: &[datafusion::arrow::record_batch::RecordBatch],
) -> Vec<(String, i64, i64, i64, String)> {
@@ -1586,6 +1608,35 @@ mod vector_search_tests {
assert_eq!(ids, vec![0, 1, 2]);
}
+ #[tokio::test]
+ async fn test_vector_search_lateral_join_uses_query_vectors() {
+ let (ctx, _tmp) = create_java_vindex_vector_search_context().await;
+ let query_batch = build_vector_batch(
+ vec![10, 20],
+ vec![vec![1.0, 0.0, 0.0, 0.0], vec![0.0, 1.0, 0.0, 0.0]],
+ );
+ let query_table = MemTable::try_new(query_batch.schema(),
vec![vec![query_batch]])
+ .expect("Failed to create query vector table");
+ ctx.register_temp_table("paimon.default.queries",
Arc::new(query_table))
+ .expect("Failed to register query vector table");
+
+ let batches = ctx
+ .sql(
+ "SELECT q.id AS query_id, r.id AS result_id \
+ FROM paimon.default.queries q \
+ CROSS JOIN LATERAL
vector_search('paimon.default.test_java_vindex_vector', 'embedding',
q.embedding, 2) AS r \
+ ORDER BY query_id, result_id",
+ )
+ .await
+ .expect("lateral vector_search SQL should parse")
+ .collect()
+ .await
+ .expect("lateral vector_search query should execute");
+
+ let rows = extract_query_result_ids(&batches);
+ assert_eq!(rows, vec![(10, 0), (10, 1), (20, 1), (20, 2)]);
+ }
+
// Manual run with a local Lumina native library:
// LUMINA_LIB_PATH=/path/to/liblumina_py.so cargo test -p
paimon-datafusion \
// vector_search_tests::test_lumina_build_then_vector_search_query \