use crate::Interests;
use crate::config::ReceiverConfig;
use crate::control::{
Controllable, NodeControlMsg, RuntimeCtrlMsgReceiver, pipeline_completion_msg_channel,
runtime_ctrl_msg_channel,
};
use crate::error::Error;
use crate::local::message::{LocalReceiver, LocalSender};
use crate::message::{Receiver, Sender};
use crate::node::NodeWithPDataSender;
use crate::receiver::ReceiverWrapper;
use crate::shared::message::{SharedReceiver, SharedSender};
use crate::testing::{CtrlMsgCounters, setup_test_runtime};
use otel_arrow_dfe_channel::error::RecvError;
use otel_arrow_dfe_config::authorized_identity_policy::AuthorizedIdentityPolicy;
use otel_arrow_dfe_config::transport_headers_policy::HeaderCapturePolicy;
use otel_arrow_dfe_telemetry::reporter::MetricsReporter;
use serde_json::Value;
use std::fmt::Debug;
use std::future::Future;
use std::marker::PhantomData;
use std::time::{Duration, Instant};
use tokio::task::LocalSet;
use tokio::time::sleep;
pub struct TestContext<PData> {
control_sender: Sender<NodeControlMsg<PData>>,
}
pub struct NotSendValidateContext<PData> {
pdata_receiver: Receiver<PData>,
counters: CtrlMsgCounters,
control_sender: Sender<NodeControlMsg<PData>>,
}
pub struct SendValidateContext<PData> {
pdata_receiver: tokio::sync::mpsc::Receiver<PData>,
counters: CtrlMsgCounters,
}
impl<PData> TestContext<PData> {
pub async fn send_control_msg(&self, msg: NodeControlMsg<PData>) -> Result<(), Error> {
self.control_sender
.send(msg)
.await
.map_err(|e| Error::RuntimeMsgError {
error: e.to_string(),
})
}
pub async fn send_timer_tick(&self) -> Result<(), Error> {
self.send_control_msg(NodeControlMsg::TimerTick {}).await
}
pub async fn send_config(&self, config: Value) -> Result<(), Error> {
self.send_control_msg(NodeControlMsg::Config { config })
.await
}
pub async fn send_shutdown(&self, deadline: Instant, reason: &str) -> Result<(), Error> {
self.send_control_msg(NodeControlMsg::Shutdown {
deadline,
reason: reason.to_owned(),
})
.await
}
pub async fn sleep(&self, duration: Duration) {
sleep(duration).await;
}
}
impl<PData> NotSendValidateContext<PData> {
pub async fn recv(&mut self) -> Result<PData, RecvError> {
self.pdata_receiver.recv().await
}
#[must_use]
pub fn counters(&self) -> CtrlMsgCounters {
self.counters.clone()
}
pub async fn send_control_msg(&self, msg: NodeControlMsg<PData>) -> Result<(), Error> {
self.control_sender
.send(msg)
.await
.map_err(|e| Error::RuntimeMsgError {
error: e.to_string(),
})
}
}
impl<PData> SendValidateContext<PData> {
pub async fn recv(&mut self) -> Result<PData, Error> {
self.pdata_receiver
.recv()
.await
.ok_or(Error::ChannelRecvError(RecvError::Closed))
}
#[must_use]
pub fn counters(&self) -> CtrlMsgCounters {
self.counters.clone()
}
}
pub struct TestRuntime<PData> {
config: ReceiverConfig,
rt: tokio::runtime::Runtime,
local_tasks: LocalSet,
counter: CtrlMsgCounters,
_pd: PhantomData<PData>,
}
pub struct TestPhase<PData> {
rt: tokio::runtime::Runtime,
local_tasks: LocalSet,
control_sender: Sender<NodeControlMsg<PData>>,
receiver: ReceiverWrapper<PData>,
counters: CtrlMsgCounters,
}
pub struct ValidationPhase<PData> {
rt: tokio::runtime::Runtime,
local_tasks: LocalSet,
counters: CtrlMsgCounters,
pdata_receiver: Receiver<PData>,
control_sender: Sender<NodeControlMsg<PData>>,
run_receiver_handle: tokio::task::JoinHandle<()>,
run_test_handle: tokio::task::JoinHandle<()>,
#[allow(unused_variables)]
#[allow(dead_code)]
runtime_ctrl_msg_receiver: RuntimeCtrlMsgReceiver<PData>,
}
impl<PData: Clone + Debug + 'static> TestRuntime<PData> {
#[must_use]
pub fn new() -> Self {
let config = ReceiverConfig::new("test_receiver");
let (rt, local_tasks) = setup_test_runtime();
Self {
config,
rt,
local_tasks,
counter: CtrlMsgCounters::new(),
_pd: PhantomData,
}
}
pub const fn config(&self) -> &ReceiverConfig {
&self.config
}
pub fn counters(&self) -> CtrlMsgCounters {
self.counter.clone()
}
pub fn set_receiver(self, receiver: ReceiverWrapper<PData>) -> TestPhase<PData> {
let control_sender = receiver.control_sender();
TestPhase {
rt: self.rt,
local_tasks: self.local_tasks,
receiver,
control_sender,
counters: self.counter,
}
}
}
impl<PData: Clone + Debug + 'static> Default for TestRuntime<PData> {
fn default() -> Self {
Self::new()
}
}
impl<PData: Debug + 'static> TestPhase<PData> {
#[must_use]
pub fn with_capture_policy(mut self, policy: Option<HeaderCapturePolicy>) -> Self {
self.receiver = self
.receiver
.with_capture_policy(policy.map(|policy| policy.compile(|_| true)));
self
}
#[must_use]
pub fn with_authorized_identity_policy(
mut self,
policy: Option<AuthorizedIdentityPolicy>,
) -> Self {
self.receiver = self.receiver.with_authorized_identity_policy(policy);
self
}
pub fn run_test<F, Fut>(mut self, f: F) -> ValidationPhase<PData>
where
F: FnOnce(TestContext<PData>) -> Fut + 'static,
Fut: Future<Output = ()> + 'static,
{
let (node_id, pdata_sender, pdata_receiver) = match &self.receiver {
ReceiverWrapper::Local {
node_id,
runtime_config,
..
} => {
let (sender, receiver) = otel_arrow_dfe_channel::mpsc::Channel::new(
runtime_config.output_pdata_channel.capacity,
);
(
node_id.clone(),
Sender::Local(LocalSender::mpsc(sender)),
Receiver::Local(LocalReceiver::mpsc(receiver)),
)
}
ReceiverWrapper::Shared {
node_id,
runtime_config,
..
} => {
let (sender, receiver) =
tokio::sync::mpsc::channel(runtime_config.output_pdata_channel.capacity);
(
node_id.clone(),
Sender::Shared(SharedSender::mpsc(sender)),
Receiver::Shared(SharedReceiver::mpsc(receiver)),
)
}
};
let (runtime_ctrl_msg_tx, runtime_ctrl_msg_rx) = runtime_ctrl_msg_channel(10);
let (pipeline_completion_msg_tx, _pipeline_completion_msg_rx) =
pipeline_completion_msg_channel(10);
self.receiver
.set_pdata_sender(node_id, "".into(), pdata_sender)
.expect("Failed to set pdata sender");
let control_sender_for_validation = self.control_sender.clone();
let (_metrics_rx, metrics_reporter) = MetricsReporter::create_new_and_receiver(1);
let final_metrics_reporter = metrics_reporter.clone();
let run_receiver_handle = self.local_tasks.spawn_local(async move {
let terminal_state = self
.receiver
.start(
runtime_ctrl_msg_tx,
pipeline_completion_msg_tx,
metrics_reporter,
Interests::empty(),
super::create_test_pipeline_runtime_services(),
)
.await
.expect("Receiver event loop failed");
for snapshot in terminal_state.into_metrics() {
let _ = final_metrics_reporter.try_report_snapshot(snapshot);
}
});
let control_sender_for_test = self.control_sender.clone();
let context = TestContext {
control_sender: control_sender_for_test,
};
let run_test_handle = self.local_tasks.spawn_local(async move {
f(context).await;
});
ValidationPhase {
rt: self.rt,
local_tasks: self.local_tasks,
counters: self.counters,
pdata_receiver,
control_sender: control_sender_for_validation,
run_receiver_handle,
run_test_handle,
runtime_ctrl_msg_receiver: runtime_ctrl_msg_rx,
}
}
}
impl<PData> ValidationPhase<PData> {
pub fn run_validation<F, Fut, T>(self, future_fn: F) -> T
where
F: FnOnce(NotSendValidateContext<PData>) -> Fut,
Fut: Future<Output = T>,
{
let ValidationPhase {
rt,
local_tasks,
counters,
pdata_receiver,
run_receiver_handle,
run_test_handle,
runtime_ctrl_msg_receiver: _,
control_sender,
} = self;
let context = NotSendValidateContext {
pdata_receiver,
counters,
control_sender,
};
rt.block_on(local_tasks);
rt.block_on(run_receiver_handle)
.expect("Receiver task failed");
rt.block_on(run_test_handle).expect("Test task failed");
rt.block_on(future_fn(context))
}
pub fn run_validation_concurrent<F, Fut, T>(self, future_fn: F) -> T
where
F: FnOnce(NotSendValidateContext<PData>) -> Fut + 'static,
Fut: Future<Output = T> + 'static,
T: 'static,
{
let context = NotSendValidateContext {
pdata_receiver: self.pdata_receiver,
counters: self.counters,
control_sender: self.control_sender,
};
let validation_handle = self.local_tasks.spawn_local(future_fn(context));
self.rt.block_on(self.local_tasks);
self.rt
.block_on(self.run_receiver_handle)
.expect("Receiver task failed");
self.rt
.block_on(self.run_test_handle)
.expect("Test task failed");
self.rt
.block_on(validation_handle)
.expect("Validation task failed")
}
}