use std::io;
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use bytes::BytesMut;
use parking_lot::Mutex;
use pea2pea::{
Config, ConnectionSide, Node as P2pNode, Pea2Pea,
protocols::{Reading, Writing},
};
use tokio::sync::broadcast;
use tokio::task::JoinHandle;
use crate::codec::Codec;
use crate::config::{Lane, ShapeConfig, ShapingScope};
use crate::error::Error;
use crate::frame::ID_SIZE;
use crate::scheduler::Scheduler;
use crate::shaper::{Shaper, Target};
#[derive(Clone)]
pub struct Node {
p2p: P2pNode,
shaper: Arc<Shaper>,
scheduler: Arc<Scheduler>,
scheduler_handle: Arc<Mutex<Option<JoinHandle<()>>>>,
}
impl Node {
pub fn new(config: ShapeConfig) -> Self {
let p2p_config = Config {
name: config.name.clone(),
listener_addr: config.listener_addr,
max_connections: config.max_connections,
max_connections_per_ip: config.max_connections_per_ip,
reuse_listener_port: config.reuse_listener_port,
..Config::default()
};
assert!(config.frame_size > 0, "frame_size must be non-zero");
assert!(
config.frame_size >= ID_SIZE,
"frame_size ({} bytes) must be at least ID_SIZE ({ID_SIZE} bytes) \
so every frame can carry its identifier prefix",
config.frame_size,
);
assert!(
config.frame_size <= config.max_frame_size,
"frame_size ({} bytes) must not exceed max_frame_size ({} bytes)",
config.frame_size,
config.max_frame_size,
);
if matches!(config.scope, ShapingScope::Global) {
assert!(
config.fanout > 0,
"fanout must be non-zero when scope is Global"
);
}
let p2p = P2pNode::new(p2p_config);
let shaper = Arc::new(Shaper::new(&config));
let scheduler = Arc::new(Scheduler::new(shaper.clone(), config.strategy));
Self {
p2p,
shaper,
scheduler,
scheduler_handle: Arc::new(Mutex::new(None)),
}
}
pub async fn spawn(&self) -> io::Result<Option<SocketAddr>> {
self.enable_reading().await;
self.enable_writing().await;
let addr = self.p2p.toggle_listener().await?;
let scheduler = self.scheduler.clone();
let node = self.clone();
let handle = tokio::spawn(async move {
scheduler.run(node).await;
});
*self.scheduler_handle.lock() = Some(handle);
Ok(addr)
}
pub fn send_shaped(&self, peer: SocketAddr, payload: &[u8]) -> Result<[u8; ID_SIZE], Error> {
self.shaper
.enqueue(Lane::High, Target::Unicast(peer), payload)
}
pub fn broadcast_shaped(&self, payload: &[u8]) -> Result<[u8; ID_SIZE], Error> {
self.shaper.enqueue(Lane::High, Target::Broadcast, payload)
}
pub fn send_shaped_low(
&self,
peer: SocketAddr,
payload: &[u8],
) -> Result<[u8; ID_SIZE], Error> {
self.shaper
.enqueue(Lane::Low, Target::Unicast(peer), payload)
}
pub fn broadcast_shaped_low(&self, payload: &[u8]) -> Result<[u8; ID_SIZE], Error> {
self.shaper.enqueue(Lane::Low, Target::Broadcast, payload)
}
pub fn relay_shaped(&self, frame: BytesMut) -> Result<(), Error> {
self.shaper.enqueue_raw(Lane::Low, Target::Broadcast, frame)
}
pub fn subscribe(&self) -> broadcast::Receiver<BytesMut> {
self.shaper.incoming().subscribe()
}
pub async fn connect(&self, addr: SocketAddr) -> io::Result<()> {
self.p2p.connect(addr).await
}
pub async fn disconnect(&self, addr: SocketAddr) -> bool {
self.p2p.disconnect(addr).await
}
pub fn connected_peers(&self) -> Vec<SocketAddr> {
self.p2p.connected_addrs()
}
pub async fn local_addr(&self) -> io::Result<SocketAddr> {
self.p2p.listening_addr().await
}
pub fn p2p(&self) -> &P2pNode {
&self.p2p
}
pub fn config(&self) -> &ShapeConfig {
self.shaper.config()
}
pub fn shaper(&self) -> Arc<Shaper> {
self.shaper.clone()
}
pub async fn shutdown(&self) {
self.shaper.shutting_down().store(true, Ordering::SeqCst);
self.scheduler.wake();
self.p2p.shut_down().await;
let handle = self.scheduler_handle.lock().take();
if let Some(handle) = handle {
let _ = handle.await;
}
}
}
impl Pea2Pea for Node {
fn node(&self) -> &P2pNode {
&self.p2p
}
}
impl Reading for Node {
type Message = BytesMut;
type Codec = Codec;
fn codec(&self, _addr: SocketAddr, _side: ConnectionSide) -> Self::Codec {
Codec::new(
self.shaper.config().frame_size,
self.shaper.config().max_frame_size,
)
}
async fn process_message(&self, _source: SocketAddr, message: Self::Message) {
self.shaper.handle_incoming(message);
}
}
impl Writing for Node {
type Message = BytesMut;
type Codec = Codec;
fn codec(&self, _addr: SocketAddr, _side: ConnectionSide) -> Self::Codec {
Codec::new(
self.shaper.config().frame_size,
self.shaper.config().max_frame_size,
)
}
}