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