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