use crate::flow::{Flow, FlowManager, NodeKind};
use crate::message::Message;
use crate::node::{ChannelOrigin, NodeContext, NodeError, NodeType};
use async_trait::async_trait;
use channel_plugin::message::{ChannelMessage, MessageContent, MessageDirection};
use channel_plugin::plugin::ChannelPlugin;
use dashmap::DashMap;
use schemars::{schema_for, JsonSchema};
use schemars::schema::RootSchema;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tracing::{error, warn};
use super::flow_router::{ChannelFlowRouter, ScriptFlowRouter};
use super::manager::{ChannelManager, IncomingHandler};
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename = "channel")]
pub struct ChannelNode {
pub channel_name: String,
pub flow_name: String,
pub node_id: String,
pub poll_messages: bool,
pub send_messages: bool,
#[serde(rename = "router")]
pub router_config: FlowRouterConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(tag = "type", rename = "router")]
pub enum FlowRouterConfig {
#[serde(rename = "channel")]
Channel(ChannelFlowRouter),
#[serde(rename = "script")]
Script(ScriptFlowRouter),
}
impl ChannelNode {
pub async fn handle_message(&self, msg: &ChannelMessage, fm: &FlowManager) {
let input = Message::new_uuid(&msg.channel, serde_json::to_value(msg).unwrap());
let channel_origin = ChannelOrigin::new(msg.channel.clone(), msg.from.clone());
if let Some(report) = fm
.process_message(&self.flow_name, &self.node_id, input, Some(channel_origin))
.await
{
let payload_json = serde_json::to_string(&report)
.expect("cannot serialize report");
tracing::event!(
target: "request", tracing::Level::INFO,
flow = %self.flow_name,
node = %self.node_id,
report = %payload_json,
"flow run completed"
);
}
}
}
#[derive(Clone)]
pub struct ChannelsRegistry {
map: Arc<DashMap<String, Vec<ChannelNode>>>,
flow_manager: Arc<FlowManager>,
}
impl ChannelsRegistry {
pub async fn new(flow_manager: Arc<FlowManager>, _channel_manager: Arc<ChannelManager>) -> Arc<Self> {
let me = Arc::new(Self {
map: Arc::new(DashMap::new()),
flow_manager: flow_manager.clone(),
});
let registry = me.clone();
flow_manager
.subscribe_flow_added(Arc::new(move |flow_id: &str, flow: &Flow| {
for (node_name, cfg) in flow.nodes().iter() {
if let NodeKind::Channel { cfg} = &cfg.kind {
registry.register(ChannelNode {
channel_name: cfg.channel_name.clone(),
flow_name: flow_id.to_string(),
node_id: node_name.clone(),
poll_messages: cfg.channel_in.clone(),
send_messages: cfg.channel_out.clone(),
router_config: FlowRouterConfig::Channel(ChannelFlowRouter::new()),
});
}
}
}))
.await;
me
}
pub fn subscribe(&self) {
}
pub fn register(&self, node: ChannelNode) {
self.map
.entry(node.channel_name.clone())
.or_default()
.push(node);
}
}
#[async_trait]
impl IncomingHandler for ChannelsRegistry {
async fn handle_incoming(&self, msg: ChannelMessage) {
if let Some(nodes) = self.map.get(&msg.channel) {
if nodes.is_empty() {
error!(
channel = %msg.channel,
"received message but channel has no nodes configured"
);
} else {
for node in nodes.iter().cloned() {
node.handle_message(&msg, &self.flow_manager).await;
}
}
} else {
error!(
channel = %msg.channel,
"received message but no flows bound for this channel"
);
}
}
}
#[typetag::serde]
impl NodeType for ChannelNode {
fn type_name(&self) -> String {
self.channel_name.clone()
}
fn schema(&self) -> RootSchema {
schema_for!(ChannelNode)
}
fn process(&self, input: Message, ctx: &mut NodeContext) -> Result<Message, NodeError> {
let mut plugin = ctx
.channel_manager()
.channel(&self.channel_name)
.ok_or_else(|| NodeError::Internal(format!("no such channel: {}", self.channel_name)))?;
let send_result = if let Ok(mut cm) = serde_json::from_value::<ChannelMessage>(input.payload().clone())
{
cm.channel = self.channel_name.clone();
cm.direction = MessageDirection::Outgoing;
if cm.to.is_empty() {
if let Some(channel_origin) = ctx.channel_origin() {
cm.to = vec![channel_origin.participant()];
} else {
let error = format!("No to field was specified so don't know where to send the message to in channel {} with session id {:?}",cm.channel, input.session_id());
error!(error);
return Err(NodeError::InvalidInput(error));
}
}
plugin.send(cm)
} else {
let text = input.payload().to_string();
let to = if let Some(channel_origin) = ctx.channel_origin() {
vec![channel_origin.participant()]
} else {
let error = format!("No to field was specified so don't know where to send the message to in channel {} with session id {:?}",plugin.name(), input.session_id());
error!(error);
return Err(NodeError::InvalidInput(error));
};
let cm = ChannelMessage {
to: to.clone(),
channel: self.channel_name.clone(),
session_id: input.session_id().clone(),
direction: MessageDirection::Outgoing,
content: Some(MessageContent::Text(text)),
..Default::default()
};
plugin.send(cm)
};
if let Err(e) = send_result {
warn!(error = ?e, "failed to send to channel {}", self.channel_name);
}
Ok(input)
}
fn clone_box(&self) -> Box<dyn NodeType> {
Box::new(self.clone())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::channel::manager::{ChannelManager, HostLogger};
use crate::config::{ConfigManager, MapConfigManager};
use crate::{executor::Executor, flow::FlowManager, logger::OpenTelemetryLogger,
secret::EmptySecretsManager, state::InMemoryState,};
use crate::secret::SecretsManager;
use crate::logger::Logger;
use channel_plugin::message::{ChannelMessage, MessageDirection};
#[tokio::test]
async fn test_registry_dispatches_safely() {
let store = InMemoryState::new();
let secrets = SecretsManager(EmptySecretsManager::new());
let logger = Logger(Box::new(OpenTelemetryLogger::new()));
let exec = Executor::new(secrets.clone(), logger);
let config = ConfigManager(MapConfigManager::new());
let host_logger = HostLogger::new();
let cm = ChannelManager::new(config, secrets.clone(), host_logger).await.expect("could not create channel manager");
let fm = FlowManager::new(store, exec, cm.clone(), secrets);
let reg = ChannelsRegistry::new(fm,cm).await;
let mut msg = ChannelMessage::default();
msg.channel = "foo".into();
msg.direction = MessageDirection::Incoming;
reg.handle_incoming(msg.clone()).await;
let node = ChannelNode {
channel_name: "foo".into(),
flow_name: "flow_x".into(),
node_id: "node_id".into(),
poll_messages: true,
send_messages: false,
router_config: FlowRouterConfig::Channel(ChannelFlowRouter::default()),
};
reg.register(node);
reg.handle_incoming(msg).await;
}
}