use std::collections::{HashMap, HashSet, VecDeque};
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
use bytes::BytesMut;
use parking_lot::Mutex;
use tokio::sync::{Notify, broadcast};
use crate::config::{CoverGenerator, Lane, ShapeConfig, ShapingScope};
use crate::error::Error;
use crate::frame::{ID_SIZE, random_cover};
pub const SUBSCRIBER_CAPACITY: usize = 1024;
#[derive(Clone, Debug)]
#[non_exhaustive]
pub enum Target {
Unicast(SocketAddr),
Broadcast,
}
#[derive(Clone, Debug)]
pub struct PendingFrame {
pub bytes: BytesMut,
pub target: Target,
}
#[derive(Default)]
struct PeerLanes {
high: VecDeque<PendingFrame>,
low: VecDeque<PendingFrame>,
}
pub struct Shaper {
config: ShapeConfig,
incoming: broadcast::Sender<BytesMut>,
high_lane: Mutex<VecDeque<PendingFrame>>,
low_lane: Mutex<VecDeque<PendingFrame>>,
peer_lanes: Mutex<HashMap<SocketAddr, PeerLanes>>,
pc_peers: Mutex<Vec<SocketAddr>>,
shutting_down: AtomicBool,
cover_generator: Arc<dyn CoverGenerator>,
enqueue_notify: Arc<Notify>,
}
impl Shaper {
pub fn new(config: &ShapeConfig) -> Self {
assert!(
config.high_lane_capacity > 0,
"high_lane_capacity must be non-zero",
);
assert!(
config.low_lane_capacity > 0,
"low_lane_capacity must be non-zero",
);
assert!(
config.frame_size >= ID_SIZE,
"frame_size ({} bytes) must be at least ID_SIZE ({ID_SIZE} bytes)",
config.frame_size,
);
let (incoming, _) = broadcast::channel(SUBSCRIBER_CAPACITY);
let cover_generator: Arc<dyn CoverGenerator> = match &config.cover_generator {
Some(g) => g.clone(),
None => Arc::new(random_cover as fn(&ShapeConfig) -> BytesMut),
};
Self {
config: config.clone(),
incoming,
high_lane: Mutex::new(VecDeque::with_capacity(config.high_lane_capacity)),
low_lane: Mutex::new(VecDeque::with_capacity(config.low_lane_capacity)),
peer_lanes: Mutex::new(HashMap::new()),
pc_peers: Mutex::new(Vec::new()),
shutting_down: AtomicBool::new(false),
cover_generator,
enqueue_notify: Arc::new(Notify::new()),
}
}
pub fn enqueue_notify(&self) -> Arc<Notify> {
self.enqueue_notify.clone()
}
pub fn config(&self) -> &ShapeConfig {
&self.config
}
pub fn incoming(&self) -> broadcast::Sender<BytesMut> {
self.incoming.clone()
}
pub fn shutting_down(&self) -> &AtomicBool {
&self.shutting_down
}
pub fn enqueue(
&self,
lane: Lane,
target: Target,
payload: &[u8],
) -> Result<[u8; ID_SIZE], Error> {
let max_payload = self.config.frame_size - ID_SIZE;
if payload.len() > max_payload {
return Err(Error::PayloadTooLarge {
size: payload.len(),
max: max_payload,
});
}
let (id, msg) = crate::frame::build_frame(&self.config, payload);
self.enqueue_raw(lane, target, msg)?;
Ok(id)
}
pub fn enqueue_raw(&self, lane: Lane, target: Target, frame: BytesMut) -> Result<(), Error> {
if frame.len() != self.config.frame_size {
return Err(Error::FrameSizeMismatch {
size: frame.len(),
expected: self.config.frame_size,
});
}
let result = match self.config.scope {
ShapingScope::Global => self.push_global(
lane,
PendingFrame {
bytes: frame,
target,
},
),
ShapingScope::PerConnection { .. } => match target {
Target::Unicast(peer) => self.push_peer(peer, lane, frame),
Target::Broadcast => {
self.fan_out_per_connection(lane, frame);
Ok(())
}
},
};
self.enqueue_notify.notify_one();
result
}
fn push_global(&self, lane: Lane, pframe: PendingFrame) -> Result<(), Error> {
match lane {
Lane::High => {
let mut l = self.high_lane.lock();
if l.len() >= self.config.high_lane_capacity {
return Err(Error::LaneFull);
}
l.push_back(pframe);
}
Lane::Low => {
let mut l = self.low_lane.lock();
while l.len() >= self.config.low_lane_capacity {
l.pop_back();
}
l.push_front(pframe);
}
}
Ok(())
}
fn push_peer(&self, peer: SocketAddr, lane: Lane, frame: BytesMut) -> Result<(), Error> {
let pframe = PendingFrame {
bytes: frame,
target: Target::Unicast(peer),
};
let mut map = self.peer_lanes.lock();
let lanes = map.entry(peer).or_default();
match lane {
Lane::High => {
if lanes.high.len() >= self.config.high_lane_capacity {
return Err(Error::LaneFull);
}
lanes.high.push_back(pframe);
}
Lane::Low => {
while lanes.low.len() >= self.config.low_lane_capacity {
lanes.low.pop_back();
}
lanes.low.push_front(pframe);
}
}
Ok(())
}
fn fan_out_per_connection(&self, lane: Lane, frame: BytesMut) {
let peers = self.pc_peers.lock().clone();
if peers.is_empty() {
return;
}
let mut map = self.peer_lanes.lock();
for peer in peers {
let pframe = PendingFrame {
bytes: frame.clone(),
target: Target::Unicast(peer),
};
let lanes = map.entry(peer).or_default();
match lane {
Lane::High => {
while lanes.high.len() >= self.config.high_lane_capacity {
lanes.high.pop_front();
}
lanes.high.push_back(pframe);
}
Lane::Low => {
while lanes.low.len() >= self.config.low_lane_capacity {
lanes.low.pop_back();
}
lanes.low.push_front(pframe);
}
}
}
}
pub fn next_frame(&self) -> Option<PendingFrame> {
if let Some(f) = self.high_lane.lock().pop_front() {
return Some(f);
}
self.low_lane.lock().pop_front()
}
pub fn next_frame_for(&self, peer: SocketAddr) -> Option<PendingFrame> {
let mut map = self.peer_lanes.lock();
let lanes = map.get_mut(&peer)?;
if let Some(f) = lanes.high.pop_front() {
return Some(f);
}
lanes.low.pop_front()
}
pub fn refresh_pc_peers(&self, peers: &[SocketAddr]) {
let mut cache = self.pc_peers.lock();
cache.clear();
cache.extend_from_slice(peers);
}
pub fn pc_peers_snapshot(&self) -> Vec<SocketAddr> {
self.pc_peers.lock().clone()
}
pub fn prune_peer_lanes(&self, connected: &[SocketAddr]) {
let mut map = self.peer_lanes.lock();
if map.is_empty() {
return;
}
let keep: HashSet<SocketAddr> = connected.iter().copied().collect();
map.retain(|addr, _| keep.contains(addr));
}
pub fn queued(&self) -> usize {
let global = self.high_lane.lock().len() + self.low_lane.lock().len();
let per_peer: usize = self
.peer_lanes
.lock()
.values()
.map(|l| l.high.len() + l.low.len())
.sum();
global + per_peer
}
pub fn cover(&self) -> PendingFrame {
let bytes = self.cover_generator.cover(&self.config);
debug_assert_eq!(
bytes.len(),
self.config.frame_size,
"cover generator returned a frame of the wrong size",
);
PendingFrame {
bytes,
target: Target::Broadcast,
}
}
pub fn finalize_real(&self, mut frame: PendingFrame) -> PendingFrame {
frame.bytes = self
.cover_generator
.finalize_real(&self.config, frame.bytes);
debug_assert_eq!(
frame.bytes.len(),
self.config.frame_size,
"cover generator finalize_real returned a frame of the wrong size",
);
frame
}
pub fn handle_incoming(&self, message: BytesMut) {
if message.len() != self.config.frame_size {
return;
}
let _ = self.incoming.send(message);
}
}