use std::any::TypeId;
use std::marker::PhantomData;
use std::sync::Arc;
use arrayvec::ArrayVec;
use flowscope::driver::{BroadcastSlotHandle, SlotHandle, SlotMessage};
use rustc_hash::FxHashMap;
use crate::ctx::Ctx;
use crate::error::Result as NetringResult;
use crate::error::{BuildError, Result};
use crate::monitor::async_handler::{AsyncHandler, BoxFuture};
use crate::monitor::dispatcher::{
AsyncHandlerSlot, BoxedAsyncHandler, BoxedHandler, Dispatcher, DynAsyncHandler, HandlerSlot,
MAX_EVENT_TYPES,
};
use crate::monitor::handler::Handler;
use crate::protocol::Protocol;
use crate::protocol::event_typed::Event;
#[derive(Default)]
pub struct HandlerRegistry {
by_type: FxHashMap<TypeId, Vec<BoxedHandler>>,
async_by_type: FxHashMap<TypeId, Vec<BoxedAsyncHandler>>,
required_protocols: FxHashMap<TypeId, (TypeId, &'static str)>,
}
impl HandlerRegistry {
pub fn register<E, H, M>(&mut self, handler: H)
where
E: Event,
H: Handler<E, M>,
M: 'static,
{
let boxed: BoxedHandler = Arc::new(move |ptr, ctx| {
let typed: &E::Payload = unsafe { &*(ptr as *const E::Payload) };
handler.call(typed, ctx)
});
self.by_type
.entry(TypeId::of::<E::Payload>())
.or_default()
.push(boxed);
if let Some(p_id) = E::protocol_marker() {
self.required_protocols
.entry(TypeId::of::<E::Payload>())
.or_insert((p_id, E::protocol_name()));
}
}
pub fn register_async<E, H>(&mut self, handler: H)
where
E: Event,
H: AsyncHandler<E>,
{
let boxed: BoxedAsyncHandler = Arc::new(AsyncHandlerWrapper::<E, H>::new(handler));
self.async_by_type
.entry(TypeId::of::<E::Payload>())
.or_default()
.push(boxed);
if let Some(p_id) = E::protocol_marker() {
self.required_protocols
.entry(TypeId::of::<E::Payload>())
.or_insert((p_id, E::protocol_name()));
}
}
pub fn required_protocols(&self) -> impl Iterator<Item = (TypeId, &'static str)> + '_ {
self.required_protocols.values().copied()
}
pub fn type_count(&self) -> usize {
let mut ids: std::collections::HashSet<TypeId> = std::collections::HashSet::new();
ids.extend(self.by_type.keys().copied());
ids.extend(self.async_by_type.keys().copied());
ids.len()
}
pub fn handler_count(&self) -> usize {
self.by_type.values().map(|v| v.len()).sum()
}
pub fn async_handler_count(&self) -> usize {
self.async_by_type.values().map(|v| v.len()).sum()
}
pub fn into_dispatcher(mut self) -> std::result::Result<Dispatcher, BuildError> {
let mut all_types: Vec<TypeId> = self
.by_type
.keys()
.copied()
.chain(self.async_by_type.keys().copied())
.collect();
all_types.sort_unstable_by_key(|t| format!("{t:?}"));
all_types.dedup();
if all_types.len() > MAX_EVENT_TYPES {
return Err(BuildError::TooManyEventTypes {
limit: MAX_EVENT_TYPES,
actual: all_types.len(),
});
}
let mut slot_by_type: ArrayVec<(TypeId, u8), MAX_EVENT_TYPES> = ArrayVec::new();
let mut slots: Vec<Vec<HandlerSlot>> = Vec::with_capacity(all_types.len());
let mut async_slots: Vec<Vec<AsyncHandlerSlot>> = Vec::with_capacity(all_types.len());
for (i, type_id) in all_types.into_iter().enumerate() {
slot_by_type.push((type_id, i as u8));
slots.push(
self.by_type
.remove(&type_id)
.unwrap_or_default()
.into_iter()
.map(|h| HandlerSlot { handler: h })
.collect(),
);
async_slots.push(
self.async_by_type
.remove(&type_id)
.unwrap_or_default()
.into_iter()
.map(|h| AsyncHandlerSlot { handler: h })
.collect(),
);
}
Ok(Dispatcher::new(
slot_by_type,
slots.into_boxed_slice(),
async_slots.into_boxed_slice(),
))
}
}
struct AsyncHandlerWrapper<E, H> {
handler: H,
_marker: PhantomData<fn() -> E>,
}
impl<E, H> AsyncHandlerWrapper<E, H> {
fn new(handler: H) -> Self {
Self {
handler,
_marker: PhantomData,
}
}
}
impl<E, H> DynAsyncHandler for AsyncHandlerWrapper<E, H>
where
E: Event,
H: AsyncHandler<E>,
{
fn call(&self, ptr: *const ()) -> BoxFuture<NetringResult<()>> {
let typed: &E::Payload = unsafe { &*(ptr as *const E::Payload) };
self.handler.call(typed)
}
}
pub trait ProtocolSlot: Send {
fn drain_and_dispatch(&mut self, dispatcher: &mut Dispatcher, ctx: &mut Ctx<'_>) -> Result<()>;
}
pub struct TypedProtocolSlot<P: Protocol> {
handle: SlotHandle<P::Message, flowscope::extract::FiveTupleKey>,
scratch: Vec<SlotMessage<P::Message, flowscope::extract::FiveTupleKey>>,
_marker: PhantomData<fn() -> P>,
}
impl<P: Protocol> TypedProtocolSlot<P> {
pub fn new(handle: SlotHandle<P::Message, flowscope::extract::FiveTupleKey>) -> Self {
Self {
handle,
scratch: Vec::new(),
_marker: PhantomData,
}
}
pub fn handle(&self) -> &SlotHandle<P::Message, flowscope::extract::FiveTupleKey> {
&self.handle
}
}
impl<P: Protocol> ProtocolSlot for TypedProtocolSlot<P> {
fn drain_and_dispatch(&mut self, dispatcher: &mut Dispatcher, ctx: &mut Ctx<'_>) -> Result<()> {
self.scratch.clear();
let n = self.handle.drain(&mut self.scratch);
if n == 0 {
return Ok(());
}
let saved_flow = ctx.flow;
let saved_ts = ctx.ts;
for slot_msg in self.scratch.drain(..) {
ctx.flow = Some(slot_msg.key);
ctx.ts = slot_msg.ts;
dispatcher.dispatch::<P::Message>(&slot_msg.message, ctx)?;
}
ctx.flow = saved_flow;
ctx.ts = saved_ts;
Ok(())
}
}
#[cfg(feature = "icmp")]
pub struct IcmpSlot {
handle: SlotHandle<flowscope::icmp::IcmpMessage, flowscope::extract::FiveTupleKey>,
scratch: Vec<SlotMessage<flowscope::icmp::IcmpMessage, flowscope::extract::FiveTupleKey>>,
}
#[cfg(feature = "icmp")]
impl IcmpSlot {
pub fn new(
handle: SlotHandle<flowscope::icmp::IcmpMessage, flowscope::extract::FiveTupleKey>,
) -> Self {
Self {
handle,
scratch: Vec::new(),
}
}
}
#[cfg(feature = "icmp")]
impl ProtocolSlot for IcmpSlot {
fn drain_and_dispatch(&mut self, dispatcher: &mut Dispatcher, ctx: &mut Ctx<'_>) -> Result<()> {
use crate::protocol::event_typed::{IcmpError, classify_icmp_error};
self.scratch.clear();
let n = self.handle.drain(&mut self.scratch);
if n == 0 {
return Ok(());
}
let saved_flow = ctx.flow;
let saved_ts = ctx.ts;
for slot_msg in self.scratch.drain(..) {
ctx.flow = Some(slot_msg.key);
ctx.ts = slot_msg.ts;
dispatcher.dispatch::<flowscope::icmp::IcmpMessage>(&slot_msg.message, ctx)?;
if let Some(kind) = classify_icmp_error(&slot_msg.message) {
let inner = slot_msg.message.error_inner().map(|(_, i)| i);
let correlated_flow =
inner.and_then(flowscope::extract::FiveTupleKey::from_inner_canonical);
let stats = inner.and_then(|i| ctx.lookup_icmp_flow(i).map(|(_, s)| s));
let err = IcmpError {
family: slot_msg.message.family,
kind,
correlated_flow,
stats,
ts: slot_msg.ts,
};
dispatcher.dispatch::<IcmpError>(&err, ctx)?;
}
}
ctx.flow = saved_flow;
ctx.ts = saved_ts;
Ok(())
}
}
pub struct TypedBroadcastProtocolSlot<P: Protocol>
where
P::Message: Send + Sync + Clone + 'static,
{
handle: BroadcastSlotHandle<P::Message, flowscope::extract::FiveTupleKey>,
scratch: Vec<SlotMessage<P::Message, flowscope::extract::FiveTupleKey>>,
_marker: PhantomData<fn() -> P>,
}
impl<P: Protocol> TypedBroadcastProtocolSlot<P>
where
P::Message: Send + Sync + Clone + 'static,
{
pub fn new(handle: BroadcastSlotHandle<P::Message, flowscope::extract::FiveTupleKey>) -> Self {
Self {
handle,
scratch: Vec::new(),
_marker: PhantomData,
}
}
}
impl<P: Protocol> ProtocolSlot for TypedBroadcastProtocolSlot<P>
where
P::Message: Send + Sync + Clone + 'static,
{
fn drain_and_dispatch(&mut self, dispatcher: &mut Dispatcher, ctx: &mut Ctx<'_>) -> Result<()> {
self.scratch.clear();
let n = self.handle.drain(&mut self.scratch);
if n == 0 {
return Ok(());
}
let saved_flow = ctx.flow;
let saved_ts = ctx.ts;
for slot_msg in self.scratch.drain(..) {
ctx.flow = Some(slot_msg.key);
ctx.ts = slot_msg.ts;
dispatcher.dispatch::<P::Message>(&slot_msg.message, ctx)?;
}
ctx.flow = saved_flow;
ctx.ts = saved_ts;
Ok(())
}
}
#[cfg(test)]
mod tests {
use flowscope::Timestamp;
use super::*;
use crate::anomaly::sink::NoopSink;
use crate::ctx::{CounterRegistry, SourceIdx, StateMap};
use crate::error::Error;
use crate::protocol::builtin::Tcp;
use crate::protocol::event_typed::FlowStarted;
fn fresh_ctx<'a>(
state: &'a mut StateMap,
sink: &'a mut NoopSink,
counters: &'a mut CounterRegistry,
flow_states: &'a mut crate::ctx::FlowStateRegistry,
) -> Ctx<'a> {
Ctx {
flow: None,
ts: Timestamp::new(0, 0),
source: SourceIdx(0),
monitor_name: None,
state_map: state,
sink,
counters,
flow_states,
label_table: crate::ctx::default_label_table(),
tracker: None,
}
}
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))
}
#[test]
fn register_one_handler_then_dispatch_fires_it() {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
let counter = Arc::new(AtomicU32::new(0));
let c = Arc::clone(&counter);
let mut reg = HandlerRegistry::default();
reg.register::<FlowStarted<Tcp>, _, _>(move |_evt: &FlowStarted<Tcp>| {
c.fetch_add(1, Ordering::Relaxed);
Ok(())
});
assert_eq!(reg.type_count(), 1);
assert_eq!(reg.handler_count(), 1);
let mut disp = reg.into_dispatcher().unwrap();
let mut s = StateMap::default();
let mut sink = NoopSink;
let mut cr = CounterRegistry::default();
let mut fs = crate::ctx::FlowStateRegistry::default();
let mut ctx = fresh_ctx(&mut s, &mut sink, &mut cr, &mut fs);
let evt = dummy_flow_started();
disp.dispatch::<FlowStarted<Tcp>>(&evt, &mut ctx).unwrap();
assert_eq!(counter.load(Ordering::Relaxed), 1);
}
#[test]
fn register_two_handlers_for_same_event_both_fire() {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
let a = Arc::new(AtomicU32::new(0));
let b = Arc::new(AtomicU32::new(0));
let a_h = Arc::clone(&a);
let b_h = Arc::clone(&b);
let mut reg = HandlerRegistry::default();
reg.register::<FlowStarted<Tcp>, _, _>(move |_evt: &FlowStarted<Tcp>| {
a_h.fetch_add(1, Ordering::Relaxed);
Ok(())
});
reg.register::<FlowStarted<Tcp>, _, _>(move |_evt: &FlowStarted<Tcp>| {
b_h.fetch_add(1, Ordering::Relaxed);
Ok(())
});
assert_eq!(reg.type_count(), 1);
assert_eq!(reg.handler_count(), 2);
let mut disp = reg.into_dispatcher().unwrap();
let mut s = StateMap::default();
let mut sink = NoopSink;
let mut cr = CounterRegistry::default();
let mut fs = crate::ctx::FlowStateRegistry::default();
let mut ctx = fresh_ctx(&mut s, &mut sink, &mut cr, &mut fs);
let evt = dummy_flow_started();
disp.dispatch::<FlowStarted<Tcp>>(&evt, &mut ctx).unwrap();
assert_eq!(a.load(Ordering::Relaxed), 1);
assert_eq!(b.load(Ordering::Relaxed), 1);
}
#[test]
fn handler_error_short_circuits_remaining_handlers() {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
let after = Arc::new(AtomicU32::new(0));
let after_h = Arc::clone(&after);
let mut reg = HandlerRegistry::default();
reg.register::<FlowStarted<Tcp>, _, _>(|_evt: &FlowStarted<Tcp>| {
Err(Error::Config("boom".into()))
});
reg.register::<FlowStarted<Tcp>, _, _>(move |_evt: &FlowStarted<Tcp>| {
after_h.fetch_add(1, Ordering::Relaxed);
Ok(())
});
let mut disp = reg.into_dispatcher().unwrap();
let mut s = StateMap::default();
let mut sink = NoopSink;
let mut cr = CounterRegistry::default();
let mut fs = crate::ctx::FlowStateRegistry::default();
let mut ctx = fresh_ctx(&mut s, &mut sink, &mut cr, &mut fs);
let evt = dummy_flow_started();
let res = disp.dispatch::<FlowStarted<Tcp>>(&evt, &mut ctx);
assert!(res.is_err());
assert_eq!(
after.load(Ordering::Relaxed),
0,
"second handler must not fire after first errored"
);
}
#[test]
fn too_many_event_types_errors_at_build() {
macro_rules! synth {
($($name:ident),+ $(,)?) => {
$(
#[derive(Debug)]
struct $name;
impl Event for $name { type Payload = $name; }
)+
};
}
synth!(
E0, E1, E2, E3, E4, E5, E6, E7, E8, E9, E10, E11, E12, E13, E14, E15, E16
);
let mut reg = HandlerRegistry::default();
reg.register::<E0, _, _>(|_: &E0| Ok(()));
reg.register::<E1, _, _>(|_: &E1| Ok(()));
reg.register::<E2, _, _>(|_: &E2| Ok(()));
reg.register::<E3, _, _>(|_: &E3| Ok(()));
reg.register::<E4, _, _>(|_: &E4| Ok(()));
reg.register::<E5, _, _>(|_: &E5| Ok(()));
reg.register::<E6, _, _>(|_: &E6| Ok(()));
reg.register::<E7, _, _>(|_: &E7| Ok(()));
reg.register::<E8, _, _>(|_: &E8| Ok(()));
reg.register::<E9, _, _>(|_: &E9| Ok(()));
reg.register::<E10, _, _>(|_: &E10| Ok(()));
reg.register::<E11, _, _>(|_: &E11| Ok(()));
reg.register::<E12, _, _>(|_: &E12| Ok(()));
reg.register::<E13, _, _>(|_: &E13| Ok(()));
reg.register::<E14, _, _>(|_: &E14| Ok(()));
reg.register::<E15, _, _>(|_: &E15| Ok(()));
reg.register::<E16, _, _>(|_: &E16| Ok(()));
let err = reg.into_dispatcher().unwrap_err();
match err {
BuildError::TooManyEventTypes { limit, actual } => {
assert_eq!(limit, MAX_EVENT_TYPES);
assert_eq!(actual, MAX_EVENT_TYPES + 1);
}
other => panic!("expected TooManyEventTypes, got {other:?}"),
}
}
#[cfg(feature = "icmp")]
#[test]
fn icmp_slot_synthesises_icmp_error_from_a_real_frame() {
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use flowscope::driver::Driver;
use flowscope::extract::FiveTuple;
use crate::protocol::Protocol;
use crate::protocol::builtin::Icmp;
use crate::protocol::event_typed::IcmpError;
let mut builder = Driver::builder(FiveTuple::bidirectional());
let handle = Icmp::register(&mut builder).expect("icmp registers");
let mut slot = Icmp::make_slot(handle);
let mut driver = builder.build();
let seen = Arc::new(AtomicU32::new(0));
let s = Arc::clone(&seen);
let mut reg = HandlerRegistry::default();
reg.register::<IcmpError, _, crate::monitor::handler::PayloadCtx>(
move |err: &IcmpError, _ctx: &mut Ctx<'_>| {
assert_eq!(err.kind.as_str(), "port_unreachable");
assert!(err.correlated_flow.is_some(), "inner 5-tuple joined");
s.fetch_add(1, Ordering::Relaxed);
Ok(())
},
);
let mut dispatcher = reg.into_dispatcher().unwrap();
use etherparse::{Ethernet2Header, IpNumber, Ipv4Header};
let mut inner = Vec::new();
inner.extend_from_slice(&[0x45, 0, 0x00, 0x28, 0, 0, 0, 0, 64, 6, 0, 0]);
inner.extend_from_slice(&[10, 0, 0, 1]); inner.extend_from_slice(&[10, 0, 0, 2]); inner.extend_from_slice(&12345u16.to_be_bytes()); inner.extend_from_slice(&80u16.to_be_bytes()); inner.extend_from_slice(&[0, 0, 0, 1]); let mut icmp = vec![3u8, 3, 0, 0, 0, 0, 0, 0]; icmp.extend_from_slice(&inner);
let ip = Ipv4Header::new(
icmp.len() as u16,
64,
IpNumber::ICMP,
[192, 0, 2, 1],
[192, 0, 2, 2],
)
.unwrap();
let eth = Ethernet2Header {
destination: [2u8; 6],
source: [1u8; 6],
ether_type: etherparse::EtherType::IPV4,
};
let mut frame = Vec::new();
eth.write(&mut frame).unwrap();
ip.write(&mut frame).unwrap();
frame.extend_from_slice(&icmp);
let ts = Timestamp::new(1, 0);
let view = flowscope::PacketView::new(&frame, ts);
let mut events = Vec::new();
driver.track_into(view, &mut events);
let mut state = StateMap::default();
let mut sink = NoopSink;
let mut counters = CounterRegistry::default();
let mut flow_states = crate::ctx::FlowStateRegistry::default();
let mut ctx = fresh_ctx(&mut state, &mut sink, &mut counters, &mut flow_states);
ctx.tracker = Some(driver.tracker());
slot.drain_and_dispatch(&mut dispatcher, &mut ctx).unwrap();
assert_eq!(seen.load(Ordering::Relaxed), 1, "one IcmpError synthesised");
}
}