1use std::any::Any;
16use std::pin::Pin;
17use std::sync::Arc;
18use std::task::{Context, Poll};
19
20use ahash::{HashSet, RandomState};
21use datafusion::arrow::array::UInt64Array;
22use datafusion::arrow::datatypes::SchemaRef;
23use datafusion::arrow::record_batch::RecordBatch;
24use datafusion::common::{DFSchema, DFSchemaRef};
25use datafusion::error::{DataFusionError, Result as DataFusionResult};
26use datafusion::execution::context::TaskContext;
27use datafusion::logical_expr::{EmptyRelation, Expr, LogicalPlan, UserDefinedLogicalNodeCore};
28use datafusion::physical_expr::EquivalenceProperties;
29use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
30use datafusion::physical_plan::metrics::{BaselineMetrics, ExecutionPlanMetricsSet, MetricsSet};
31use datafusion::physical_plan::{
32 DisplayAs, DisplayFormatType, Distribution, ExecutionPlan, Partitioning, PlanProperties,
33 RecordBatchStream, SendableRecordBatchStream, hash_utils,
34};
35use datafusion_expr::col;
36use datatypes::arrow::compute;
37use futures::{Stream, StreamExt, ready};
38use greptime_proto::substrait_extension as pb;
39use prost::Message;
40use snafu::ResultExt;
41
42use crate::error::{DataFusionPlanningSnafu, DeserializeSnafu, Result};
43
44#[derive(Debug, PartialEq, Eq, Hash)]
59pub struct UnionDistinctOn {
60 left: LogicalPlan,
61 right: LogicalPlan,
62 compare_key_indices: Vec<usize>,
65 ts_col_idx: usize,
66 output_schema: DFSchemaRef,
67}
68
69impl UnionDistinctOn {
70 pub fn name() -> &'static str {
71 "UnionDistinctOn"
72 }
73
74 pub fn try_new(
75 left: LogicalPlan,
76 right: LogicalPlan,
77 compare_key_indices: Vec<usize>,
78 ts_col_idx: usize,
79 ) -> DataFusionResult<Self> {
80 let output_schema =
81 Self::validate_children(&left, &right, &compare_key_indices, ts_col_idx)?;
82 Ok(Self {
83 left,
84 right,
85 compare_key_indices,
86 ts_col_idx,
87 output_schema,
88 })
89 }
90
91 fn validate_children(
92 left: &LogicalPlan,
93 right: &LogicalPlan,
94 compare_key_indices: &[usize],
95 ts_col_idx: usize,
96 ) -> DataFusionResult<DFSchemaRef> {
97 let left_schema = left.schema();
98 let right_schema = right.schema();
99 let left_fields = left_schema.fields();
100 let right_fields = right_schema.fields();
101
102 if left_fields.len() != right_fields.len() {
103 return Err(DataFusionError::Plan(format!(
104 "UnionDistinctOn inputs have different field counts: left={}, right={}",
105 left_fields.len(),
106 right_fields.len()
107 )));
108 }
109
110 for (column_type, index) in compare_key_indices
111 .iter()
112 .map(|index| ("compare key", *index))
113 .chain(std::iter::once(("timestamp", ts_col_idx)))
114 {
115 if index >= left_fields.len() || index >= right_fields.len() {
116 return Err(DataFusionError::Plan(format!(
117 "UnionDistinctOn {column_type} index {index} is out of bounds for inputs with {} fields",
118 left_fields.len()
119 )));
120 }
121 }
122
123 for (index, (left_field, right_field)) in left_fields.iter().zip(right_fields).enumerate() {
124 if left_field.data_type() != right_field.data_type() {
125 return Err(DataFusionError::Plan(format!(
126 "UnionDistinctOn input field at index {index} has incompatible data types: left={:?}, right={:?}",
127 left_field.data_type(),
128 right_field.data_type()
129 )));
130 }
131 }
132
133 let output_fields = left_fields
134 .iter()
135 .zip(right_fields)
136 .enumerate()
137 .map(|(index, (left_field, right_field))| {
138 let (qualifier, _) = left_schema.qualified_field(index);
139 (
140 qualifier.cloned(),
141 Arc::new(
142 left_field
143 .as_ref()
144 .clone()
145 .with_nullable(left_field.is_nullable() || right_field.is_nullable()),
146 ),
147 )
148 })
149 .collect();
150 let output_schema =
151 DFSchema::new_with_metadata(output_fields, left_schema.metadata().clone()).map_err(
152 |error| {
153 DataFusionError::Plan(format!(
154 "Failed to construct UnionDistinctOn output schema: {error}"
155 ))
156 },
157 )?;
158
159 Ok(Arc::new(output_schema))
160 }
161
162 pub fn to_execution_plan(
163 &self,
164 left_exec: Arc<dyn ExecutionPlan>,
165 right_exec: Arc<dyn ExecutionPlan>,
166 ) -> Arc<dyn ExecutionPlan> {
167 let output_schema: SchemaRef = self.output_schema.inner().clone();
168 let properties = Arc::new(PlanProperties::new(
169 EquivalenceProperties::new(output_schema.clone()),
170 Partitioning::UnknownPartitioning(1),
171 EmissionType::Incremental,
172 Boundedness::Bounded,
173 ));
174 Arc::new(UnionDistinctOnExec {
175 left: left_exec,
176 right: right_exec,
177 compare_key_indices: self.compare_key_indices.clone(),
178 ts_col_idx: self.ts_col_idx,
179 output_schema,
180 metric: ExecutionPlanMetricsSet::new(),
181 properties,
182 random_state: RandomState::new(),
183 })
184 }
185
186 pub fn serialize(&self) -> Vec<u8> {
187 let compare_key_indices = self
188 .compare_key_indices
189 .iter()
190 .map(|index| u64::try_from(*index).expect("usize always fits in u64"))
191 .collect();
192 let ts_col_idx = u64::try_from(self.ts_col_idx).expect("usize always fits in u64");
193
194 pb::UnionDistinctOn {
195 compare_key_indices,
196 ts_col_idx,
197 }
198 .encode_to_vec()
199 }
200
201 pub fn deserialize(bytes: &[u8]) -> Result<Self> {
202 let pb_union = pb::UnionDistinctOn::decode(bytes).context(DeserializeSnafu)?;
203 let placeholder_plan = LogicalPlan::EmptyRelation(EmptyRelation {
204 produce_one_row: false,
205 schema: Arc::new(DFSchema::empty()),
206 });
207
208 let compare_key_indices = pb_union
209 .compare_key_indices
210 .into_iter()
211 .map(|index| {
212 usize::try_from(index).map_err(|_| {
213 DataFusionError::Plan(format!(
214 "UnionDistinctOn compare key index {index} does not fit in usize"
215 ))
216 })
217 })
218 .collect::<DataFusionResult<Vec<_>>>()
219 .context(DataFusionPlanningSnafu)?;
220 let ts_col_idx = usize::try_from(pb_union.ts_col_idx)
221 .map_err(|_| {
222 DataFusionError::Plan(format!(
223 "UnionDistinctOn timestamp index {} does not fit in usize",
224 pb_union.ts_col_idx
225 ))
226 })
227 .context(DataFusionPlanningSnafu)?;
228
229 Ok(Self {
230 left: placeholder_plan.clone(),
231 right: placeholder_plan,
232 compare_key_indices,
233 ts_col_idx,
234 output_schema: Arc::new(DFSchema::empty()),
235 })
236 }
237}
238
239impl PartialOrd for UnionDistinctOn {
240 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
241 match self.left.partial_cmp(&other.left) {
243 Some(core::cmp::Ordering::Equal) => {}
244 ord => return ord,
245 }
246 match self.right.partial_cmp(&other.right) {
247 Some(core::cmp::Ordering::Equal) => {}
248 ord => return ord,
249 }
250 match self
251 .compare_key_indices
252 .partial_cmp(&other.compare_key_indices)
253 {
254 Some(core::cmp::Ordering::Equal) => {}
255 ord => return ord,
256 }
257 self.ts_col_idx.partial_cmp(&other.ts_col_idx)
258 }
259}
260
261impl UserDefinedLogicalNodeCore for UnionDistinctOn {
262 fn name(&self) -> &str {
263 Self::name()
264 }
265
266 fn inputs(&self) -> Vec<&LogicalPlan> {
267 vec![&self.left, &self.right]
268 }
269
270 fn schema(&self) -> &DFSchemaRef {
271 &self.output_schema
272 }
273
274 fn expressions(&self) -> Vec<Expr> {
275 let fields = self.left.schema().fields();
276 let mut exprs = self
277 .compare_key_indices
278 .iter()
279 .filter_map(|index| fields.get(*index).map(|field| col(field.name())))
280 .collect::<Vec<_>>();
281 if !self.compare_key_indices.contains(&self.ts_col_idx)
282 && let Some(field) = fields.get(self.ts_col_idx)
283 {
284 exprs.push(col(field.name()));
285 }
286 exprs
287 }
288
289 fn necessary_children_exprs(&self, _output_columns: &[usize]) -> Option<Vec<Vec<usize>>> {
290 let left_len = self.left.schema().fields().len();
291 let right_len = self.right.schema().fields().len();
292 Some(vec![
293 (0..left_len).collect::<Vec<_>>(),
294 (0..right_len).collect::<Vec<_>>(),
295 ])
296 }
297
298 fn fmt_for_explain(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
299 let fields = self.left.schema().fields();
300 let display_column = |index: usize| match fields.get(index) {
301 Some(field) => format!("{}@{index}", field.name()),
302 None => format!("@{index}"),
303 };
304 let compare_keys = self
305 .compare_key_indices
306 .iter()
307 .map(|index| display_column(*index))
308 .collect::<Vec<_>>();
309 write!(
310 f,
311 "UnionDistinctOn: on col={compare_keys:?}, ts_col={}",
312 display_column(self.ts_col_idx)
313 )
314 }
315
316 fn with_exprs_and_inputs(
317 &self,
318 _exprs: Vec<Expr>,
319 inputs: Vec<LogicalPlan>,
320 ) -> DataFusionResult<Self> {
321 if inputs.len() != 2 {
322 return Err(DataFusionError::Internal(
323 "UnionDistinctOn must have exactly 2 inputs".to_string(),
324 ));
325 }
326
327 let mut inputs = inputs.into_iter();
328 let left = inputs.next().unwrap();
329 let right = inputs.next().unwrap();
330
331 Self::try_new(
332 left,
333 right,
334 self.compare_key_indices.clone(),
335 self.ts_col_idx,
336 )
337 }
338}
339
340#[derive(Debug)]
341pub struct UnionDistinctOnExec {
342 left: Arc<dyn ExecutionPlan>,
343 right: Arc<dyn ExecutionPlan>,
344 compare_key_indices: Vec<usize>,
345 ts_col_idx: usize,
346 output_schema: SchemaRef,
347 metric: ExecutionPlanMetricsSet,
348 properties: Arc<PlanProperties>,
349
350 random_state: RandomState,
352}
353
354impl ExecutionPlan for UnionDistinctOnExec {
355 fn as_any(&self) -> &dyn Any {
356 self
357 }
358
359 fn schema(&self) -> SchemaRef {
360 self.output_schema.clone()
361 }
362
363 fn required_input_distribution(&self) -> Vec<Distribution> {
364 vec![Distribution::SinglePartition, Distribution::SinglePartition]
365 }
366
367 fn properties(&self) -> &Arc<PlanProperties> {
368 &self.properties
369 }
370
371 fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
372 vec![&self.left, &self.right]
373 }
374
375 fn with_new_children(
376 self: Arc<Self>,
377 children: Vec<Arc<dyn ExecutionPlan>>,
378 ) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
379 assert_eq!(children.len(), 2);
380
381 let left = children[0].clone();
382 let right = children[1].clone();
383 Ok(Arc::new(UnionDistinctOnExec {
384 left,
385 right,
386 compare_key_indices: self.compare_key_indices.clone(),
387 ts_col_idx: self.ts_col_idx,
388 output_schema: self.output_schema.clone(),
389 metric: self.metric.clone(),
390 properties: self.properties.clone(),
391 random_state: self.random_state.clone(),
392 }))
393 }
394
395 fn execute(
396 &self,
397 partition: usize,
398 context: Arc<TaskContext>,
399 ) -> DataFusionResult<SendableRecordBatchStream> {
400 let left_stream = self.left.execute(partition, context.clone())?;
401
402 let mut key_indices = self.compare_key_indices.clone();
403 key_indices.push(self.ts_col_idx);
404
405 Ok(Box::pin(UnionDistinctOnStream {
406 left: left_stream,
407 right_plan: self.right.clone(),
408 right_partition: partition,
409 right_context: context,
410 right: None,
411 compare_keys: key_indices,
412 output_schema: self.output_schema.clone(),
413 random_state: self.random_state.clone(),
414 lhs_signatures: HashSet::default(),
415 hashes: Vec::new(),
416 phase: StreamPhase::Left,
417 metric: BaselineMetrics::new(&self.metric, partition),
418 }))
419 }
420
421 fn metrics(&self) -> Option<MetricsSet> {
422 Some(self.metric.clone_inner())
423 }
424
425 fn name(&self) -> &str {
426 "UnionDistinctOnExec"
427 }
428}
429
430impl DisplayAs for UnionDistinctOnExec {
431 fn fmt_as(&self, t: DisplayFormatType, f: &mut std::fmt::Formatter) -> std::fmt::Result {
432 match t {
433 DisplayFormatType::Default
434 | DisplayFormatType::Verbose
435 | DisplayFormatType::TreeRender => {
436 write!(
437 f,
438 "UnionDistinctOnExec: on col={:?}, ts_col={}",
439 self.compare_key_indices, self.ts_col_idx
440 )
441 }
442 }
443 }
444}
445
446pub struct UnionDistinctOnStream {
447 left: SendableRecordBatchStream,
448 right_plan: Arc<dyn ExecutionPlan>,
449 right_partition: usize,
450 right_context: Arc<TaskContext>,
451 right: Option<SendableRecordBatchStream>,
452 compare_keys: Vec<usize>,
454 output_schema: SchemaRef,
455 random_state: RandomState,
456 lhs_signatures: HashSet<u64>,
457 hashes: Vec<u64>,
458 phase: StreamPhase,
459 metric: BaselineMetrics,
460}
461
462#[derive(Debug, Clone, Copy, PartialEq, Eq)]
463enum StreamPhase {
464 Left,
465 Right,
466 Done,
467}
468
469impl UnionDistinctOnStream {
470 fn hash_batch(&mut self, batch: &RecordBatch) -> DataFusionResult<()> {
471 let arrays = self
472 .compare_keys
473 .iter()
474 .map(|index| batch.column(*index).clone())
475 .collect::<Vec<_>>();
476 self.hashes.clear();
477 self.hashes.resize(batch.num_rows(), 0);
478 hash_utils::create_hashes(&arrays, &self.random_state, &mut self.hashes)?;
479 Ok(())
480 }
481
482 fn filter_rhs_batch(&mut self, batch: RecordBatch) -> DataFusionResult<Option<RecordBatch>> {
483 self.hash_batch(&batch)?;
484 let mut survivor_indices: Option<Vec<usize>> = None;
485 for (index, hash) in self.hashes.iter().enumerate() {
486 if self.lhs_signatures.contains(hash) {
487 survivor_indices.get_or_insert_with(|| (0..index).collect());
488 } else if let Some(indices) = &mut survivor_indices {
489 indices.push(index);
490 }
491 }
492
493 match survivor_indices {
494 None => Ok(Some(with_schema(batch, self.output_schema.clone())?)),
495 Some(indices) if indices.is_empty() => Ok(None),
496 Some(indices) => Ok(Some(with_schema(
497 take_batch(&batch, &indices)?,
498 self.output_schema.clone(),
499 )?)),
500 }
501 }
502
503 fn terminal_error(&mut self, error: DataFusionError) -> Poll<Option<<Self as Stream>::Item>> {
504 self.phase = StreamPhase::Done;
505 Poll::Ready(Some(Err(error)))
506 }
507
508 fn poll_right(&mut self, cx: &mut Context<'_>) -> Poll<Option<<Self as Stream>::Item>> {
509 if self.right.is_none() {
510 let right = match self
511 .right_plan
512 .execute(self.right_partition, self.right_context.clone())
513 {
514 Ok(right) => right,
515 Err(error) => return self.terminal_error(error),
516 };
517 self.right = Some(right);
518 }
519
520 let right = self.right.as_mut().expect("right stream is initialized");
521 match ready!(right.poll_next_unpin(cx)) {
522 Some(Ok(batch)) if self.lhs_signatures.is_empty() => {
523 match with_schema(batch, self.output_schema.clone()) {
524 Ok(batch) => Poll::Ready(Some(Ok(batch))),
525 Err(error) => self.terminal_error(error),
526 }
527 }
528 Some(Ok(batch)) => match self.filter_rhs_batch(batch) {
529 Ok(Some(batch)) => Poll::Ready(Some(Ok(batch))),
530 Ok(None) => {
531 cx.waker().wake_by_ref();
534 Poll::Pending
535 }
536 Err(error) => self.terminal_error(error),
537 },
538 Some(Err(error)) => self.terminal_error(error),
539 None => {
540 self.phase = StreamPhase::Done;
541 Poll::Ready(None)
542 }
543 }
544 }
545
546 fn poll_impl(&mut self, cx: &mut Context<'_>) -> Poll<Option<<Self as Stream>::Item>> {
547 match self.phase {
548 StreamPhase::Left => match ready!(self.left.poll_next_unpin(cx)) {
549 Some(Ok(batch)) => {
550 if batch.num_rows() > 0 {
551 if let Err(error) = self.hash_batch(&batch) {
552 return self.terminal_error(error);
553 }
554 self.lhs_signatures.extend(self.hashes.iter().copied());
555 }
556 match with_schema(batch, self.output_schema.clone()) {
557 Ok(batch) => Poll::Ready(Some(Ok(batch))),
558 Err(error) => self.terminal_error(error),
559 }
560 }
561 Some(Err(error)) => self.terminal_error(error),
562 None => {
563 self.phase = StreamPhase::Right;
564 self.poll_right(cx)
565 }
566 },
567 StreamPhase::Right => self.poll_right(cx),
568 StreamPhase::Done => Poll::Ready(None),
569 }
570 }
571}
572
573impl RecordBatchStream for UnionDistinctOnStream {
574 fn schema(&self) -> SchemaRef {
575 self.output_schema.clone()
576 }
577}
578
579impl Stream for UnionDistinctOnStream {
580 type Item = DataFusionResult<RecordBatch>;
581
582 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
583 let poll = self.poll_impl(cx);
584 self.metric.record_poll(poll)
585 }
586}
587
588fn with_schema(batch: RecordBatch, schema: SchemaRef) -> DataFusionResult<RecordBatch> {
589 RecordBatch::try_new(schema, batch.columns().to_vec())
590 .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))
591}
592
593fn take_batch(batch: &RecordBatch, indices: &[usize]) -> DataFusionResult<RecordBatch> {
595 if batch.num_rows() == indices.len() {
596 return Ok(batch.clone());
597 }
598
599 let indices_array = UInt64Array::from_iter(indices.iter().map(|index| *index as u64));
600 let arrays = batch
601 .columns()
602 .iter()
603 .map(|array| compute::take(array, &indices_array, None))
604 .collect::<std::result::Result<Vec<_>, _>>()
605 .map_err(|error| DataFusionError::ArrowError(Box::new(error), None))?;
606
607 RecordBatch::try_new(batch.schema(), arrays)
608 .map_err(|error| DataFusionError::ArrowError(Box::new(error), None))
609}
610
611#[cfg(test)]
612mod test {
613 use std::collections::VecDeque;
614 use std::sync::Arc;
615 use std::sync::atomic::{AtomicUsize, Ordering};
616 use std::task::{Wake, Waker};
617
618 use datafusion::arrow::array::{Array, Float64Array, Int32Array, Int64Array, StringArray};
619 use datafusion::arrow::datatypes::{DataType, Field, Schema, SchemaRef};
620 use datafusion::common::ToDFSchema;
621 use datafusion::datasource::memory::MemorySourceConfig;
622 use datafusion::datasource::source::DataSourceExec;
623 use datafusion::logical_expr::{EmptyRelation, LogicalPlan};
624 use datafusion::prelude::SessionContext;
625 use futures::StreamExt;
626
627 use super::*;
628
629 #[test]
630 fn pruning_should_keep_all_columns_for_exec() {
631 let schema = Arc::new(Schema::new(vec![
632 Field::new("ts", DataType::Int32, false),
633 Field::new("k", DataType::Int32, false),
634 Field::new("v", DataType::Int32, false),
635 ]));
636 let df_schema = schema.to_dfschema_ref().unwrap();
637 let left = LogicalPlan::EmptyRelation(EmptyRelation {
638 produce_one_row: false,
639 schema: df_schema.clone(),
640 });
641 let right = LogicalPlan::EmptyRelation(EmptyRelation {
642 produce_one_row: false,
643 schema: df_schema.clone(),
644 });
645 let plan = UnionDistinctOn::try_new(left, right, vec![1], 0).unwrap();
646
647 let output_columns = [2usize];
649 let required = plan.necessary_children_exprs(&output_columns).unwrap();
650 assert_eq!(required.len(), 2);
651 assert_eq!(required[0].as_slice(), &[0, 1, 2]);
652 assert_eq!(required[1].as_slice(), &[0, 1, 2]);
653 }
654
655 #[test]
656 fn test_take_batch() {
657 let schema = Schema::new(vec![
658 Field::new("a", DataType::Int32, false),
659 Field::new("b", DataType::Int32, false),
660 ]);
661
662 let batch = RecordBatch::try_new(
663 Arc::new(schema.clone()),
664 vec![
665 Arc::new(Int32Array::from(vec![1, 2, 3])),
666 Arc::new(Int32Array::from(vec![4, 5, 6])),
667 ],
668 )
669 .unwrap();
670
671 let indices = vec![0, 2];
672 let result = take_batch(&batch, &indices).unwrap();
673
674 let expected = RecordBatch::try_new(
675 Arc::new(schema),
676 vec![
677 Arc::new(Int32Array::from(vec![1, 3])),
678 Arc::new(Int32Array::from(vec![4, 6])),
679 ],
680 )
681 .unwrap();
682
683 assert_eq!(result, expected);
684 }
685
686 fn empty_plan(schema: datafusion::common::DFSchemaRef) -> LogicalPlan {
687 LogicalPlan::EmptyRelation(EmptyRelation {
688 produce_one_row: false,
689 schema,
690 })
691 }
692
693 fn schemas(left: SchemaRef, right: SchemaRef) -> (LogicalPlan, LogicalPlan) {
694 (
695 empty_plan(left.to_dfschema_ref().unwrap()),
696 empty_plan(right.to_dfschema_ref().unwrap()),
697 )
698 }
699
700 fn schema(prefix: &str, key: &str, nullable: bool) -> SchemaRef {
701 Arc::new(Schema::new(vec![
702 Field::new(format!("{prefix}_ts"), DataType::Int64, false),
703 Field::new(format!("{prefix}_{key}"), DataType::Utf8, nullable),
704 Field::new(format!("{prefix}_value"), DataType::Float64, nullable),
705 ]))
706 }
707
708 fn batch(schema: SchemaRef, ts: i64, key: Option<&str>, value: Option<f64>) -> RecordBatch {
709 RecordBatch::try_new(
710 schema,
711 vec![
712 Arc::new(Int64Array::from(vec![ts])),
713 Arc::new(StringArray::from(vec![key])),
714 Arc::new(Float64Array::from(vec![value])),
715 ],
716 )
717 .unwrap()
718 }
719
720 fn source_exec(batch: RecordBatch) -> Arc<dyn ExecutionPlan> {
721 source_exec_batches(batch.schema(), vec![batch])
722 }
723
724 async fn execute(
725 plan: &UnionDistinctOn,
726 left: RecordBatch,
727 right: RecordBatch,
728 ) -> Vec<RecordBatch> {
729 datafusion::physical_plan::collect(
730 plan.to_execution_plan(source_exec(left), source_exec(right)),
731 SessionContext::default().task_ctx(),
732 )
733 .await
734 .unwrap()
735 }
736
737 fn source_exec_batches(schema: SchemaRef, batches: Vec<RecordBatch>) -> Arc<dyn ExecutionPlan> {
738 Arc::new(DataSourceExec::new(Arc::new(
739 MemorySourceConfig::try_new(&[batches], schema, None).unwrap(),
740 )))
741 }
742
743 async fn execute_batches(
744 plan: &UnionDistinctOn,
745 left_schema: SchemaRef,
746 left: Vec<RecordBatch>,
747 right_schema: SchemaRef,
748 right: Vec<RecordBatch>,
749 ) -> Vec<RecordBatch> {
750 datafusion::physical_plan::collect(
751 plan.to_execution_plan(
752 source_exec_batches(left_schema, left),
753 source_exec_batches(right_schema, right),
754 ),
755 SessionContext::default().task_ctx(),
756 )
757 .await
758 .unwrap()
759 }
760
761 fn simple_schema(prefix: &str, nullable: bool) -> SchemaRef {
762 schema(prefix, "label", nullable)
763 }
764
765 fn simple_batch(schema: SchemaRef, rows: &[(i64, &str, f64)]) -> RecordBatch {
766 RecordBatch::try_new(
767 schema,
768 vec![
769 Arc::new(Int64Array::from_iter_values(rows.iter().map(|row| row.0))),
770 Arc::new(StringArray::from_iter_values(rows.iter().map(|row| row.1))),
771 Arc::new(Float64Array::from_iter_values(rows.iter().map(|row| row.2))),
772 ],
773 )
774 .unwrap()
775 }
776
777 fn simple_rows(batches: &[RecordBatch]) -> Vec<(i64, String, f64)> {
778 batches
779 .iter()
780 .flat_map(|batch| {
781 let timestamps = batch
782 .column(0)
783 .as_any()
784 .downcast_ref::<Int64Array>()
785 .unwrap();
786 let labels = batch
787 .column(1)
788 .as_any()
789 .downcast_ref::<StringArray>()
790 .unwrap();
791 let values = batch
792 .column(2)
793 .as_any()
794 .downcast_ref::<Float64Array>()
795 .unwrap();
796 (0..batch.num_rows()).map(move |row| {
797 (
798 timestamps.value(row),
799 labels.value(row).to_string(),
800 values.value(row),
801 )
802 })
803 })
804 .collect()
805 }
806
807 fn simple_plan(left_schema: SchemaRef, right_schema: SchemaRef) -> UnionDistinctOn {
808 let (left, right) = schemas(left_schema, right_schema);
809 UnionDistinctOn::try_new(left, right, vec![1], 0).unwrap()
810 }
811
812 #[derive(Clone, Debug)]
813 enum TestEvent {
814 Batch(RecordBatch),
815 Error(&'static str),
816 }
817
818 struct TestStream {
819 schema: SchemaRef,
820 events: VecDeque<TestEvent>,
821 polls: Arc<AtomicUsize>,
822 }
823
824 impl Stream for TestStream {
825 type Item = DataFusionResult<RecordBatch>;
826
827 fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
828 self.polls.fetch_add(1, Ordering::SeqCst);
829 match self.events.pop_front() {
830 Some(TestEvent::Batch(batch)) => Poll::Ready(Some(Ok(batch))),
831 Some(TestEvent::Error(message)) => {
832 Poll::Ready(Some(Err(DataFusionError::Internal(message.to_string()))))
833 }
834 None => Poll::Ready(None),
835 }
836 }
837 }
838
839 impl RecordBatchStream for TestStream {
840 fn schema(&self) -> SchemaRef {
841 self.schema.clone()
842 }
843 }
844
845 struct CountingWake(AtomicUsize);
846
847 impl Wake for CountingWake {
848 fn wake(self: Arc<Self>) {
849 self.0.fetch_add(1, Ordering::SeqCst);
850 }
851
852 fn wake_by_ref(self: &Arc<Self>) {
853 self.0.fetch_add(1, Ordering::SeqCst);
854 }
855 }
856
857 fn test_stream(
858 schema: SchemaRef,
859 events: Vec<TestEvent>,
860 polls: Arc<AtomicUsize>,
861 ) -> SendableRecordBatchStream {
862 Box::pin(TestStream {
863 schema,
864 events: events.into(),
865 polls,
866 })
867 }
868
869 #[derive(Debug)]
870 struct TestExec {
871 schema: SchemaRef,
872 events: Vec<TestEvent>,
873 polls: Arc<AtomicUsize>,
874 executions: Arc<AtomicUsize>,
875 execute_error: Option<&'static str>,
876 properties: Arc<PlanProperties>,
877 }
878
879 impl TestExec {
880 fn new(
881 schema: SchemaRef,
882 events: Vec<TestEvent>,
883 polls: Arc<AtomicUsize>,
884 executions: Arc<AtomicUsize>,
885 ) -> Self {
886 Self {
887 properties: Arc::new(PlanProperties::new(
888 EquivalenceProperties::new(schema.clone()),
889 Partitioning::UnknownPartitioning(1),
890 EmissionType::Incremental,
891 Boundedness::Bounded,
892 )),
893 schema,
894 events,
895 polls,
896 executions,
897 execute_error: None,
898 }
899 }
900
901 fn with_execute_error(mut self, error: &'static str) -> Self {
902 self.execute_error = Some(error);
903 self
904 }
905 }
906
907 impl DisplayAs for TestExec {
908 fn fmt_as(&self, _t: DisplayFormatType, _f: &mut std::fmt::Formatter) -> std::fmt::Result {
909 Ok(())
910 }
911 }
912
913 impl ExecutionPlan for TestExec {
914 fn name(&self) -> &str {
915 "TestExec"
916 }
917
918 fn as_any(&self) -> &dyn Any {
919 self
920 }
921
922 fn properties(&self) -> &Arc<PlanProperties> {
923 &self.properties
924 }
925
926 fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
927 vec![]
928 }
929
930 fn with_new_children(
931 self: Arc<Self>,
932 _children: Vec<Arc<dyn ExecutionPlan>>,
933 ) -> DataFusionResult<Arc<dyn ExecutionPlan>> {
934 Ok(self)
935 }
936
937 fn execute(
938 &self,
939 _partition: usize,
940 _context: Arc<TaskContext>,
941 ) -> DataFusionResult<SendableRecordBatchStream> {
942 self.executions.fetch_add(1, Ordering::SeqCst);
943 if let Some(error) = self.execute_error {
944 return Err(DataFusionError::Internal(error.to_string()));
945 }
946 Ok(test_stream(
947 self.schema.clone(),
948 self.events.clone(),
949 self.polls.clone(),
950 ))
951 }
952 }
953
954 fn test_union_stream(
955 left: SendableRecordBatchStream,
956 right: Arc<dyn ExecutionPlan>,
957 output_schema: SchemaRef,
958 ) -> UnionDistinctOnStream {
959 let metrics = ExecutionPlanMetricsSet::new();
960 UnionDistinctOnStream {
961 left,
962 right_plan: right,
963 right_partition: 0,
964 right_context: SessionContext::default().task_ctx(),
965 right: None,
966 compare_keys: vec![1, 0],
967 output_schema,
968 random_state: RandomState::new(),
969 lhs_signatures: HashSet::default(),
970 hashes: Vec::new(),
971 phase: StreamPhase::Left,
972 metric: BaselineMetrics::new(&metrics, 0),
973 }
974 }
975
976 fn series_rows(batches: &[RecordBatch]) -> Vec<(i64, &str, f64, &str)> {
977 batches
978 .iter()
979 .flat_map(|batch| {
980 let timestamps = batch
981 .column(0)
982 .as_any()
983 .downcast_ref::<Int64Array>()
984 .unwrap();
985 let series = batch
986 .column(1)
987 .as_any()
988 .downcast_ref::<StringArray>()
989 .unwrap();
990 let values = batch
991 .column(2)
992 .as_any()
993 .downcast_ref::<Float64Array>()
994 .unwrap();
995 let keys = batch
996 .column(3)
997 .as_any()
998 .downcast_ref::<StringArray>()
999 .unwrap();
1000 (0..batch.num_rows()).map(move |row| {
1001 (
1002 timestamps.value(row),
1003 series.value(row),
1004 values.value(row),
1005 keys.value(row),
1006 )
1007 })
1008 })
1009 .collect()
1010 }
1011
1012 #[tokio::test]
1013 async fn serialize_deserialize_and_execute_with_different_input_names() {
1014 let (left_schema, right_schema) =
1015 (schema("left", "job", false), schema("right", "job", false));
1016 let (left_plan, right_plan) = schemas(left_schema.clone(), right_schema.clone());
1017 let decoded = UnionDistinctOn::deserialize(
1018 &UnionDistinctOn::try_new(left_plan.clone(), right_plan.clone(), vec![1], 0)
1019 .unwrap()
1020 .serialize(),
1021 )
1022 .unwrap()
1023 .with_exprs_and_inputs(vec![], vec![left_plan, right_plan])
1024 .unwrap();
1025 assert_eq!(
1026 (decoded.compare_key_indices.as_slice(), decoded.ts_col_idx),
1027 (&[1usize][..], 0)
1028 );
1029 assert_eq!(
1030 decoded.output_schema,
1031 left_schema.clone().to_dfschema_ref().unwrap()
1032 );
1033 let result = execute(
1034 &decoded,
1035 batch(left_schema.clone(), 1, Some("left"), Some(10.0)),
1036 batch(right_schema, 2, Some("right"), Some(20.0)),
1037 )
1038 .await;
1039 assert!(result.iter().all(|batch| batch.schema() == left_schema));
1040 }
1041
1042 #[tokio::test]
1043 async fn execute_widens_nullable_fields_and_emits_rhs_nulls() {
1044 let (left_schema, right_schema) = (
1045 schema("left", "label", false),
1046 schema("right", "label", true),
1047 );
1048 let (left_plan, right_plan) = schemas(left_schema.clone(), right_schema.clone());
1049 let plan = UnionDistinctOn::try_new(left_plan, right_plan, vec![1, 2], 0).unwrap();
1050 let declared_schema = plan.output_schema.inner().clone();
1051 assert!(
1052 declared_schema.fields()[1..]
1053 .iter()
1054 .all(|field| field.is_nullable())
1055 );
1056 let result = execute(
1057 &plan,
1058 batch(left_schema, 1, Some("present"), Some(10.0)),
1059 batch(right_schema, 2, None, None),
1060 )
1061 .await;
1062 assert_eq!(result.len(), 2);
1063 assert!(result.iter().all(|batch| batch.schema() == declared_schema));
1064 assert!(result[1].column(1).is_null(0));
1065 assert!(result[1].column(2).is_null(0));
1066 }
1067
1068 #[tokio::test]
1069 async fn empty_lhs_preserves_distinct_rhs_series_with_same_normalized_key() {
1070 let left_schema = Arc::new(Schema::new(vec![
1071 Field::new("ts", DataType::Int64, false),
1072 Field::new("series", DataType::Utf8, false),
1073 Field::new("value", DataType::Float64, false),
1074 Field::new("__normalized_absent_label", DataType::Utf8, false),
1075 ]));
1076 let right_schema = Arc::new(Schema::new(vec![
1077 Field::new("rhs_ts", DataType::Int64, false),
1078 Field::new("rhs_series", DataType::Utf8, false),
1079 Field::new("rhs_value", DataType::Float64, false),
1080 Field::new("rhs_normalized_absent_label", DataType::Utf8, false),
1081 ]));
1082 let (left, right) = schemas(left_schema.clone(), right_schema.clone());
1083 let plan = UnionDistinctOn::try_new(left, right, vec![3], 0).unwrap();
1084 let declared_schema = plan.output_schema.inner().clone();
1085 let rhs = RecordBatch::try_new(
1086 right_schema,
1087 vec![
1088 Arc::new(Int64Array::from(vec![1_000, 1_000])),
1089 Arc::new(StringArray::from(vec!["series_a", "series_b"])),
1090 Arc::new(Float64Array::from(vec![10.0, 20.0])),
1091 Arc::new(StringArray::from(vec!["", ""])),
1092 ],
1093 )
1094 .unwrap();
1095
1096 let result = execute(&plan, RecordBatch::new_empty(left_schema), rhs).await;
1097 assert!(result.iter().all(|batch| batch.schema() == declared_schema));
1098 assert_eq!(result.iter().map(RecordBatch::num_rows).sum::<usize>(), 2);
1099
1100 let mut rows = series_rows(&result);
1101 rows.sort_by_key(|row| row.1);
1102 assert_eq!(
1103 rows,
1104 vec![(1_000, "series_a", 10.0, ""), (1_000, "series_b", 20.0, "")]
1105 );
1106 }
1107
1108 #[tokio::test]
1109 async fn lhs_signature_suppresses_all_matching_rhs_rows_and_preserves_lhs() {
1110 let left_schema = Arc::new(Schema::new(vec![
1111 Field::new("ts", DataType::Int64, false),
1112 Field::new("series", DataType::Utf8, false),
1113 Field::new("value", DataType::Float64, false),
1114 Field::new("__normalized_absent_label", DataType::Utf8, false),
1115 ]));
1116 let right_schema = Arc::new(Schema::new(vec![
1117 Field::new("rhs_ts", DataType::Int64, false),
1118 Field::new("rhs_series", DataType::Utf8, false),
1119 Field::new("rhs_value", DataType::Float64, false),
1120 Field::new("rhs_normalized_absent_label", DataType::Utf8, false),
1121 ]));
1122 let (left, right) = schemas(left_schema.clone(), right_schema.clone());
1123 let plan = UnionDistinctOn::try_new(left, right, vec![3], 0).unwrap();
1124 let declared_schema = plan.output_schema.inner().clone();
1125 let lhs = RecordBatch::try_new(
1126 left_schema,
1127 vec![
1128 Arc::new(Int64Array::from(vec![1_000, 2_000])),
1129 Arc::new(StringArray::from(vec!["lhs_match", "lhs_only"])),
1130 Arc::new(Float64Array::from(vec![1.0, 2.0])),
1131 Arc::new(StringArray::from(vec!["", "left-only"])),
1132 ],
1133 )
1134 .unwrap();
1135 let rhs = RecordBatch::try_new(
1136 right_schema,
1137 vec![
1138 Arc::new(Int64Array::from(vec![1_000, 1_000, 3_000])),
1139 Arc::new(StringArray::from(vec![
1140 "rhs_match_a",
1141 "rhs_match_b",
1142 "rhs_unmatched",
1143 ])),
1144 Arc::new(Float64Array::from(vec![10.0, 20.0, 30.0])),
1145 Arc::new(StringArray::from(vec!["", "", ""])),
1146 ],
1147 )
1148 .unwrap();
1149
1150 let result = execute(&plan, lhs, rhs).await;
1151 assert_eq!(result.iter().map(RecordBatch::num_rows).sum::<usize>(), 3);
1152 assert!(result.iter().all(|batch| batch.schema() == declared_schema));
1153 let mut rows = series_rows(&result);
1154 rows.sort_by_key(|row| row.1);
1155 assert_eq!(
1156 rows,
1157 vec![
1158 (1_000, "lhs_match", 1.0, ""),
1159 (2_000, "lhs_only", 2.0, "left-only"),
1160 (3_000, "rhs_unmatched", 30.0, ""),
1161 ]
1162 );
1163 }
1164
1165 #[tokio::test]
1166 async fn empty_lhs_preserves_duplicate_rhs_rows_across_batches() {
1167 let left_schema = simple_schema("left", false);
1168 let right_schema = simple_schema("right", false);
1169 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1170 let result = execute_batches(
1171 &plan,
1172 left_schema,
1173 vec![],
1174 right_schema.clone(),
1175 vec![
1176 simple_batch(
1177 right_schema.clone(),
1178 &[(1, "same", 10.0), (1, "same", 11.0)],
1179 ),
1180 simple_batch(right_schema, &[(1, "same", 12.0)]),
1181 ],
1182 )
1183 .await;
1184 assert_eq!(
1185 simple_rows(&result),
1186 vec![
1187 (1, "same".to_string(), 10.0),
1188 (1, "same".to_string(), 11.0),
1189 (1, "same".to_string(), 12.0),
1190 ]
1191 );
1192 }
1193
1194 #[tokio::test]
1195 async fn lhs_signature_suppresses_rhs_duplicates_across_batches() {
1196 let left_schema = simple_schema("left", false);
1197 let right_schema = simple_schema("right", false);
1198 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1199 let result = execute_batches(
1200 &plan,
1201 left_schema.clone(),
1202 vec![simple_batch(left_schema, &[(1, "match", 1.0)])],
1203 right_schema.clone(),
1204 vec![
1205 simple_batch(right_schema.clone(), &[(1, "match", 10.0)]),
1206 simple_batch(right_schema, &[(1, "match", 11.0)]),
1207 ],
1208 )
1209 .await;
1210 assert_eq!(simple_rows(&result), vec![(1, "match".to_string(), 1.0)]);
1211 }
1212
1213 #[tokio::test]
1214 async fn hash_scratch_reset_suppresses_null_label_after_same_sized_lhs_batches() {
1215 let left_schema = simple_schema("left", true);
1216 let right_schema = simple_schema("right", true);
1217 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1218 let lhs_poison = RecordBatch::try_new(
1219 left_schema.clone(),
1220 vec![
1221 Arc::new(Int64Array::from(vec![1])),
1222 Arc::new(StringArray::from(vec![Some("poison")])),
1223 Arc::new(Float64Array::from(vec![1.0])),
1224 ],
1225 )
1226 .unwrap();
1227 let lhs_null = RecordBatch::try_new(
1228 left_schema.clone(),
1229 vec![
1230 Arc::new(Int64Array::from(vec![42])),
1231 Arc::new(StringArray::from(vec![None::<&str>])),
1232 Arc::new(Float64Array::from(vec![2.0])),
1233 ],
1234 )
1235 .unwrap();
1236 let rhs_null = RecordBatch::try_new(
1237 right_schema.clone(),
1238 vec![
1239 Arc::new(Int64Array::from(vec![42])),
1240 Arc::new(StringArray::from(vec![None::<&str>])),
1241 Arc::new(Float64Array::from(vec![999.0])),
1242 ],
1243 )
1244 .unwrap();
1245
1246 let result = execute_batches(
1247 &plan,
1248 left_schema,
1249 vec![lhs_poison, lhs_null],
1250 right_schema,
1251 vec![rhs_null],
1252 )
1253 .await;
1254 assert_eq!(result.iter().map(RecordBatch::num_rows).sum::<usize>(), 2);
1255 let values = result
1256 .iter()
1257 .flat_map(|batch| {
1258 batch
1259 .column(2)
1260 .as_any()
1261 .downcast_ref::<Float64Array>()
1262 .unwrap()
1263 .values()
1264 .iter()
1265 .copied()
1266 })
1267 .collect::<Vec<_>>();
1268 assert_eq!(values, vec![1.0, 2.0]);
1269 }
1270
1271 #[tokio::test]
1272 async fn mixed_rhs_duplicates_retain_unmatched_rows_in_order() {
1273 let left_schema = simple_schema("left", false);
1274 let right_schema = simple_schema("right", false);
1275 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1276 let result = execute_batches(
1277 &plan,
1278 left_schema.clone(),
1279 vec![simple_batch(left_schema, &[(1, "match", 1.0)])],
1280 right_schema.clone(),
1281 vec![simple_batch(
1282 right_schema,
1283 &[
1284 (1, "match", 10.0),
1285 (1, "keep", 11.0),
1286 (1, "keep", 12.0),
1287 (1, "match", 13.0),
1288 (1, "later", 14.0),
1289 ],
1290 )],
1291 )
1292 .await;
1293 assert_eq!(
1294 simple_rows(&result),
1295 vec![
1296 (1, "match".to_string(), 1.0),
1297 (1, "keep".to_string(), 11.0),
1298 (1, "keep".to_string(), 12.0),
1299 (1, "later".to_string(), 14.0),
1300 ]
1301 );
1302 }
1303
1304 #[tokio::test]
1305 async fn same_label_at_different_timestamps_is_unmatched() {
1306 let left_schema = simple_schema("left", false);
1307 let right_schema = simple_schema("right", false);
1308 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1309 let result = execute_batches(
1310 &plan,
1311 left_schema.clone(),
1312 vec![simple_batch(left_schema, &[(1, "label", 1.0)])],
1313 right_schema.clone(),
1314 vec![simple_batch(right_schema, &[(2, "label", 2.0)])],
1315 )
1316 .await;
1317 assert_eq!(
1318 simple_rows(&result),
1319 vec![(1, "label".to_string(), 1.0), (2, "label".to_string(), 2.0)]
1320 );
1321 }
1322
1323 #[tokio::test]
1324 async fn empty_input_combinations_complete_cleanly() {
1325 let left_schema = simple_schema("left", false);
1326 let right_schema = simple_schema("right", false);
1327 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1328
1329 let empty_left = execute_batches(
1330 &plan,
1331 left_schema.clone(),
1332 vec![],
1333 right_schema.clone(),
1334 vec![simple_batch(right_schema.clone(), &[(1, "rhs", 1.0)])],
1335 )
1336 .await;
1337 assert_eq!(simple_rows(&empty_left), vec![(1, "rhs".to_string(), 1.0)]);
1338
1339 let empty_right = execute_batches(
1340 &plan,
1341 left_schema.clone(),
1342 vec![simple_batch(left_schema.clone(), &[(1, "lhs", 1.0)])],
1343 right_schema.clone(),
1344 vec![],
1345 )
1346 .await;
1347 assert_eq!(simple_rows(&empty_right), vec![(1, "lhs".to_string(), 1.0)]);
1348
1349 let both_empty = execute_batches(&plan, left_schema, vec![], right_schema, vec![]).await;
1350 assert!(both_empty.is_empty());
1351 }
1352
1353 #[tokio::test]
1354 async fn zero_row_batches_on_either_side_are_handled() {
1355 let left_schema = simple_schema("left", false);
1356 let right_schema = simple_schema("right", false);
1357 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1358
1359 let zero_row_left = execute_batches(
1360 &plan,
1361 left_schema.clone(),
1362 vec![simple_batch(left_schema.clone(), &[])],
1363 right_schema.clone(),
1364 vec![simple_batch(right_schema.clone(), &[(1, "rhs", 1.0)])],
1365 )
1366 .await;
1367 assert_eq!(
1368 simple_rows(&zero_row_left),
1369 vec![(1, "rhs".to_string(), 1.0)]
1370 );
1371
1372 let zero_row_right = execute_batches(
1373 &plan,
1374 left_schema.clone(),
1375 vec![simple_batch(left_schema, &[(1, "lhs", 1.0)])],
1376 right_schema.clone(),
1377 vec![simple_batch(right_schema, &[])],
1378 )
1379 .await;
1380 assert_eq!(
1381 simple_rows(&zero_row_right),
1382 vec![(1, "lhs".to_string(), 1.0)]
1383 );
1384 }
1385
1386 #[tokio::test]
1387 async fn metrics_record_output_rows_and_completion() {
1388 let left_schema = simple_schema("left", false);
1389 let right_schema = simple_schema("right", false);
1390 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1391 let exec = plan.to_execution_plan(
1392 source_exec(simple_batch(left_schema, &[(1, "lhs", 1.0)])),
1393 source_exec(simple_batch(right_schema, &[(2, "rhs", 2.0)])),
1394 );
1395
1396 let output =
1397 datafusion::physical_plan::collect(exec.clone(), SessionContext::default().task_ctx())
1398 .await
1399 .unwrap();
1400 assert_eq!(simple_rows(&output).len(), 2);
1401
1402 let metrics = exec.metrics().unwrap();
1403 assert_eq!(metrics.output_rows(), Some(2));
1404 assert!(metrics.iter().any(|metric| {
1405 metric.value().name() == "end_timestamp" && metric.value().as_usize() > 0
1406 }));
1407 }
1408
1409 #[tokio::test]
1410 async fn output_keeps_lhs_before_rhs_and_input_order() {
1411 let left_schema = simple_schema("left", false);
1412 let right_schema = simple_schema("right", false);
1413 let plan = simple_plan(left_schema.clone(), right_schema.clone());
1414 let result = execute_batches(
1415 &plan,
1416 left_schema.clone(),
1417 vec![
1418 simple_batch(left_schema.clone(), &[(1, "lhs-a", 1.0)]),
1419 simple_batch(left_schema, &[(2, "lhs-b", 2.0)]),
1420 ],
1421 right_schema.clone(),
1422 vec![
1423 simple_batch(right_schema.clone(), &[(3, "rhs-a", 3.0)]),
1424 simple_batch(right_schema, &[(4, "rhs-b", 4.0)]),
1425 ],
1426 )
1427 .await;
1428 assert_eq!(
1429 simple_rows(&result),
1430 vec![
1431 (1, "lhs-a".to_string(), 1.0),
1432 (2, "lhs-b".to_string(), 2.0),
1433 (3, "rhs-a".to_string(), 3.0),
1434 (4, "rhs-b".to_string(), 4.0),
1435 ]
1436 );
1437 }
1438
1439 #[test]
1440 fn fully_filtered_rhs_batch_wakes_before_later_unmatched_batch() {
1441 let schema = simple_schema("stream", false);
1442 let left_polls = Arc::new(AtomicUsize::new(0));
1443 let right_polls = Arc::new(AtomicUsize::new(0));
1444 let right_executions = Arc::new(AtomicUsize::new(0));
1445 let left = test_stream(
1446 schema.clone(),
1447 vec![TestEvent::Batch(simple_batch(
1448 schema.clone(),
1449 &[(1, "match", 1.0)],
1450 ))],
1451 left_polls,
1452 );
1453 let right = Arc::new(TestExec::new(
1454 schema.clone(),
1455 vec![
1456 TestEvent::Batch(simple_batch(schema.clone(), &[(1, "match", 2.0)])),
1457 TestEvent::Batch(simple_batch(schema.clone(), &[(2, "keep", 3.0)])),
1458 ],
1459 right_polls.clone(),
1460 right_executions.clone(),
1461 ));
1462 let mut stream = Box::pin(test_union_stream(left, right, schema));
1463 let wake = Arc::new(CountingWake(AtomicUsize::new(0)));
1464 let waker = Waker::from(wake.clone());
1465 let mut context = Context::from_waker(&waker);
1466
1467 assert!(matches!(
1468 stream.as_mut().poll_next(&mut context),
1469 Poll::Ready(Some(Ok(_)))
1470 ));
1471 assert!(matches!(
1472 stream.as_mut().poll_next(&mut context),
1473 Poll::Pending
1474 ));
1475 assert_eq!(right_executions.load(Ordering::SeqCst), 1);
1476 assert_eq!(right_polls.load(Ordering::SeqCst), 1);
1477 assert_eq!(wake.0.load(Ordering::SeqCst), 1);
1478 assert!(matches!(
1479 stream.as_mut().poll_next(&mut context),
1480 Poll::Ready(Some(Ok(batch)))
1481 if simple_rows(std::slice::from_ref(&batch))
1482 == vec![(2, "keep".to_string(), 3.0)]
1483 ));
1484 }
1485
1486 #[tokio::test]
1487 async fn rhs_is_not_polled_before_lhs_eof_and_drop_preserves_backpressure() {
1488 let schema = simple_schema("stream", false);
1489 let left_polls = Arc::new(AtomicUsize::new(0));
1490 let right_polls = Arc::new(AtomicUsize::new(0));
1491 let left = test_stream(
1492 schema.clone(),
1493 vec![TestEvent::Batch(simple_batch(
1494 schema.clone(),
1495 &[(1, "lhs", 1.0)],
1496 ))],
1497 left_polls.clone(),
1498 );
1499 let right_executions = Arc::new(AtomicUsize::new(0));
1500 let right = Arc::new(TestExec::new(
1501 schema.clone(),
1502 vec![TestEvent::Batch(simple_batch(
1503 schema.clone(),
1504 &[(2, "rhs", 2.0)],
1505 ))],
1506 right_polls.clone(),
1507 right_executions.clone(),
1508 ));
1509 let mut stream = Box::pin(test_union_stream(left, right, schema));
1510
1511 let first = stream.next().await.unwrap().unwrap();
1512 assert_eq!(simple_rows(&[first]), vec![(1, "lhs".to_string(), 1.0)]);
1513 assert_eq!(left_polls.load(Ordering::SeqCst), 1);
1514 assert_eq!(right_polls.load(Ordering::SeqCst), 0);
1515 assert_eq!(right_executions.load(Ordering::SeqCst), 0);
1516 drop(stream);
1517 assert_eq!(right_polls.load(Ordering::SeqCst), 0);
1518 assert_eq!(right_executions.load(Ordering::SeqCst), 0);
1519 }
1520
1521 #[tokio::test]
1522 async fn lhs_and_delayed_rhs_errors_propagate() {
1523 let schema = simple_schema("stream", false);
1524 let polls = Arc::new(AtomicUsize::new(0));
1525 let mut left_error = Box::pin(test_union_stream(
1526 test_stream(
1527 schema.clone(),
1528 vec![TestEvent::Error("left failed")],
1529 polls.clone(),
1530 ),
1531 Arc::new(TestExec::new(
1532 schema.clone(),
1533 vec![],
1534 polls.clone(),
1535 Arc::new(AtomicUsize::new(0)),
1536 )),
1537 schema.clone(),
1538 ));
1539 assert!(left_error.next().await.unwrap().is_err());
1540 assert!(left_error.next().await.is_none());
1541
1542 let mut right_error = Box::pin(test_union_stream(
1543 test_stream(
1544 schema.clone(),
1545 vec![TestEvent::Batch(simple_batch(
1546 schema.clone(),
1547 &[(1, "lhs", 1.0)],
1548 ))],
1549 polls.clone(),
1550 ),
1551 Arc::new(TestExec::new(
1552 schema.clone(),
1553 vec![TestEvent::Error("right failed")],
1554 polls,
1555 Arc::new(AtomicUsize::new(0)),
1556 )),
1557 schema.clone(),
1558 ));
1559 assert!(right_error.next().await.unwrap().is_ok());
1560 assert!(right_error.next().await.unwrap().is_err());
1561 assert!(right_error.next().await.is_none());
1562
1563 let right_execute_error = Arc::new(
1564 TestExec::new(
1565 schema.clone(),
1566 vec![],
1567 Arc::new(AtomicUsize::new(0)),
1568 Arc::new(AtomicUsize::new(0)),
1569 )
1570 .with_execute_error("right execute failed"),
1571 );
1572 let mut right_execute_error = Box::pin(test_union_stream(
1573 test_stream(schema.clone(), vec![], Arc::new(AtomicUsize::new(0))),
1574 right_execute_error,
1575 schema,
1576 ));
1577 assert!(right_execute_error.next().await.unwrap().is_err());
1578 assert!(right_execute_error.next().await.is_none());
1579 }
1580
1581 #[tokio::test]
1582 async fn rhs_partial_and_full_batches_rebind_to_declared_left_schema() {
1583 let left_schema = schema("left", "label", false);
1584 let right_schema = schema("right", "label", true);
1585 let (left, right) = schemas(left_schema.clone(), right_schema.clone());
1586 let plan = UnionDistinctOn::try_new(left, right, vec![1], 0).unwrap();
1587 let declared_schema = plan.output_schema.inner().clone();
1588 let partial_rhs = RecordBatch::try_new(
1589 right_schema.clone(),
1590 vec![
1591 Arc::new(Int64Array::from(vec![1, 2])),
1592 Arc::new(StringArray::from(vec![Some("match"), None])),
1593 Arc::new(Float64Array::from(vec![Some(10.0), None])),
1594 ],
1595 )
1596 .unwrap();
1597 let full_rhs = batch(right_schema.clone(), 3, Some("full"), Some(30.0));
1598 let result = execute_batches(
1599 &plan,
1600 left_schema.clone(),
1601 vec![batch(left_schema, 1, Some("match"), Some(1.0))],
1602 right_schema,
1603 vec![partial_rhs, full_rhs],
1604 )
1605 .await;
1606 assert!(result.iter().all(|batch| batch.schema() == declared_schema));
1607 assert_eq!(result.iter().map(RecordBatch::num_rows).sum::<usize>(), 3);
1608 assert!(result[1].column(1).is_null(0));
1609 }
1610
1611 #[test]
1612 fn malformed_indices_and_incompatible_inputs_fail_before_execution() {
1613 let schema = Arc::new(Schema::new(vec![
1614 Field::new("ts", DataType::Int64, false),
1615 Field::new("job", DataType::Utf8, false),
1616 ]));
1617 let invalid = |compare_key_indices, ts_col_idx| {
1618 let decoded = UnionDistinctOn::deserialize(
1619 &pb::UnionDistinctOn {
1620 compare_key_indices,
1621 ts_col_idx,
1622 }
1623 .encode_to_vec(),
1624 )
1625 .unwrap();
1626 let (left, right) = schemas(schema.clone(), schema.clone());
1627 decoded
1628 .with_exprs_and_inputs(vec![], vec![left, right])
1629 .is_err()
1630 };
1631 assert!(invalid(vec![2], 0));
1632 assert!(invalid(vec![1], 2));
1633 let incompatible_schema = Arc::new(Schema::new(vec![
1634 Field::new("other_ts", DataType::Int64, false),
1635 Field::new("other_job", DataType::Int64, false),
1636 ]));
1637 let (left, right) = schemas(schema, incompatible_schema);
1638 assert!(UnionDistinctOn::try_new(left, right, vec![1], 0).is_err());
1639 }
1640}