use crate::StampOutputPort;
use crate::error::{Error, TypedError};
use crate::node::NodeId;
use otel_arrow_dfe_channel::error::SendError;
use otel_arrow_dfe_config::PortName;
use std::collections::HashMap;
use std::future::Future;
pub trait OutputSend: Clone {
type Data;
fn output_send(
&self,
msg: Self::Data,
) -> impl Future<Output = Result<(), SendError<Self::Data>>>;
fn try_output_send(&self, msg: Self::Data) -> Result<(), SendError<Self::Data>>;
}
impl<T> OutputSend for crate::message::Sender<T> {
type Data = T;
async fn output_send(&self, msg: T) -> Result<(), SendError<T>> {
self.send(msg).await
}
fn try_output_send(&self, msg: T) -> Result<(), SendError<T>> {
self.try_send(msg)
}
}
impl<T> OutputSend for crate::shared::message::SharedSender<T> {
type Data = T;
async fn output_send(&self, msg: T) -> Result<(), SendError<T>> {
self.send(msg).await
}
fn try_output_send(&self, msg: T) -> Result<(), SendError<T>> {
self.try_send(msg)
}
}
#[derive(Clone)]
pub struct OutputRouter<S> {
node_id: NodeId,
ports: HashMap<PortName, (S, u16)>,
default: Option<(PortName, S, u16)>,
}
impl<S: Clone> OutputRouter<S> {
#[must_use]
pub fn new(
node_id: NodeId,
msg_senders: HashMap<PortName, S>,
default_port: Option<PortName>,
) -> Self {
let mut entries: Vec<(PortName, S)> = msg_senders.into_iter().collect();
entries.sort_by(|(a, _), (b, _)| a.cmp(b));
let ports: HashMap<PortName, (S, u16)> = entries
.into_iter()
.enumerate()
.map(|(i, (name, sender))| (name, (sender, i as u16)))
.collect();
let default = if let Some(ref port) = default_port {
ports
.get(port)
.cloned()
.map(|(sender, idx)| (port.clone(), sender, idx))
} else if ports.len() == 1 {
ports
.iter()
.next()
.map(|(name, (sender, idx))| (name.clone(), sender.clone(), *idx))
} else {
None
};
Self {
node_id,
ports,
default,
}
}
#[must_use]
pub fn connected_ports(&self) -> Vec<PortName> {
let mut ports: Vec<_> = self.ports.keys().cloned().collect();
ports.sort();
ports
}
#[must_use]
pub fn default_port(&self) -> Option<PortName> {
self.default.as_ref().map(|(name, _, _)| name.clone())
}
}
impl<S: OutputSend> OutputRouter<S> {
#[inline]
pub async fn send_default(&self, data: S::Data) -> Result<(), TypedError<S::Data>> {
match &self.default {
Some((_, sender, _)) => sender
.output_send(data)
.await
.map_err(TypedError::ChannelSendError),
None => Err(TypedError::Error(Error::NoDefaultOutputPort {
node: self.node_id.clone(),
})),
}
}
#[inline]
pub fn try_send_default(&self, data: S::Data) -> Result<(), TypedError<S::Data>> {
match &self.default {
Some((_, sender, _)) => sender
.try_output_send(data)
.map_err(TypedError::ChannelSendError),
None => Err(TypedError::Error(Error::NoDefaultOutputPort {
node: self.node_id.clone(),
})),
}
}
#[inline]
pub async fn send_to<P: Into<PortName>>(
&self,
port: P,
data: S::Data,
) -> Result<(), TypedError<S::Data>> {
let port_name: PortName = port.into();
match self.ports.get(&port_name) {
Some((sender, _)) => sender
.output_send(data)
.await
.map_err(TypedError::ChannelSendError),
None => Err(TypedError::Error(Error::UnknownOutputPort {
node: self.node_id.clone(),
port: port_name,
})),
}
}
#[inline]
pub fn try_send_to<P: Into<PortName>>(
&self,
port: P,
data: S::Data,
) -> Result<(), TypedError<S::Data>> {
let port_name: PortName = port.into();
match self.ports.get(&port_name) {
Some((sender, _)) => sender
.try_output_send(data)
.map_err(TypedError::ChannelSendError),
None => Err(TypedError::Error(Error::UnknownOutputPort {
node: self.node_id.clone(),
port: port_name,
})),
}
}
}
impl<S: OutputSend> OutputRouter<S>
where
S::Data: StampOutputPort,
{
#[inline]
pub async fn send_default_stamped(&self, mut data: S::Data) -> Result<(), TypedError<S::Data>> {
match &self.default {
Some((_, sender, idx)) => {
data.stamp_output_port_index(self.node_id.index, *idx);
sender
.output_send(data)
.await
.map_err(TypedError::ChannelSendError)
}
None => Err(TypedError::Error(Error::NoDefaultOutputPort {
node: self.node_id.clone(),
})),
}
}
#[inline]
pub fn try_send_default_stamped(&self, mut data: S::Data) -> Result<(), TypedError<S::Data>> {
match &self.default {
Some((_, sender, idx)) => {
data.stamp_output_port_index(self.node_id.index, *idx);
sender
.try_output_send(data)
.map_err(TypedError::ChannelSendError)
}
None => Err(TypedError::Error(Error::NoDefaultOutputPort {
node: self.node_id.clone(),
})),
}
}
#[inline]
pub async fn send_to_stamped<P: Into<PortName>>(
&self,
port: P,
mut data: S::Data,
) -> Result<(), TypedError<S::Data>> {
let port_name: PortName = port.into();
match self.ports.get(&port_name) {
Some((sender, idx)) => {
data.stamp_output_port_index(self.node_id.index, *idx);
sender
.output_send(data)
.await
.map_err(TypedError::ChannelSendError)
}
None => Err(TypedError::Error(Error::UnknownOutputPort {
node: self.node_id.clone(),
port: port_name,
})),
}
}
#[inline]
pub fn try_send_to_stamped<P: Into<PortName>>(
&self,
port: P,
mut data: S::Data,
) -> Result<(), TypedError<S::Data>> {
let port_name: PortName = port.into();
match self.ports.get(&port_name) {
Some((sender, idx)) => {
data.stamp_output_port_index(self.node_id.index, *idx);
sender
.try_output_send(data)
.map_err(TypedError::ChannelSendError)
}
None => Err(TypedError::Error(Error::UnknownOutputPort {
node: self.node_id.clone(),
port: port_name,
})),
}
}
}