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>,
}
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),
}
}
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 mut handles = Vec::with_capacity(self.num_shards);
for cpu in 0..self.num_shards {
let stop = Arc::clone(&stop);
let build = Arc::clone(&build_shard);
let handle = std::thread::Builder::new()
.name(format!("netring-shard-{cpu}"))
.spawn(move || -> Result<()> {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.map_err(crate::error::Error::Io)?;
let monitor = build(cpu)?;
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 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(e) = first_err {
return Err(e);
}
Ok(())
}
}
#[derive(Copy, Clone)]
enum RunMode {
Deadline(Instant),
Signal,
}
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)
.finish()
}
}