Skip to main content

promql/extension_plan/
union_distinct_on.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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/// A special kind of `UNION`(`OR` in PromQL) operator, for PromQL specific use case.
47///
48/// This operator is similar to `UNION` from SQL, but it only accepts two inputs. The
49/// most different part is that it treat left child and right child differently:
50/// - All columns from left child will be outputted.
51/// - Only check collisions (when not distinct) on the columns specified by `compare_keys`.
52/// - Rows from the right child with the same comparison signature are all preserved.
53/// - If a signature occurs in the left child, all matching rows from the right child are discarded.
54/// - Output preserves each input's batch and row order, with all left rows before right rows.
55/// - The output schema is based on the left child schema, with nullability widened from both inputs.
56///
57/// Execution streams the left child first while retaining its comparison signatures. The
58/// right child is polled only after the left child completes, then rows whose signatures were
59/// observed on the left are omitted.
60#[derive(Debug, PartialEq, Eq, Hash)]
61pub struct UnionDistinctOn {
62    left: LogicalPlan,
63    right: LogicalPlan,
64    /// The columns to compare for equality.
65    /// TIME INDEX is included.
66    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        // Compare fields in order excluding output_schema
244        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    /// Shared deterministic hash state for the hashing algorithm.
353    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    /// Include time index
461    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                    // One fully filtered input batch has been consumed. Yield so a sequence
540                    // of filtered batches cannot monopolize a single downstream poll.
541                    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
601/// Utility function to take ordered rows from a record batch.
602fn 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        // Simulate a parent projection requesting only one output column.
656        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}