use std::net::SocketAddr;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::time::Duration;
use parking_lot::Mutex;
use rand::Rng;
use tokio::sync::Notify;
use tokio::time;
use tracing::debug;
use crate::config::{ShapingScope, ShapingStrategy};
use crate::node::Node;
use crate::shaper::{PendingFrame, Shaper, Target};
use pea2pea::{protocols::Writing, Pea2Pea};
pub struct Scheduler {
shaper: Arc<Shaper>,
strategy: ShapingStrategy,
wake: Notify,
cursor: Mutex<usize>,
}
impl Scheduler {
pub fn new(shaper: Arc<Shaper>, strategy: ShapingStrategy) -> Self {
Self {
shaper,
strategy,
wake: Notify::new(),
cursor: Mutex::new(0),
}
}
pub fn wake(&self) {
self.wake.notify_one();
}
pub async fn run(&self, node: Node) {
match self.strategy {
ShapingStrategy::Constant { interval } => self.run_constant(node, interval).await,
ShapingStrategy::Poisson { rate } => self.run_poisson(node, rate).await,
}
}
async fn run_constant(&self, node: Node, interval: Duration) {
let mut ticker = time::interval(interval);
ticker.set_missed_tick_behavior(time::MissedTickBehavior::Skip);
ticker.tick().await;
loop {
tokio::select! {
_ = ticker.tick() => {}
_ = self.wake.notified() => break,
}
if self.is_shutting_down() {
break;
}
self.send_one(&node).await;
}
}
async fn run_poisson(&self, node: Node, rate: f64) {
loop {
let dur = {
let mut rng = rand::rng();
let u: f64 = rng.random::<f64>().clamp(f64::MIN_POSITIVE, 1.0);
let secs = (-u.ln() / rate).min(60.0);
Duration::from_secs_f64(secs)
};
tokio::select! {
_ = time::sleep(dur) => {}
_ = self.wake.notified() => break,
}
if self.is_shutting_down() {
break;
}
self.send_one(&node).await;
}
}
fn is_shutting_down(&self) -> bool {
self.shaper.shutting_down().load(Ordering::SeqCst)
}
async fn send_one(&self, node: &Node) {
let peers = node.connected_peers();
if peers.is_empty() {
return;
}
let frame = self
.shaper
.next_frame()
.unwrap_or_else(|| self.shaper.cover());
self.dispatch(node, &peers, frame);
}
fn dispatch(&self, node: &Node, peers: &[SocketAddr], frame: PendingFrame) {
if let Target::Unicast(peer) = &frame.target {
if peers.contains(peer) {
self.send_one_to(node, *peer, &frame);
return;
}
debug!(
parent: node.node().span(),
"unicast target {peer} is no longer connected; falling back to scope-based dispatch"
);
}
match self.shaper.config().scope {
ShapingScope::Global => {
let fanout = self.shaper.config().fanout.min(peers.len()).max(1);
for peer in self.pick_targets(peers, fanout) {
self.send_one_to(node, peer, &frame);
}
}
ShapingScope::PerConnection { randomize } => {
let pick = self.pick_round_robin(peers, randomize);
self.send_one_to(node, pick, &frame);
}
}
}
fn send_one_to(&self, node: &Node, peer: SocketAddr, frame: &PendingFrame) {
if let Err(e) = node.unicast_fast(peer, frame.bytes.clone()) {
debug!(parent: node.node().span(), "send to {peer} failed: {e}");
}
}
fn pick_targets(&self, peers: &[SocketAddr], fanout: usize) -> Vec<SocketAddr> {
let mut rng = rand::rng();
let mut cursor = self.cursor.lock();
let n = peers.len();
if n == 0 {
return Vec::new();
}
let mut pool: Vec<usize> = (0..n).collect();
let mut chosen: Vec<SocketAddr> = Vec::with_capacity(fanout);
for _ in 0..fanout.min(n) {
let remaining = pool.len();
if remaining == 0 {
break;
}
let offset = rng.random_range(0..remaining);
let pick = (*cursor + offset) % remaining;
*cursor = cursor.wrapping_add(1);
let pos = pool.remove(pick);
chosen.push(peers[pos]);
}
chosen
}
fn pick_round_robin(&self, peers: &[SocketAddr], randomize: bool) -> SocketAddr {
let n = peers.len();
debug_assert!(n > 0);
let mut cursor = self.cursor.lock();
let pick = if randomize {
let mut rng = rand::rng();
let offset = rng.random_range(0..n);
(*cursor + offset) % n
} else {
*cursor % n
};
*cursor = cursor.wrapping_add(1);
peers[pick]
}
}