use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
use crate::config::FanoutMode;
use crate::error::Result;
use crate::monitor::Monitor;
pub struct ShardedRunner {
iface: String,
mode: FanoutMode,
group_id: u16,
num_shards: usize,
build_shard: Arc<dyn Fn(usize) -> Result<Monitor> + Send + Sync + 'static>,
layer_specs: Vec<Box<dyn crate::layer::LayerSpec>>,
merges: Vec<crate::monitor::merge::MergeSpec>,
pin_cpus: bool,
}
impl ShardedRunner {
pub fn new<F>(
iface: impl Into<String>,
mode: FanoutMode,
group_id: u16,
num_shards: usize,
build_shard: F,
) -> Self
where
F: Fn(usize) -> Result<Monitor> + Send + Sync + 'static,
{
Self {
iface: iface.into(),
mode,
group_id,
num_shards: num_shards.max(1),
build_shard: Arc::new(build_shard),
layer_specs: Vec::new(),
merges: Vec::new(),
pin_cpus: false,
}
}
pub fn pin_cpus(mut self, on: bool) -> Self {
self.pin_cpus = on;
self
}
pub fn merge_state<T, F>(mut self, period: Duration, merge: F) -> Self
where
T: Default + Send + 'static,
F: FnMut(&mut T, T) + Send + 'static,
{
self.merges
.push(crate::monitor::merge::MergeSpec::new::<T, F>(period, merge));
self
}
pub fn state_auto_merge<T>(mut self, period: Duration) -> Self
where
T: std::ops::AddAssign + Default + Send + 'static,
{
self.merges
.push(crate::monitor::merge::MergeSpec::new::<T, _>(
period,
|p: &mut T, t: T| *p += t,
));
self
}
pub fn on_merge<T, G>(mut self, observe: G) -> Self
where
T: 'static,
G: Fn(&T) + Send + 'static,
{
let tid = std::any::TypeId::of::<T>();
if let Some(spec) = self.merges.iter_mut().find(|s| s.type_id() == tid) {
spec.set_observe::<T, G>(observe);
}
self
}
pub fn layer<L: crate::layer::LayerSpec>(mut self, spec: L) -> Self {
self.layer_specs.push(Box::new(spec));
self
}
pub fn shard_count(&self) -> usize {
self.num_shards
}
pub fn interface(&self) -> &str {
&self.iface
}
pub fn fanout(&self) -> (FanoutMode, u16) {
(self.mode, self.group_id)
}
pub fn run_until(self, deadline: Instant) -> Result<()> {
self.run_inner(RunMode::Deadline(deadline))
}
pub fn run_for(self, duration: Duration) -> Result<()> {
let deadline = Instant::now() + duration;
self.run_until(deadline)
}
pub fn run_until_signal(self) -> Result<()> {
self.run_inner(RunMode::Signal)
}
fn run_inner(self, mode: RunMode) -> Result<()> {
let stop = Arc::new(AtomicBool::new(false));
let build_shard = self.build_shard;
let pin_cpus = self.pin_cpus;
let layer_specs = Arc::new(self.layer_specs);
let mut handles = Vec::with_capacity(self.num_shards);
let merges = self.merges;
let (merge_txs, mut merge_rxs): (Vec<_>, Vec<Option<_>>) = if merges.is_empty() {
(Vec::new(), Vec::new())
} else {
(0..self.num_shards)
.map(|_| {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
(tx, Some(rx))
})
.unzip()
};
for cpu in 0..self.num_shards {
let stop = Arc::clone(&stop);
let build = Arc::clone(&build_shard);
let layer_specs = Arc::clone(&layer_specs);
let merge_rx = merge_rxs.get_mut(cpu).and_then(Option::take);
let handle = std::thread::Builder::new()
.name(format!("netring-shard-{cpu}"))
.spawn(move || -> Result<()> {
if pin_cpus && !pin_current_thread_to_core(cpu) {
tracing::warn!(shard = cpu, "could not set CPU affinity for shard");
}
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(crate::error::Error::Io)?;
let mut monitor = build(cpu)?;
for spec in layer_specs.iter() {
monitor.wrap_sink(spec.instantiate());
}
if let Some(rx) = merge_rx {
monitor.set_merge_rx(rx);
}
match mode {
RunMode::Deadline(deadline) => {
let now = Instant::now();
let dur = deadline.saturating_duration_since(now);
rt.block_on(monitor.run_for(dur))?;
}
RunMode::Signal => {
rt.block_on(monitor.run_until_signal())?;
}
}
let _ = stop;
Ok(())
})
.map_err(crate::error::Error::Io)?;
handles.push(handle);
}
let merge_handle = if merges.is_empty() {
None
} else {
let stop = Arc::clone(&stop);
let handle = std::thread::Builder::new()
.name("netring-merge".to_string())
.spawn(move || crate::monitor::merge::merge_worker(merge_txs, merges, stop))
.map_err(crate::error::Error::Io)?;
Some(handle)
};
let mut first_err: Option<crate::error::Error> = None;
for h in handles {
match h.join() {
Ok(Ok(())) => {}
Ok(Err(e)) => {
if first_err.is_none() {
first_err = Some(e);
}
}
Err(panic) => {
if first_err.is_none() {
first_err = Some(crate::error::Error::Io(std::io::Error::other(format!(
"shard thread panicked: {panic:?}"
))));
}
}
}
}
stop.store(true, Ordering::Relaxed);
if let Some(h) = merge_handle {
let _ = h.join();
}
if let Some(e) = first_err {
return Err(e);
}
Ok(())
}
}
#[derive(Copy, Clone)]
enum RunMode {
Deadline(Instant),
Signal,
}
fn pin_current_thread_to_core(index: usize) -> bool {
let n = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1)
.max(1);
let core = index % n;
unsafe {
let mut set: libc::cpu_set_t = std::mem::zeroed();
libc::CPU_SET(core, &mut set);
libc::sched_setaffinity(0, std::mem::size_of::<libc::cpu_set_t>(), &set) == 0
}
}
impl std::fmt::Debug for ShardedRunner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ShardedRunner")
.field("iface", &self.iface)
.field("mode", &self.mode)
.field("group_id", &self.group_id)
.field("num_shards", &self.num_shards)
.field("pin_cpus", &self.pin_cpus)
.finish()
}
}
#[cfg(test)]
mod tests {
#[test]
fn pin_current_thread_sets_single_core_affinity() {
let n = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
if n < 2 {
return; }
if !super::pin_current_thread_to_core(0) {
return; }
unsafe {
let mut got: libc::cpu_set_t = std::mem::zeroed();
let r = libc::sched_getaffinity(0, std::mem::size_of::<libc::cpu_set_t>(), &mut got);
assert_eq!(r, 0, "sched_getaffinity failed");
assert!(libc::CPU_ISSET(0, &got), "core 0 should be set");
assert_eq!(
libc::CPU_COUNT(&got),
1,
"affinity should be pinned to exactly one core"
);
let mut all: libc::cpu_set_t = std::mem::zeroed();
for c in 0..n {
libc::CPU_SET(c, &mut all);
}
libc::sched_setaffinity(0, std::mem::size_of::<libc::cpu_set_t>(), &all);
}
}
}