use std::{
collections::HashMap,
fmt,
io::{self, ErrorKind},
net::SocketAddr,
ops::Deref,
sync::{
Arc,
atomic::{AtomicBool, AtomicUsize, Ordering::*},
},
time::{Duration, Instant},
};
use futures_util::stream::{FuturesUnordered, StreamExt};
use parking_lot::Mutex;
use tokio::{
net::{TcpListener, TcpSocket, TcpStream},
sync::{Notify, RwLock, Semaphore, oneshot, watch},
task::JoinHandle,
time::{sleep, timeout},
};
use tracing::*;
#[cfg(doc)]
use crate::protocols::{Handshake, OnConnect, OnDisconnect, Reading, Writing};
use crate::{
Config, Heuristics, Stats,
connections::{
Connection, ConnectionGuard, ConnectionInfo, ConnectionSide, Connections, DisconnectOrigin,
create_connection_span,
},
protocols::{Protocol, Protocols},
};
macro_rules! enable_protocol {
($handler_type:ident, $node:expr, $conn:expr) => {
if let Some(handler) = $node.protocols.$handler_type.get() {
let (conn_returner, conn_retriever) = oneshot::channel();
handler.trigger(($conn, conn_returner)).await;
match crate::protocols::await_handler_response(conn_retriever, handler.closed()).await {
Some(Ok(conn)) => conn,
None => return Err(shutting_down_error()),
Some(e) => return e,
}
} else {
$conn
}
};
}
async fn wait_for_drain(notify: &Notify, drained: impl Fn() -> bool) {
loop {
let notified = notify.notified();
tokio::pin!(notified);
notified.as_mut().enable(); if drained() {
break;
}
notified.await;
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ShuttingDown;
impl fmt::Display for ShuttingDown {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("shutting down")
}
}
impl std::error::Error for ShuttingDown {}
impl ShuttingDown {
pub fn caused(err: &io::Error) -> bool {
err.get_ref().is_some_and(|inner| inner.is::<Self>())
}
}
pub(crate) fn shutting_down_error() -> io::Error {
io::Error::other(ShuttingDown)
}
static SEQUENTIAL_NODE_ID: AtomicUsize = AtomicUsize::new(0);
#[derive(Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) enum NodeTask {
Listener,
OnDisconnect,
Handshake,
OnConnect,
Reading,
Writing,
}
#[derive(Clone)]
pub struct Node(Arc<InnerNode>);
impl Deref for Node {
type Target = Arc<InnerNode>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
#[doc(hidden)]
pub struct InnerNode {
span: Span,
config: Config,
listening_addr: RwLock<Option<SocketAddr>>,
pub(crate) protocols: Protocols,
pub(crate) connections: Connections,
connecting_permits: Arc<Semaphore>,
stats: Stats,
heuristics: Heuristics,
pub(crate) tasks: Mutex<HashMap<NodeTask, JoinHandle<()>>>,
pub(crate) shutdown: ShutdownState,
}
pub(crate) struct ShutdownState {
flag: AtomicBool,
handler_signal: watch::Sender<bool>,
sched_in_flight: AtomicUsize,
sched_drained: Notify,
}
impl Default for ShutdownState {
fn default() -> Self {
Self {
flag: Default::default(),
handler_signal: watch::Sender::new(false),
sched_in_flight: Default::default(),
sched_drained: Notify::new(),
}
}
}
impl ShutdownState {
pub(crate) fn is_underway(&self) -> bool {
self.flag.load(Acquire)
}
fn begin(&self) {
self.flag.store(true, Release);
}
fn signal_handlers(&self) {
let _ = self.handler_signal.send(true);
}
pub(crate) fn handler_signal(&self) -> watch::Receiver<bool> {
self.handler_signal.subscribe()
}
fn sched_started(&self) {
self.sched_in_flight.fetch_add(1, AcqRel);
}
fn sched_finished(&self) {
if self.sched_in_flight.fetch_sub(1, AcqRel) == 1 {
self.sched_drained.notify_waiters();
}
}
async fn wait_sched_drained(&self) {
wait_for_drain(&self.sched_drained, || {
self.sched_in_flight.load(Acquire) == 0
})
.await;
}
}
impl Node {
pub fn new(mut config: Config) -> Self {
assert!(
config.max_connections != 0,
"Config::max_connections must not be 0"
);
assert!(
config.max_connections_per_ip != 0,
"Config::max_connections_per_ip must not be 0"
);
assert!(
config.max_connecting != 0,
"Config::max_connecting must not be 0"
);
if config.name.is_none() {
config.name = Some(SEQUENTIAL_NODE_ID.fetch_add(1, Relaxed).to_string());
}
let span = create_span(config.name.as_deref().unwrap());
if config.max_connecting > config.max_connections {
debug!(
parent: &span,
"clamping max_connecting ({}) to max_connections ({})",
config.max_connecting, config.max_connections,
);
config.max_connecting = config.max_connections;
}
let connecting_permits = Arc::new(Semaphore::new(config.max_connecting as usize));
let node = Node(Arc::new(InnerNode {
span,
config,
listening_addr: Default::default(),
shutdown: Default::default(),
protocols: Default::default(),
connections: Default::default(),
connecting_permits,
stats: Default::default(),
heuristics: Default::default(),
tasks: Default::default(),
}));
debug!(parent: node.span(), "the node is ready");
node
}
pub async fn toggle_listener(&self) -> io::Result<Option<SocketAddr>> {
let mut listening_addr = self.listening_addr.write().await;
if let Some(old_listening_addr) = *listening_addr {
let Some(listener_task) = self.tasks.lock().remove(&NodeTask::Listener) else {
return Err(shutting_down_error());
};
listener_task.abort();
trace!(parent: self.span(), "aborted the listening task");
debug!(parent: self.span(), "no longer listening on {old_listening_addr}");
*listening_addr = None;
Ok(None)
} else {
let listener_addr = self.config().listener_addr.ok_or_else(|| {
error!(parent: self.span(), "the listener was toggled on, but Config::listener_addr is not set");
ErrorKind::AddrNotAvailable
})?;
trace!(parent: self.span(), "attempting to listen on {listener_addr}");
let socket = match listener_addr {
SocketAddr::V4(_) => TcpSocket::new_v4()?,
SocketAddr::V6(_) => TcpSocket::new_v6()?,
};
socket.set_reuseaddr(true)?;
#[cfg(all(unix, not(any(target_os = "solaris", target_os = "illumos"))))]
if self.config().reuse_listener_port {
socket.set_reuseport(true)?;
}
#[cfg(not(all(unix, not(any(target_os = "solaris", target_os = "illumos")))))]
if self.config().reuse_listener_port {
return Err(io::Error::new(
ErrorKind::Unsupported,
"Config::reuse_listener_port is set, but SO_REUSEPORT is not supported on this platform",
));
}
socket.bind(listener_addr)?;
let listener = socket.listen(self.config().listener_backlog)?; let port = listener.local_addr()?.port(); let new_listening_addr = (listener_addr.ip(), port).into();
self.start_listening(listener).await?;
debug!(parent: self.span(), "listening on {new_listening_addr}");
*listening_addr = Some(new_listening_addr);
Ok(Some(new_listening_addr))
}
}
async fn start_listening(&self, listener: TcpListener) -> io::Result<()> {
let (tx, rx) = oneshot::channel();
let node = self.clone();
let listening_task = UnregisteredTask::new(tokio::spawn(async move {
trace!(parent: node.span(), "spawned the listening task");
if tx.send(()).is_err() {
error!(parent: node.span(), "listener setup interrupted; shutting down the listening task");
return;
}
let inbound_permits = node.connecting_permits.clone();
loop {
match listener.accept().await {
Ok((stream, addr)) => {
let permit = match inbound_permits.clone().acquire_owned().await {
Ok(p) => p,
Err(_) => {
error!(parent: node.span(), "inbound permit semaphore closed unexpectedly");
debug_assert!(
false,
"acquiring an owned listening semaphore failed"
);
return;
}
};
let node = node.clone();
tokio::spawn(async move {
let _permit = permit;
node.handle_connection_request(stream, addr).await.inspect_err(|e|
match e.kind() {
ErrorKind::QuotaExceeded | ErrorKind::AlreadyExists => {
debug!(parent: node.span(), "rejecting connection from {addr}: {e}");
}
_ if node.shutdown.is_underway() => {
debug!(parent: node.span(), "dropping the connection from {addr}: {e}");
}
_ => {
error!(parent: node.span(), "couldn't accept a connection from {addr}: {e}");
}
}
)
});
}
Err(e) => {
match e.kind() {
ErrorKind::ConnectionAborted | ErrorKind::ConnectionReset => {
debug!(parent: node.span(), "transient accept error: {e}");
}
_ => {
node.heuristics.register_accept_error();
error!(parent: node.span(), "couldn't accept a connection: {e}");
sleep(Duration::from_millis(500)).await;
}
}
}
}
}
}));
let _ = rx.await;
listening_task.register(self, NodeTask::Listener)?;
Ok(())
}
async fn handle_connection_request(
&self,
stream: TcpStream,
addr: SocketAddr,
) -> io::Result<()> {
if self.shutdown.is_underway() {
return Err(shutting_down_error());
}
let guard = self.check_and_reserve(addr).inspect_err(|_| {
self.heuristics.register_inbound_rejection();
})?;
self.adapt_stream(stream, addr, ConnectionSide::Responder, guard)
.await
}
#[inline]
pub fn name(&self) -> &str {
self.config.name.as_deref().unwrap()
}
#[inline]
pub fn config(&self) -> &Config {
&self.config
}
#[inline]
pub fn stats(&self) -> &Stats {
&self.stats
}
#[inline]
pub fn heuristics(&self) -> &Heuristics {
&self.heuristics
}
#[inline]
pub fn span(&self) -> &Span {
&self.span
}
pub async fn listening_addr(&self) -> io::Result<SocketAddr> {
self.listening_addr
.read()
.await
.as_ref()
.copied()
.ok_or_else(|| ErrorKind::AddrNotAvailable.into())
}
async fn enable_protocols(&self, conn: Connection) -> io::Result<Connection> {
let mut conn = enable_protocol!(handshake, self, conn);
if let Some(stream) = conn.stream.take() {
let (reader, writer) = stream.into_split();
conn.reader = Some(Box::new(reader));
conn.writer = Some(Box::new(writer));
}
let conn = enable_protocol!(reading, self, conn);
let conn = enable_protocol!(writing, self, conn);
Ok(conn)
}
async fn adapt_stream(
&self,
stream: TcpStream,
peer_addr: SocketAddr,
own_side: ConnectionSide,
guard: ConnectionGuard<'_>,
) -> io::Result<()> {
let conn_span = create_connection_span(peer_addr, self.span());
debug!(parent: &conn_span, "establishing connection as the {own_side:?}");
if own_side == ConnectionSide::Initiator && enabled!(Level::TRACE) {
if let Ok(addr) = stream.local_addr() {
trace!(parent: &conn_span, "the peer is connected on port {}", addr.port());
} else {
warn!(parent: &conn_span, "couldn't determine the peer-side port");
}
}
let connection = Connection::new(peer_addr, stream, !own_side, conn_span.clone());
let mut connection = self.enable_protocols(connection).await?;
let conn_ready_tx = connection.readiness_notifier.take();
let conn_id = connection.id;
let sched_guard = self
.protocols
.on_connect
.get()
.is_some()
.then(|| SchedulingGuard::new(self.clone()));
self.connections.add(connection, guard, &self.shutdown)?;
if let Some(tx) = conn_ready_tx {
let _ = tx.send(());
}
debug!(parent: &conn_span, "fully connected");
if sched_guard.is_some() {
trace!(parent: &conn_span, "executing OnConnect logic...");
let node = self.clone();
let scheduling_task = tokio::spawn(async move {
let _sched_guard = sched_guard;
let Some(handler) = node.protocols.on_connect.get() else {
return; };
let (sender, receiver) = oneshot::channel();
handler.trigger(((peer_addr, conn_id), sender)).await;
if let Some((handle, abortable)) =
crate::protocols::await_handler_response(receiver, handler.closed()).await
{
if !abortable {
drop(handle);
} else if let Some(conn) = node
.connections
.active
.write()
.get_mut(&peer_addr)
.filter(|conn| conn.id == conn_id)
{
conn.tasks.push(handle);
} else {
handle.abort();
}
}
});
let _ = scheduling_task.await;
}
Ok(())
}
async fn create_stream(
&self,
addr: SocketAddr,
socket: Option<TcpSocket>,
) -> io::Result<TcpStream> {
match timeout(
Duration::from_millis(self.config().connection_timeout_ms.into()),
self.create_stream_inner(addr, socket),
)
.await
{
Ok(Ok(stream)) => Ok(stream),
Ok(err) => err,
Err(err) => Err(io::Error::new(ErrorKind::TimedOut, err)),
}
}
async fn create_stream_inner(
&self,
addr: SocketAddr,
socket: Option<TcpSocket>,
) -> io::Result<TcpStream> {
if let Some(socket) = socket {
socket.connect(addr).await
} else {
TcpStream::connect(addr).await
}
}
pub async fn connect(&self, addr: SocketAddr) -> io::Result<()> {
self.connect_inner(addr, None)
.await
.inspect_err(|e| error!(parent: self.span(), "couldn't connect to {addr}: {e}"))
}
pub async fn connect_using_socket(
&self,
addr: SocketAddr,
socket: TcpSocket,
) -> io::Result<()> {
self.connect_inner(addr, Some(socket))
.await
.inspect_err(|e| error!(parent: self.span(), "couldn't connect to {addr}: {e}"))
}
async fn connect_inner(&self, addr: SocketAddr, socket: Option<TcpSocket>) -> io::Result<()> {
if self.shutdown.is_underway() {
return Err(shutting_down_error());
}
if let Ok(listening_addr) = self.listening_addr().await
&& (addr == listening_addr
|| listening_addr.ip().is_unspecified()
&& addr.ip().is_loopback()
&& addr.port() == listening_addr.port())
{
return Err(io::Error::new(
ErrorKind::AddrInUse,
format!("can't connect to node's own listening address ({addr})"),
));
}
if self.connections.is_connected(addr) {
return Err(io::Error::new(
ErrorKind::AlreadyExists,
"already connected",
));
}
let _permit = self
.connecting_permits
.clone()
.try_acquire_owned()
.map_err(|_| {
self.heuristics.register_connect_budget_rejection();
io::Error::new(
ErrorKind::QuotaExceeded,
format!(
"maximum number ({}) of pending connections reached",
self.config.max_connecting
),
)
})?;
let guard = self.check_and_reserve(addr)?;
let stream = self.create_stream(addr, socket).await?;
self.adapt_stream(stream, addr, ConnectionSide::Initiator, guard)
.await
}
pub async fn disconnect(&self, addr: SocketAddr) -> bool {
self.disconnect_w_origin(addr, DisconnectOrigin::User, None)
.await
}
pub(crate) async fn disconnect_w_origin(
&self,
addr: SocketAddr,
origin: DisconnectOrigin,
conn_id: Option<u64>,
) -> bool {
let Some(finalizer) = self.claim_disconnect(addr, conn_id) else {
return false;
};
let conn_span = create_connection_span(addr, self.span());
debug!(parent: &conn_span, "disconnecting (origin: {origin:?})...");
if let Some(handler) = self.protocols.on_disconnect.get() {
trace!(parent: &conn_span, "executing OnDisconnect logic...");
let (sender, receiver) = oneshot::channel();
handler.trigger(((addr, origin), sender)).await;
if let Some((handle, waiter)) =
crate::protocols::await_handler_response(receiver, handler.closed()).await
{
if let Some(conn) = self.connections.active.write().get_mut(&addr) {
conn.tasks.push(handle);
} else {
debug_assert!(
false,
"disconnect of {addr} claimed, yet the connection vanished before OnDisconnect registration"
);
handle.abort();
}
let _ = waiter.await;
}
}
drop(finalizer);
debug!(parent: &conn_span, "fully disconnected");
true
}
fn claim_disconnect(
&self,
addr: SocketAddr,
conn_id: Option<u64>,
) -> Option<DisconnectFinalizer<'_>> {
let active = self.connections.active.read();
let conn = active.get(&addr)?;
if conn_id.is_some_and(|id| conn.id != id) {
return None;
}
if conn.disconnecting.swap(true, AcqRel) {
return None;
}
Some(DisconnectFinalizer { node: self, addr })
}
pub fn connected_addrs(&self) -> Vec<SocketAddr> {
self.connections.addrs()
}
pub fn is_connected(&self, addr: SocketAddr) -> bool {
self.connections.is_connected(addr)
}
pub fn is_connecting(&self, addr: SocketAddr) -> bool {
self.connections.limits.lock().connecting.contains(&addr)
}
pub fn num_connected(&self) -> usize {
self.connections.num_connected()
}
pub fn num_connecting(&self) -> usize {
self.connections.limits.lock().connecting.len()
}
pub fn connection_info(&self, addr: SocketAddr) -> Option<ConnectionInfo> {
self.connections.get_info(addr)
}
pub fn connection_infos(&self) -> HashMap<SocketAddr, ConnectionInfo> {
self.connections.infos()
}
fn check_and_reserve(&self, addr: SocketAddr) -> io::Result<ConnectionGuard<'_>> {
let mut limits = self.connections.limits.lock();
let active = self.connections.active.read();
if active.contains_key(&addr) {
return Err(io::Error::new(
ErrorKind::AlreadyExists,
"already connected",
));
}
let num_ip_conns = limits.ip_count(addr);
let per_ip_limit = self.config.max_connections_per_ip as usize;
if num_ip_conns >= per_ip_limit {
return Err(io::Error::new(
ErrorKind::QuotaExceeded,
format!("maximum number ({per_ip_limit}) of per-IP connections reached"),
));
}
let num_connecting = limits.connecting.len();
let connecting_limit = self.config.max_connecting as usize;
if num_connecting >= connecting_limit {
return Err(io::Error::new(
ErrorKind::QuotaExceeded,
format!("maximum number ({connecting_limit}) of pending connections reached"),
));
}
let num_connected = active.len();
let connection_limit = self.config.max_connections as usize;
if num_connected + num_connecting >= connection_limit {
return Err(io::Error::new(
ErrorKind::QuotaExceeded,
format!("maximum number ({connection_limit}) of connections reached"),
));
}
if limits.connecting.contains(&addr) {
return Err(io::Error::new(
ErrorKind::AlreadyExists,
"already connecting",
));
}
limits.reserve(addr);
Ok(ConnectionGuard {
addr,
connections: &self.connections,
completed: false,
})
}
pub async fn shut_down(&self) {
self.shutdown.begin();
debug!(parent: self.span(), "shutting down");
let mut tasks: HashMap<_, _> = std::mem::take(&mut *self.tasks.lock())
.into_iter()
.map(|(kind, handle)| (kind, UnregisteredTask::new(handle)))
.collect();
drop(tasks.remove(&NodeTask::Listener));
let mut disconnects: FuturesUnordered<_> = self
.connected_addrs()
.into_iter()
.map(|addr| self.disconnect_w_origin(addr, DisconnectOrigin::Shutdown, None))
.collect();
while disconnects.next().await.is_some() {}
wait_for_drain(&self.connections.drain_notify, || {
self.connections.active.read().is_empty()
})
.await;
self.shutdown.wait_sched_drained().await;
self.shutdown.signal_handlers();
const DRAIN_DEADLINE: Duration = Duration::from_secs(3);
let deadline = Instant::now() + DRAIN_DEADLINE;
for kind in [
NodeTask::Handshake,
NodeTask::Reading,
NodeTask::Writing,
NodeTask::OnConnect,
NodeTask::OnDisconnect,
] {
if let Some(task) = tasks.get_mut(&kind) {
let remaining = deadline.saturating_duration_since(Instant::now());
let _ = timeout(remaining, task.handle_mut()).await;
}
}
drop(tasks);
*self.listening_addr.write().await = None;
}
pub(crate) fn register_task(&self, kind: NodeTask, handle: JoinHandle<()>) -> io::Result<()> {
let mut tasks = self.tasks.lock();
if self.shutdown.is_underway() {
handle.abort();
return Err(shutting_down_error());
}
let prev = tasks.insert(kind, handle);
debug_assert!(prev.is_none(), "a Node task was registered more than once");
Ok(())
}
}
pub(crate) struct UnregisteredTask(Option<JoinHandle<()>>);
impl UnregisteredTask {
pub(crate) fn new(handle: JoinHandle<()>) -> Self {
Self(Some(handle))
}
pub(crate) fn register(mut self, node: &Node, kind: NodeTask) -> io::Result<()> {
node.register_task(kind, self.0.take().unwrap()) }
fn handle_mut(&mut self) -> &mut JoinHandle<()> {
self.0.as_mut().unwrap() }
}
impl Drop for UnregisteredTask {
fn drop(&mut self) {
if let Some(handle) = &self.0 {
handle.abort();
}
}
}
struct SchedulingGuard(Node);
impl SchedulingGuard {
fn new(node: Node) -> Self {
node.shutdown.sched_started();
Self(node)
}
}
impl Drop for SchedulingGuard {
fn drop(&mut self) {
self.0.shutdown.sched_finished();
}
}
struct DisconnectFinalizer<'a> {
node: &'a Node,
addr: SocketAddr,
}
impl Drop for DisconnectFinalizer<'_> {
fn drop(&mut self) {
if let Some(writing) = self.node.protocols.writing.get() {
writing.senders.write().remove(&self.addr);
}
let conn = {
let mut limits = self.node.connections.limits.lock();
let conn = self.node.connections.remove(self.addr);
limits.release_ip(self.addr);
conn
};
drop(conn);
}
}
fn create_span(node_name: &str) -> Span {
macro_rules! try_span {
($lvl:expr) => {
let s = span!($lvl, "node", name = node_name);
if !s.is_disabled() {
return s;
}
};
}
try_span!(Level::TRACE);
try_span!(Level::DEBUG);
try_span!(Level::INFO);
try_span!(Level::WARN);
error_span!("node", name = node_name)
}
#[cfg(test)]
mod config_tests {
use super::*;
#[test]
#[should_panic(expected = "Config::max_connections must not be 0")]
fn zero_max_connections_is_rejected() {
let _ = Node::new(Config {
max_connections: 0,
..Default::default()
});
}
#[test]
#[should_panic(expected = "Config::max_connections_per_ip must not be 0")]
fn zero_max_connections_per_ip_is_rejected() {
let _ = Node::new(Config {
max_connections_per_ip: 0,
..Default::default()
});
}
#[test]
#[should_panic(expected = "Config::max_connecting must not be 0")]
fn zero_max_connecting_is_rejected() {
let _ = Node::new(Config {
max_connecting: 0,
..Default::default()
});
}
#[test]
fn excessive_max_connecting_is_clamped() {
let node = Node::new(Config {
max_connections: 10,
max_connecting: 11,
..Default::default()
});
assert_eq!(node.config().max_connecting, 10);
}
}
#[cfg(test)]
mod shutdown_tests {
use std::time::Instant;
use super::*;
use crate::{Pea2Pea, protocols::OnDisconnect};
#[tokio::test]
async fn cancelled_shut_down_aborts_node_tasks() {
#[derive(Clone)]
struct SlowDisconnect(Node);
impl Pea2Pea for SlowDisconnect {
fn node(&self) -> &Node {
&self.0
}
}
impl OnDisconnect for SlowDisconnect {
async fn on_disconnect(&self, _: SocketAddr, _: DisconnectOrigin) {
sleep(Duration::from_millis(300)).await;
}
}
let slow = SlowDisconnect(Node::new(Config {
listener_addr: Some("127.0.0.1:0".parse().unwrap()),
..Default::default()
}));
slow.enable_on_disconnect().await;
let slow_addr = slow.0.toggle_listener().await.unwrap().unwrap();
let peer = Node::new(Default::default());
peer.connect(slow_addr).await.unwrap();
let deadline = Instant::now() + Duration::from_secs(2);
while slow.0.num_connected() != 1 {
assert!(
Instant::now() < deadline,
"the test connection was never registered",
);
sleep(Duration::from_millis(10)).await;
}
assert!(
timeout(Duration::from_millis(50), slow.0.shut_down())
.await
.is_err()
);
timeout(Duration::from_secs(2), slow.0.shut_down())
.await
.expect("a repeated shut_down hung");
let node = slow.0.clone();
drop(slow);
let deadline = Instant::now() + Duration::from_secs(2);
while Arc::strong_count(&node.0) > 1 {
assert!(
Instant::now() < deadline,
"the node is still referenced, most likely by leaked tasks",
);
sleep(Duration::from_millis(10)).await;
}
peer.shut_down().await;
}
}
#[cfg(test)]
mod budget_tests {
use super::*;
#[tokio::test]
async fn exhausted_connecting_budget_is_signalled() {
let config = Config {
max_connecting: 1,
listener_addr: None, ..Default::default()
};
let node = Node::new(config);
let _held = node.connecting_permits.clone().try_acquire_owned().unwrap();
let err = node
.connect("127.0.0.1:9".parse().unwrap())
.await
.unwrap_err();
assert_eq!(err.kind(), ErrorKind::QuotaExceeded);
assert_eq!(node.heuristics().connect_budget_rejections(), 1);
}
#[tokio::test]
async fn idle_listener_does_not_hold_the_connect_budget() {
let node = Node::new(Config {
max_connecting: 1,
listener_addr: Some("127.0.0.1:0".parse().unwrap()),
..Default::default()
});
node.toggle_listener().await.unwrap();
let peer = Node::new(Config {
listener_addr: Some("127.0.0.1:0".parse().unwrap()),
..Default::default()
});
let peer_addr = peer.toggle_listener().await.unwrap().unwrap();
sleep(Duration::from_millis(50)).await;
node.connect(peer_addr).await.unwrap();
assert_eq!(node.heuristics().connect_budget_rejections(), 0);
}
}