use std::sync::Arc;
use std::time::Duration;
use tokio::spawn;
use tokio::sync::{mpsc, oneshot};
use tokio::time::Instant;
use super::{
msg::{InternalMsg, Msg, RequestMsg, ResponseMsg, Tvf},
service::{ServiceError, ServiceTable},
};
pub struct Apn<M>
where
M: Sized + Clone + Tvf,
{
service_table: Arc<ServiceTable<M>>,
timeout: Duration,
trace_id: Option<tracing::span::Id>,
}
impl<M> std::fmt::Debug for Apn<M>
where
M: Sized + Clone + Tvf,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Apn")
.field("timeout", &self.timeout)
.field("trace_id", &self.trace_id)
.finish()
}
}
impl<M> Apn<M>
where
M: Sized + Clone + Tvf,
{
pub(crate) fn new(
service_table: Arc<ServiceTable<M>>,
timeout: Duration,
trace_id: Option<tracing::span::Id>,
) -> Apn<M> {
Apn {
service_table,
timeout,
trace_id,
}
}
pub fn trace_id(&self) -> Option<&tracing::span::Id> {
self.trace_id.as_ref()
}
pub async fn call(&self, service_name: &str, data: M) -> Result<ResponseMsg<M>, ServiceError> {
let Some(proc_queue) = self
.service_table
.get_proc_service(service_name)
.map(|proc_service| proc_service.proc_queue.clone())
else {
return Err(ServiceError::UnableToReachService(service_name.to_string()));
};
self.dispatch(&proc_queue, service_name, data, None).await
}
pub async fn call_with_timeout(
&self,
service_name: &str,
data: M,
timeout: Duration,
) -> Result<ResponseMsg<M>, ServiceError> {
let Some(proc_queue) = self
.service_table
.get_proc_service(service_name)
.map(|proc_service| proc_service.proc_queue.clone())
else {
return Err(ServiceError::UnableToReachService(service_name.to_string()));
};
self.dispatch(&proc_queue, service_name, data, Some(timeout))
.await
}
async fn dispatch(
&self,
proc_queue: &mpsc::Sender<InternalMsg<M>>,
service_name: &str,
data: M,
timeout: Option<Duration>,
) -> Result<ResponseMsg<M>, ServiceError> {
let (response_tx, response_rx) = oneshot::channel();
let request = match &self.trace_id {
Some(trace_id) => RequestMsg::new_with_trace_id(
service_name.to_string(),
data,
response_tx,
trace_id.clone(),
),
None => RequestMsg::new(service_name.to_string(), data, response_tx),
};
if proc_queue
.send(InternalMsg::Request(request))
.await
.is_err()
{
return Err(ServiceError::UnableToReachService(service_name.to_string()));
}
if let Some(timeout) = timeout {
match tokio::time::timeout(timeout, response_rx).await {
Ok(Ok(InternalMsg::Response(resp))) => Ok(resp),
Ok(Ok(InternalMsg::Error(err))) => Err(err.into_err()),
Ok(Ok(_)) => Err(ServiceError::ProtocolError(service_name.to_string())),
Ok(Err(_recv)) => Err(ServiceError::UnableToReachService(service_name.to_string())),
Err(_elapsed) => Err(ServiceError::Timeout(
service_name.to_string(),
timeout.as_millis() as u64,
)),
}
} else {
match response_rx.await {
Ok(InternalMsg::Response(resp)) => Ok(resp),
Ok(InternalMsg::Error(err)) => Err(err.into_err()),
Ok(_) => Err(ServiceError::ProtocolError(service_name.to_string())),
Err(_recv) => Err(ServiceError::UnableToReachService(service_name.to_string())),
}
}
}
}
impl<M> RequestMsg<M>
where
M: Sized
+ Clone
+ std::fmt::Debug
+ Tvf
+ Default
+ 'static
+ std::marker::Send
+ std::marker::Sync,
{
pub fn apn<F, Fut>(
mut self,
service_table: Arc<ServiceTable<M>>,
timeout: Duration,
automaton: F,
) where
F: FnOnce(Apn<M>, String, M) -> Fut + Send + 'static,
Fut: Future<Output = Result<M, ServiceError>> + Send + 'static,
{
let apn = Apn::new(service_table, timeout, self.get_span().id());
let service = self.get_service().clone();
let data = self.take_data().unwrap_or_default();
let deadline = Instant::now() + timeout;
spawn(async move {
let _ = match tokio::time::timeout_at(deadline, automaton(apn, service, data)).await {
Ok(result) => self.return_result_to_sender(result),
Err(_elapsed) => {
let service_name = self.get_service().to_string();
self.return_error_to_sender(
None,
ServiceError::Timeout(service_name, timeout.as_millis() as u64),
)
}
};
});
}
}
#[cfg(test)]
mod tests {
extern crate self as prosa;
use std::sync::Arc;
use std::time::Duration;
use prosa_macros::{proc, settings};
use prosa_utils::msg::{simple_string_tvf::SimpleStringTvf, tvf::Tvf};
use serde::Serialize;
use tokio::sync::mpsc;
use tokio::time::timeout;
use super::Apn;
use crate::core::{
error::BusError,
main::{Main, MainProc, MainRunnable},
msg::{InternalMsg, Msg, RequestMsg},
proc::{ProcBusParam, ProcConfig, ProcParam},
service::{ProcService, ServiceError, ServiceTable},
};
use crate::stub::adaptor::StubParotAdaptor;
use crate::stub::proc::{StubProc, StubSettings};
#[settings]
#[derive(Default, Debug, Serialize)]
struct DummySettings {}
#[tokio::test]
async fn apn_call_unreachable_service() {
let apn: Apn<SimpleStringTvf> = Apn::new(
Arc::new(ServiceTable::default()),
Duration::from_millis(50),
None,
);
let err = apn
.call("NOPE", SimpleStringTvf::default())
.await
.expect_err("unreachable service should error");
assert!(matches!(err, ServiceError::UnableToReachService(_)));
}
#[tokio::test]
async fn apn_call_timeout() {
let (bus, _main): (Main<SimpleStringTvf>, MainProc<SimpleStringTvf>) =
MainProc::create(&DummySettings::default(), None);
let (queue_tx, _queue_rx) = mpsc::channel(8);
let proc_param = ProcParam::new(1, "slow".to_string(), queue_tx.clone(), bus);
let proc_service = ProcService::new(&proc_param, queue_tx, 0);
let mut table = ServiceTable::default();
table.add_service("SLOW", proc_service);
let apn: Apn<SimpleStringTvf> = Apn::new(Arc::new(table), Duration::from_millis(30), None);
let err = apn
.call_with_timeout(
"SLOW",
SimpleStringTvf::default(),
Duration::from_millis(30),
)
.await
.expect_err("slow service should time out");
assert!(matches!(err, ServiceError::Timeout(_, _)));
drop(_queue_rx);
}
#[proc]
struct ApnTestProc {}
#[proc]
impl ApnTestProc<SimpleStringTvf> {
async fn apn_run(&mut self) -> Result<(), BusError> {
self.proc.add_proc().await?;
self.proc
.add_service_proc(vec![String::from("APN")])
.await?;
let mut sent = false;
loop {
if let Some(msg) = self.internal_rx_queue.recv().await {
match msg {
InternalMsg::Service(table) => {
self.service = table;
if !sent
&& self.service.exist_proc_service("SUB1")
&& self.service.exist_proc_service("SUB2")
&& self.service.exist_proc_service("APN")
{
sent = true;
let mut data = SimpleStringTvf::default();
data.put_string(1, "start");
if let Some(service) = self.service.get_proc_service("APN") {
service
.proc_queue
.send(InternalMsg::Request(RequestMsg::new(
String::from("APN"),
data,
self.proc.get_service_queue(),
)))
.await
.expect("APN request should be sent");
}
}
}
InternalMsg::Request(req) => {
req.apn(
self.service.clone(),
Duration::from_millis(500),
move |apn, _service, data| async move {
let mut first = apn.call("SUB1", data).await?;
let first_data = first.take_data().ok_or_else(|| {
ServiceError::ProtocolError("SUB1".to_string())
})?;
let mut resp = apn.call("SUB2", first_data).await?;
resp.take_data().ok_or_else(|| {
ServiceError::ProtocolError("SUB2".to_string())
})
},
);
}
InternalMsg::Response(resp) => {
assert_eq!("start", resp.get_data()?.get_string(1)?.into_owned());
self.proc.remove_proc(None).await?;
return Ok(());
}
InternalMsg::Error(err) => {
return Err(BusError::ProcComm(
self.get_proc_id(),
0,
format!("unexpected APN error: {:?}", err.get_err()),
));
}
_ => {}
}
}
}
}
}
#[tokio::test]
async fn apn_happy_path() {
let (bus, main) = MainProc::<SimpleStringTvf>::create(&DummySettings::default(), Some(2));
let main_task = tokio::spawn(main.run());
let stub_proc = StubProc::<SimpleStringTvf>::create(
1,
String::from("stub"),
bus.clone(),
StubSettings::new(vec![String::from("SUB1"), String::from("SUB2")]),
);
crate::core::proc::Proc::<StubParotAdaptor>::run(stub_proc).expect("stub should run");
let result = timeout(
Duration::from_secs(5),
ApnTestProc::<SimpleStringTvf>::create_raw(2, "apn_test".to_string(), bus.clone())
.apn_run(),
)
.await
.expect("APN test should not time out");
assert_eq!(Ok(()), result);
bus.stop("ProSA unit test end".into())
.await
.expect("ProSA should stop");
main_task.await.expect("Main task should end correctly");
}
}