use std::sync::Arc;
use futures::stream::{self, StreamExt};
use crate::exec::{
AccessMode, CombineAccessModes, ContextLevel, ExecOperator, ExecutionContext, FlowResult,
OperatorMetrics, ValueBatchStream, buffer_stream, monitor_stream,
};
#[derive(Debug, Clone)]
pub struct Union {
pub(crate) inputs: Vec<Arc<dyn ExecOperator>>,
pub(crate) metrics: Arc<OperatorMetrics>,
}
impl Union {
pub(crate) fn new(inputs: Vec<Arc<dyn ExecOperator>>) -> Self {
Self {
inputs,
metrics: Arc::new(OperatorMetrics::new()),
}
}
}
impl ExecOperator for Union {
fn name(&self) -> &'static str {
"Union"
}
fn attrs(&self) -> Vec<(String, String)> {
vec![("inputs".to_string(), self.inputs.len().to_string())]
}
fn required_context(&self) -> ContextLevel {
self.inputs.iter().map(|input| input.required_context()).max().unwrap_or(ContextLevel::Root)
}
fn access_mode(&self) -> AccessMode {
self.inputs.iter().map(|input| input.access_mode()).combine_all()
}
fn children(&self) -> Vec<&Arc<dyn ExecOperator>> {
self.inputs.iter().collect()
}
fn metrics(&self) -> Option<&OperatorMetrics> {
Some(&self.metrics)
}
fn execute(&self, ctx: &ExecutionContext) -> FlowResult<ValueBatchStream> {
if self.inputs.is_empty() {
return Ok(monitor_stream(Box::pin(stream::empty()), "Union", &self.metrics));
}
if self.inputs.len() == 1 {
let stream = buffer_stream(
self.inputs[0].execute(ctx)?,
self.inputs[0].access_mode(),
self.inputs[0].cardinality_hint(),
ctx.root().ctx.config.operator_buffer_size,
);
return Ok(monitor_stream(stream, "Union", &self.metrics));
}
let inputs = self.inputs.clone();
let ctx = ctx.clone();
let combined = stream::unfold(
(inputs, ctx, 0usize, Option::<ValueBatchStream>::None),
|(inputs, ctx, mut idx, mut current)| async move {
loop {
let item = match &mut current {
Some(stream) => stream.next().await,
None => None,
};
if let Some(item) = item {
return Some((item, (inputs, ctx, idx, current)));
}
if idx >= inputs.len() {
return None;
}
let i = idx;
idx += 1;
match inputs[i].execute(&ctx) {
Ok(stream) => {
current = Some(buffer_stream(
stream,
inputs[i].access_mode(),
inputs[i].cardinality_hint(),
ctx.root().ctx.config.operator_buffer_size,
))
}
Err(e) => return Some((Err(e), (inputs, ctx, idx, None))),
}
}
},
);
Ok(monitor_stream(Box::pin(combined), "Union", &self.metrics))
}
}