use crate::attributes::{
ChannelImplementation, ChannelKind, ChannelMode as ChannelModeAttribute, ChannelType,
};
use crate::channel_metrics::{
ChannelMetricsRegistry, ChannelQueueDepth, ChannelReceiverMetricSets, ChannelSenderMetricSets,
ControlChannelReceiverMetricSets, ControlChannelReceiverMetrics,
ControlChannelSenderFailureMetrics, ControlChannelSenderMetricSets,
ControlChannelSenderMetrics, LocalChannelQueueDepth, SharedChannelQueueDepth,
control_channel_id,
};
use crate::context::{ExtensionContext, PipelineContext};
use crate::entity_context::{EntityTelemetryHandle, current_node_telemetry_handle};
use crate::local::message::{LocalReceiver, LocalSender};
use crate::shared::message::{SharedReceiver, SharedSender};
use otel_arrow_dfe_channel::mpsc;
use otel_arrow_dfe_telemetry::otel_warn;
use otel_arrow_dfe_telemetry::registry::EntityKey;
use std::borrow::Cow;
pub(crate) trait ChannelMode {
const CHANNEL_MODE: ChannelModeAttribute;
const CHANNEL_IMPL: ChannelImplementation;
type ControlSender<T>;
type ControlReceiver<T>;
type InnerSender<T>;
type InnerReceiver<T>;
type QueueDepth: ChannelQueueDepth + Default;
fn try_into_inner_sender<T>(
sender: Self::ControlSender<T>,
) -> Result<Self::InnerSender<T>, Self::ControlSender<T>>;
fn try_into_inner_receiver<T>(
receiver: Self::ControlReceiver<T>,
) -> Result<Self::InnerReceiver<T>, Self::ControlReceiver<T>>;
fn from_inner_sender<T>(sender: Self::InnerSender<T>) -> Self::ControlSender<T>;
fn from_inner_receiver<T>(receiver: Self::InnerReceiver<T>) -> Self::ControlReceiver<T>;
fn attach_sender_metrics<T>(
sender: Self::InnerSender<T>,
channel_metrics: &mut ChannelMetricsRegistry,
metrics: ChannelSenderMetricSets,
queue_depth: Self::QueueDepth,
) -> Self::ControlSender<T>;
fn attach_receiver_metrics<T>(
receiver: Self::InnerReceiver<T>,
channel_metrics: &mut ChannelMetricsRegistry,
metrics: ChannelReceiverMetricSets,
capacity: u64,
queue_depth: Self::QueueDepth,
) -> Self::ControlReceiver<T>;
}
pub(crate) struct LocalMode;
pub(crate) struct SharedMode;
impl ChannelMode for LocalMode {
const CHANNEL_MODE: ChannelModeAttribute = ChannelModeAttribute::Local;
const CHANNEL_IMPL: ChannelImplementation = ChannelImplementation::Internal;
type ControlSender<T> = LocalSender<T>;
type ControlReceiver<T> = LocalReceiver<T>;
type InnerSender<T> = mpsc::Sender<T>;
type InnerReceiver<T> = mpsc::Receiver<T>;
type QueueDepth = LocalChannelQueueDepth;
fn try_into_inner_sender<T>(
sender: Self::ControlSender<T>,
) -> Result<Self::InnerSender<T>, Self::ControlSender<T>> {
sender.into_mpsc()
}
fn try_into_inner_receiver<T>(
receiver: Self::ControlReceiver<T>,
) -> Result<Self::InnerReceiver<T>, Self::ControlReceiver<T>> {
receiver.into_mpsc()
}
fn from_inner_sender<T>(sender: Self::InnerSender<T>) -> Self::ControlSender<T> {
LocalSender::mpsc(sender)
}
fn from_inner_receiver<T>(receiver: Self::InnerReceiver<T>) -> Self::ControlReceiver<T> {
LocalReceiver::mpsc(receiver)
}
fn attach_sender_metrics<T>(
sender: Self::InnerSender<T>,
channel_metrics: &mut ChannelMetricsRegistry,
metrics: ChannelSenderMetricSets,
queue_depth: Self::QueueDepth,
) -> Self::ControlSender<T> {
LocalSender::mpsc_with_metrics(sender, channel_metrics, metrics, queue_depth, None)
}
fn attach_receiver_metrics<T>(
receiver: Self::InnerReceiver<T>,
channel_metrics: &mut ChannelMetricsRegistry,
metrics: ChannelReceiverMetricSets,
capacity: u64,
queue_depth: Self::QueueDepth,
) -> Self::ControlReceiver<T> {
LocalReceiver::mpsc_with_metrics(
receiver,
channel_metrics,
metrics,
capacity,
queue_depth,
None,
)
}
}
impl ChannelMode for SharedMode {
const CHANNEL_MODE: ChannelModeAttribute = ChannelModeAttribute::Shared;
const CHANNEL_IMPL: ChannelImplementation = ChannelImplementation::Tokio;
type ControlSender<T> = SharedSender<T>;
type ControlReceiver<T> = SharedReceiver<T>;
type InnerSender<T> = tokio::sync::mpsc::Sender<T>;
type InnerReceiver<T> = tokio::sync::mpsc::Receiver<T>;
type QueueDepth = SharedChannelQueueDepth;
fn try_into_inner_sender<T>(
sender: Self::ControlSender<T>,
) -> Result<Self::InnerSender<T>, Self::ControlSender<T>> {
sender.into_mpsc()
}
fn try_into_inner_receiver<T>(
receiver: Self::ControlReceiver<T>,
) -> Result<Self::InnerReceiver<T>, Self::ControlReceiver<T>> {
receiver.into_mpsc()
}
fn from_inner_sender<T>(sender: Self::InnerSender<T>) -> Self::ControlSender<T> {
SharedSender::mpsc(sender)
}
fn from_inner_receiver<T>(receiver: Self::InnerReceiver<T>) -> Self::ControlReceiver<T> {
SharedReceiver::mpsc(receiver)
}
fn attach_sender_metrics<T>(
sender: Self::InnerSender<T>,
channel_metrics: &mut ChannelMetricsRegistry,
metrics: ChannelSenderMetricSets,
queue_depth: Self::QueueDepth,
) -> Self::ControlSender<T> {
SharedSender::mpsc_with_metrics(sender, channel_metrics, metrics, queue_depth, None)
}
fn attach_receiver_metrics<T>(
receiver: Self::InnerReceiver<T>,
channel_metrics: &mut ChannelMetricsRegistry,
metrics: ChannelReceiverMetricSets,
capacity: u64,
queue_depth: Self::QueueDepth,
) -> Self::ControlReceiver<T> {
SharedReceiver::mpsc_with_metrics(
receiver,
channel_metrics,
metrics,
capacity,
queue_depth,
None,
)
}
}
pub(crate) fn wrap_node_control_channel_metrics<M, Msg>(
name: &str,
pipeline_ctx: &PipelineContext,
channel_metrics: &mut ChannelMetricsRegistry,
channel_metrics_enabled: bool,
capacity: u64,
control_sender: M::ControlSender<Msg>,
control_receiver: M::ControlReceiver<Msg>,
) -> (M::ControlSender<Msg>, M::ControlReceiver<Msg>)
where
M: ChannelMode,
{
wrap_control_channel_metrics_inner::<M, Msg>(
channel_metrics,
channel_metrics_enabled,
capacity,
control_sender,
control_receiver,
|| {
let key = pipeline_ctx.register_node_channel_entity(
control_channel_id(name),
"input".into(),
ChannelKind::Control,
M::CHANNEL_MODE,
ChannelType::Mpsc,
M::CHANNEL_IMPL,
);
if let Some(telemetry) = current_node_telemetry_handle() {
telemetry.set_control_channel_key(key);
}
key
},
|channel_entity_key| {
(
ControlChannelSenderMetricSets {
messages: pipeline_ctx
.register_measurement_metric_set_for_entity::<ControlChannelSenderMetrics>(
channel_entity_key,
),
failures: pipeline_ctx.register_measurement_metric_set_for_entity::<
ControlChannelSenderFailureMetrics,
>(channel_entity_key),
},
ControlChannelReceiverMetricSets {
metrics: pipeline_ctx.register_metric_set_for_entity::<
ControlChannelReceiverMetrics,
>(channel_entity_key),
},
)
},
)
}
pub(crate) fn wrap_extension_control_channel_metrics<M, Msg>(
extension_id: Cow<'static, str>,
variant: crate::extension::wrapper::ExtensionVariant,
entity_handle: &EntityTelemetryHandle,
ext_ctx: &ExtensionContext,
channel_metrics: &mut ChannelMetricsRegistry,
channel_metrics_enabled: bool,
capacity: u64,
control_sender: M::ControlSender<Msg>,
control_receiver: M::ControlReceiver<Msg>,
) -> (M::ControlSender<Msg>, M::ControlReceiver<Msg>)
where
M: ChannelMode,
{
wrap_control_channel_metrics_inner::<M, Msg>(
channel_metrics,
channel_metrics_enabled,
capacity,
control_sender,
control_receiver,
|| {
let key = ext_ctx.register_extension_channel_entity(
extension_id.clone(),
variant,
control_channel_id(extension_id.as_ref()),
M::CHANNEL_MODE,
M::CHANNEL_IMPL,
);
entity_handle.track_entity(key);
key
},
|channel_entity_key| {
(
ControlChannelSenderMetricSets {
messages: entity_handle.register_measurement_metric_set_for_entity::<
ControlChannelSenderMetrics,
>(channel_entity_key),
failures: entity_handle.register_measurement_metric_set_for_entity::<
ControlChannelSenderFailureMetrics,
>(channel_entity_key),
},
ControlChannelReceiverMetricSets {
metrics: entity_handle.register_metric_set_for_entity::<
ControlChannelReceiverMetrics,
>(channel_entity_key),
},
)
},
)
}
fn wrap_control_channel_metrics_inner<M, Msg>(
channel_metrics: &mut ChannelMetricsRegistry,
channel_metrics_enabled: bool,
capacity: u64,
control_sender: M::ControlSender<Msg>,
control_receiver: M::ControlReceiver<Msg>,
register_channel: impl FnOnce() -> EntityKey,
register_metrics: impl FnOnce(
EntityKey,
) -> (
ControlChannelSenderMetricSets,
ControlChannelReceiverMetricSets,
),
) -> (M::ControlSender<Msg>, M::ControlReceiver<Msg>)
where
M: ChannelMode,
{
let control_sender = M::try_into_inner_sender(control_sender);
let control_receiver = M::try_into_inner_receiver(control_receiver);
match (control_sender, control_receiver) {
(Ok(sender), Ok(receiver)) => {
let channel_entity_key = register_channel();
if channel_metrics_enabled {
let (sender_metrics, receiver_metrics) = register_metrics(channel_entity_key);
let queue_depth = M::QueueDepth::default();
(
M::attach_sender_metrics(
sender,
channel_metrics,
ChannelSenderMetricSets::Control(sender_metrics),
queue_depth.clone(),
),
M::attach_receiver_metrics(
receiver,
channel_metrics,
ChannelReceiverMetricSets::Control(receiver_metrics),
capacity,
queue_depth,
),
)
} else {
(
M::from_inner_sender(sender),
M::from_inner_receiver(receiver),
)
}
}
(sender, receiver) => {
let sender_was_inner = sender.is_ok();
let receiver_was_inner = receiver.is_ok();
debug_assert!(
false,
"wrap_control_channel_metrics_inner: mismatched partial_wrap inputs \
(sender_was_inner={sender_was_inner}, receiver_was_inner={receiver_was_inner}); \
channel sender and receiver must originate from the same factory",
);
otel_warn!(
"channel.metrics.partial_wrap_skip",
sender_was_inner = sender_was_inner,
receiver_was_inner = receiver_was_inner,
message = "mismatched halves; skipping metric registration",
);
let sender = match sender {
Ok(sender) => M::from_inner_sender(sender),
Err(sender) => sender,
};
let receiver = match receiver {
Ok(receiver) => M::from_inner_receiver(receiver),
Err(receiver) => receiver,
};
(sender, receiver)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use otel_arrow_dfe_channel::mpmc as raw_mpmc;
use otel_arrow_dfe_channel::mpsc as raw_mpsc;
use std::cell::Cell;
use std::num::NonZeroUsize;
use std::rc::Rc;
fn mismatched_local_pair<T: 'static>() -> (LocalSender<T>, LocalReceiver<T>) {
let (mpsc_sender, _mpsc_receiver) = raw_mpsc::Channel::<T>::new(4);
let (_mpmc_sender, mpmc_receiver) =
raw_mpmc::Channel::<T>::new(NonZeroUsize::new(4).expect("non-zero capacity"));
(
LocalSender::mpsc(mpsc_sender),
LocalReceiver::mpmc(mpmc_receiver),
)
}
#[test]
#[should_panic(expected = "partial_wrap")]
fn partial_wrap_panics_in_debug_builds() {
let (sender, receiver) = mismatched_local_pair::<u8>();
let mut channel_metrics = ChannelMetricsRegistry::default();
let _ = wrap_control_channel_metrics_inner::<LocalMode, u8>(
&mut channel_metrics,
true,
4,
sender,
receiver,
|| panic!("register_channel must not be invoked on the partial-wrap path"),
|_key| panic!("register_metrics must not be invoked on the partial-wrap path"),
);
}
#[test]
fn partial_wrap_does_not_register_or_attach_metrics() {
use std::panic::{AssertUnwindSafe, catch_unwind};
let register_invocations: Rc<Cell<u32>> = Rc::new(Cell::new(0));
let metrics_invocations: Rc<Cell<u32>> = Rc::new(Cell::new(0));
let r_invocations = Rc::clone(®ister_invocations);
let m_invocations = Rc::clone(&metrics_invocations);
let mut channel_metrics = ChannelMetricsRegistry::default();
let (sender, receiver) = mismatched_local_pair::<u8>();
let prev_hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let _ = catch_unwind(AssertUnwindSafe(|| {
let _ = wrap_control_channel_metrics_inner::<LocalMode, u8>(
&mut channel_metrics,
true,
4,
sender,
receiver,
move || {
r_invocations.set(r_invocations.get() + 1);
EntityKey::default()
},
move |_key| {
m_invocations.set(m_invocations.get() + 1);
unreachable!("metrics registration must not run when partial-wrap is detected");
},
);
}));
std::panic::set_hook(prev_hook);
assert_eq!(
register_invocations.get(),
0,
"partial-wrap must not register a channel entity",
);
assert_eq!(
metrics_invocations.get(),
0,
"partial-wrap must not attach metric handles",
);
assert!(
channel_metrics.into_handles().is_empty(),
"channel metrics registry must remain empty after a partial-wrap call",
);
}
}