use std::future::Future;
use std::pin::Pin;
use crate::error::Result;
use crate::protocol::event_typed::Event;
pub type BoxFuture<T> = Pin<Box<dyn Future<Output = T>>>;
pub trait AsyncHandler<E: Event>: Send + Sync + 'static {
fn call(&self, payload: &E::Payload) -> BoxFuture<Result<()>>;
}
impl<E, F, Fut> AsyncHandler<E> for F
where
E: Event,
F: Fn(&E::Payload) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<()>> + 'static,
{
#[inline]
fn call(&self, p: &E::Payload) -> BoxFuture<Result<()>> {
Box::pin(self(p))
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use flowscope::Timestamp;
use super::*;
use crate::protocol::builtin::Tcp;
use crate::protocol::event_typed::FlowStarted;
fn dummy_flow_started() -> FlowStarted<Tcp> {
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
let key = flowscope::extract::FiveTupleKey {
proto: flowscope::L4Proto::Tcp,
a: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)), 12345),
b: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2)), 80),
};
FlowStarted::<Tcp>::new(key, Some(flowscope::L4Proto::Tcp), Timestamp::new(0, 0))
}
#[tokio::test(flavor = "current_thread")]
async fn async_closure_awaits_to_completion() {
let counter = Arc::new(AtomicU32::new(0));
let c = Arc::clone(&counter);
let handler = move |_p: &FlowStarted<Tcp>| {
let c = Arc::clone(&c);
async move {
tokio::task::yield_now().await;
c.fetch_add(1, Ordering::Relaxed);
Ok(())
}
};
let evt = dummy_flow_started();
AsyncHandler::<FlowStarted<Tcp>>::call(&handler, &evt)
.await
.unwrap();
assert_eq!(counter.load(Ordering::Relaxed), 1);
}
#[tokio::test(flavor = "current_thread")]
async fn async_handler_can_capture_arc_state() {
struct PoolStub {
calls: AtomicU32,
}
let pool = Arc::new(PoolStub {
calls: AtomicU32::new(0),
});
let pool_h = Arc::clone(&pool);
let handler = move |_p: &FlowStarted<Tcp>| {
let pool = Arc::clone(&pool_h);
async move {
tokio::task::yield_now().await;
pool.calls.fetch_add(1, Ordering::Relaxed);
Ok(())
}
};
let evt = dummy_flow_started();
AsyncHandler::<FlowStarted<Tcp>>::call(&handler, &evt)
.await
.unwrap();
AsyncHandler::<FlowStarted<Tcp>>::call(&handler, &evt)
.await
.unwrap();
assert_eq!(pool.calls.load(Ordering::Relaxed), 2);
}
}