1use std::{
17 collections::{HashMap, HashSet},
18 fmt,
19 io,
20 net::{IpAddr, SocketAddr},
21 ops::Deref,
22 sync::{
23 Arc,
24 atomic::{AtomicUsize, Ordering::*},
25 },
26 time::{Duration, Instant},
27};
28
29use anyhow::anyhow;
30#[cfg(feature = "locktick")]
31use locktick::parking_lot::Mutex;
32use once_cell::sync::OnceCell;
33#[cfg(not(feature = "locktick"))]
34use parking_lot::Mutex;
35use tokio::{
36 io::split,
37 net::{TcpListener, TcpSocket, TcpStream},
38 sync::{OwnedSemaphorePermit, Semaphore, oneshot},
39 task::{JoinHandle, JoinSet},
40 time::timeout,
41};
42use tracing::*;
43
44use crate::{
45 BannedPeers,
46 Config,
47 Stats,
48 connections::{Connection, ConnectionSide, Connections, DisconnectOrigin, canonical_ip, create_connection_span},
49 protocols::{Protocol, Protocols},
50};
51
52static SEQUENTIAL_NODE_ID: AtomicUsize = AtomicUsize::new(0);
54
55#[derive(Clone)]
57pub struct Tcp(Arc<InnerTcp>);
58
59impl Deref for Tcp {
60 type Target = Arc<InnerTcp>;
61
62 fn deref(&self) -> &Self::Target {
63 &self.0
64 }
65}
66
67pub trait ApplicationError: Send + Sync + std::fmt::Debug + std::fmt::Display + 'static {}
69
70#[allow(missing_docs)]
72#[derive(thiserror::Error, Debug)]
73pub enum ConnectError {
74 #[error("already reached the maximum number of {limit} connections")]
75 MaximumConnectionsReached { limit: u16 },
76 #[error("already reached the maximum number of {limit} connections with IP '{ip}'")]
77 MaximumConnectionsPerIpReached { ip: IpAddr, limit: u16 },
78 #[error("already connecting to node at {address:?}")]
79 AlreadyConnecting { address: SocketAddr },
80 #[error("already connected to node at {address:?}")]
81 AlreadyConnected { address: SocketAddr },
82 #[error("attempt to self-connect (at address {address:?}")]
83 SelfConnect { address: SocketAddr },
84 #[error("rejected a connection attempt from a banned IP '{ip}'")]
85 BannedIp { ip: IpAddr },
86 #[error(transparent)]
88 IoError(std::io::Error),
89 #[error("{0}")]
92 ApplicationError(Box<dyn ApplicationError>),
93 #[error(transparent)]
97 Other(#[from] Box<dyn std::error::Error + Send + Sync>),
98}
99
100impl ConnectError {
101 pub fn application<E: ApplicationError>(err: E) -> Self {
103 Self::ApplicationError(Box::new(err))
104 }
105
106 pub fn other<E: Into<Box<dyn std::error::Error + Send + Sync>>>(err: E) -> Self {
108 Self::Other(err.into())
109 }
110}
111
112impl From<ConnectError> for std::io::Error {
113 fn from(err: ConnectError) -> Self {
114 match err {
115 ConnectError::IoError(err) => err,
116 ConnectError::Other(err) => std::io::Error::other(err),
117 err => std::io::Error::other(err.to_string()),
118 }
119 }
120}
121
122impl From<std::io::Error> for ConnectError {
123 fn from(err: std::io::Error) -> Self {
124 if err.kind() == std::io::ErrorKind::Other {
126 let inner = err.into_inner().unwrap_or_else(|| anyhow!("Unknown error").into());
128 ConnectError::other(inner)
129 } else {
130 ConnectError::IoError(err)
131 }
132 }
133}
134
135#[doc(hidden)]
136pub struct InnerTcp {
137 span: Span,
139 config: Config,
141 listening_addr: OnceCell<SocketAddr>,
143 pub(crate) protocols: Protocols,
145 connecting: Mutex<HashSet<SocketAddr>>,
147 pub(crate) connections: Connections,
149 banned_peers: BannedPeers,
151 stats: Stats,
153 pub(crate) tasks: Mutex<Vec<JoinHandle<()>>>,
155}
156
157impl Tcp {
158 pub fn new(mut config: Config) -> Self {
160 if config.name.is_none() {
162 config.name = Some(SEQUENTIAL_NODE_ID.fetch_add(1, Relaxed).to_string());
163 }
164
165 let span = crate::helpers::create_span(config.name.as_deref().unwrap());
167
168 assert!(config.max_connections_per_ip != 0, "Config::max_connections_per_ip must not be 0");
171
172 let tcp = Tcp(Arc::new(InnerTcp {
174 span,
175 config,
176 listening_addr: Default::default(),
177 protocols: Default::default(),
178 connecting: Default::default(),
179 connections: Default::default(),
180 banned_peers: Default::default(),
181 stats: Stats::new(Instant::now()),
182 tasks: Default::default(),
183 }));
184
185 debug!(parent: tcp.span(), "The node is ready");
186
187 tcp
188 }
189
190 pub fn uptime(&self) -> Duration {
192 self.stats.created().elapsed()
193 }
194
195 #[inline]
197 pub fn name(&self) -> &str {
198 self.config.name.as_deref().unwrap()
200 }
201
202 #[inline]
204 pub fn config(&self) -> &Config {
205 &self.config
206 }
207
208 pub fn listening_addr(&self) -> io::Result<SocketAddr> {
211 self.listening_addr.get().copied().ok_or_else(|| io::ErrorKind::AddrNotAvailable.into())
212 }
213
214 pub fn is_connected(&self, addr: SocketAddr) -> bool {
216 self.connections.is_connected(addr)
217 }
218
219 pub fn is_connecting(&self, addr: SocketAddr) -> bool {
221 self.connecting.lock().contains(&addr)
222 }
223
224 pub fn num_connected(&self) -> usize {
226 self.connections.num_connected()
227 }
228
229 pub fn num_connecting(&self) -> usize {
231 self.connecting.lock().len()
232 }
233
234 pub fn connected_addrs(&self) -> Vec<SocketAddr> {
236 self.connections.addrs()
237 }
238
239 pub fn connecting_addrs(&self) -> Vec<SocketAddr> {
241 self.connecting.lock().iter().copied().collect()
242 }
243
244 pub fn connection_stats(&self, addr: SocketAddr) -> Option<Arc<Stats>> {
246 self.connections.stats(addr)
247 }
248
249 pub fn connection_stats_snapshot(&self) -> HashMap<SocketAddr, Arc<Stats>> {
251 self.connections.stats_snapshot()
252 }
253
254 #[inline]
256 pub fn banned_peers(&self) -> &BannedPeers {
257 &self.banned_peers
258 }
259
260 #[inline]
262 pub fn stats(&self) -> &Stats {
263 &self.stats
264 }
265
266 #[inline]
268 pub fn span(&self) -> &Span {
269 &self.span
270 }
271
272 pub async fn shut_down(&self) {
274 debug!(parent: self.span(), "Shutting down the TCP stack");
275
276 let mut tasks = std::mem::take(&mut *self.tasks.lock()).into_iter();
278
279 if let Some(listening_task) = tasks.next() {
281 listening_task.abort(); }
283
284 let mut disconnect_tasks = JoinSet::new();
286 for addr in self.connected_addrs() {
287 let node = self.clone();
288 disconnect_tasks.spawn(async move {
289 node.disconnect_w_origin(addr, DisconnectOrigin::Shutdown).await;
290 });
291 }
292 while disconnect_tasks.join_next().await.is_some() {}
293
294 for handle in tasks {
296 handle.abort();
297 }
298 }
299}
300
301impl Tcp {
302 pub async fn connect(&self, addr: SocketAddr) -> Result<(), ConnectError> {
304 if let Ok(listening_addr) = self.listening_addr() {
305 if addr == listening_addr || self.is_self_connect(addr) {
307 error!(parent: self.span(), "Attempted to self-connect ({addr})");
308 return Err(ConnectError::SelfConnect { address: addr });
309 }
310 }
311
312 self.can_add_connection(addr)?;
313
314 if self.is_connected(addr) {
315 trace!(parent: self.span(), "Already connected to {addr}");
316 return Err(ConnectError::AlreadyConnected { address: addr });
317 }
318
319 if !self.connecting.lock().insert(addr) {
320 debug!(parent: self.span(), "Already connecting to {addr}");
321 return Err(ConnectError::AlreadyConnecting { address: addr });
322 }
323
324 let timeout_duration = Duration::from_millis(self.config().connection_timeout_ms.into());
325
326 let res = if let Some(listen_ip) = self.config().listener_ip {
329 timeout(timeout_duration, self.connect_with_specific_interface(listen_ip, addr)).await
330 } else {
331 timeout(timeout_duration, TcpStream::connect(addr)).await
332 };
333
334 let stream = match res {
335 Ok(Ok(stream)) => Ok(stream),
336 Ok(err) => {
337 self.connecting.lock().remove(&addr);
338 err
339 }
340 Err(err) => {
341 self.connecting.lock().remove(&addr);
342 error!("connection timeout error: {}", err);
343 Err(io::ErrorKind::TimedOut.into())
344 }
345 }?;
346
347 let ret = self.adapt_stream(stream, addr, ConnectionSide::Initiator).await;
348
349 if let Err(ref e) = ret {
350 self.connecting.lock().remove(&addr);
351 error!(parent: self.span(), "Unable to initiate a connection with {addr}: {e}");
352 }
353
354 ret.map_err(|err| err.into())
355 }
356
357 async fn connect_with_specific_interface(&self, listen_ip: IpAddr, addr: SocketAddr) -> io::Result<TcpStream> {
358 let sock = if listen_ip.is_ipv4() { TcpSocket::new_v4()? } else { TcpSocket::new_v6()? };
359 sock.bind(SocketAddr::new(listen_ip, 0))?;
361 sock.connect(addr).await
362 }
363
364 pub async fn disconnect(&self, addr: SocketAddr) -> bool {
368 self.disconnect_w_origin(addr, DisconnectOrigin::User).await
369 }
370
371 pub(crate) async fn disconnect_w_origin(&self, addr: SocketAddr, origin: DisconnectOrigin) -> bool {
372 if let Some(conn) = self.connections.0.read().get(&addr) {
374 if conn.disconnecting.swap(true, AcqRel) {
375 return false;
377 }
378 } else {
379 return false;
381 };
382
383 if let Some(handler) = self.protocols.disconnect.get() {
384 let (sender, receiver) = oneshot::channel();
385 handler.trigger(((addr, origin), sender)).await;
386 if let Ok((handle, waiter)) = receiver.await {
387 if let Some(conn) = self.connections.0.write().get_mut(&addr) {
390 conn.tasks.push(handle);
391 }
392 let _ = waiter.await;
394 }
395 }
396
397 let conn = self.connections.remove(addr);
398 let disconnected = conn.is_some();
399
400 if let Some(conn) = conn {
401 debug!(parent: self.span(), "Disconnecting from {addr}");
402
403 drop(conn);
405
406 debug!(parent: self.span(), "Disconnected from {addr}");
407 } else {
408 warn!(parent: self.span(), "Failed to disconnect, was not connected to {addr}");
409 }
410
411 disconnected
412 }
413}
414
415impl Tcp {
416 pub async fn enable_listener(&self) -> io::Result<SocketAddr> {
418 let listener_ip =
420 self.config().listener_ip.expect("Tcp::enable_listener was called, but Config::listener_ip is not set");
421
422 let listener = self.create_listener(listener_ip).await?;
424
425 let port = listener.local_addr()?.port();
427
428 let listening_addr = (listener_ip, port).into();
430 self.listening_addr.set(listening_addr).expect("The node's listener was started more than once");
431
432 let (tx, rx) = oneshot::channel();
434
435 let inbound_permits = Arc::new(Semaphore::new(self.config.max_connections as usize));
440
441 let tcp = self.clone();
442 let listening_task = tokio::spawn(async move {
443 trace!(parent: tcp.span(), "Spawned the listening task");
444 tx.send(()).unwrap(); loop {
447 let permit = match inbound_permits.clone().acquire_owned().await {
449 Ok(p) => p,
450 Err(_) => {
451 error!(parent: tcp.span(), "Inbound permit semaphore closed unexpectedly");
453 return;
454 }
455 };
456
457 match listener.accept().await {
459 Ok((stream, addr)) => tcp.handle_connection(stream, addr, permit),
460 Err(e) => {
461 drop(permit);
463
464 match e.kind() {
465 io::ErrorKind::ConnectionAborted | io::ErrorKind::ConnectionReset => {
467 debug!(parent: tcp.span(), "Transient accept error: {e}");
468 }
469 _ => {
472 error!(parent: tcp.span(), "Couldn't accept a connection: {e}");
473 tokio::time::sleep(Duration::from_millis(500)).await;
474 }
475 }
476 }
477 }
478 }
479 });
480 self.tasks.lock().push(listening_task);
481 let _ = rx.await;
482 debug!(parent: self.span(), "Listening on {listening_addr}");
483
484 Ok(listening_addr)
485 }
486
487 async fn create_listener(&self, listener_ip: IpAddr) -> io::Result<TcpListener> {
489 debug!("Creating a TCP listener on {listener_ip}...");
490 let listener = if let Some(port) = self.config().desired_listening_port {
491 let desired_listening_addr = SocketAddr::new(listener_ip, port);
493 match TcpListener::bind(desired_listening_addr).await {
495 Ok(listener) => listener,
496 Err(e) => {
497 if self.config().allow_random_port {
498 warn!(
499 parent: self.span(),
500 "Trying any listening port, as the desired port is unavailable: {e}"
501 );
502 let random_available_addr = SocketAddr::new(listener_ip, 0);
503 TcpListener::bind(random_available_addr).await?
504 } else {
505 error!(parent: self.span(), "The desired listening port is unavailable: {e}");
506 return Err(e);
507 }
508 }
509 }
510 } else if self.config().allow_random_port {
511 let random_available_addr = SocketAddr::new(listener_ip, 0);
512 TcpListener::bind(random_available_addr).await?
513 } else {
514 panic!("As 'listener_ip' is set, either 'desired_listening_port' or 'allow_random_port' must be set");
515 };
516
517 Ok(listener)
518 }
519
520 fn handle_connection(&self, stream: TcpStream, addr: SocketAddr, permit: OwnedSemaphorePermit) {
522 debug!(parent: self.span(), "Received a connection from {addr}");
523
524 if self.can_add_connection(addr).is_err() || self.is_self_connect(addr) {
525 debug!(parent: self.span(), "Rejecting the connection from {addr}");
526 return;
527 }
528
529 self.connecting.lock().insert(addr);
530
531 let tcp = self.clone();
532 tokio::spawn(async move {
533 let _permit = permit;
535
536 if let Err(e) = tcp.adapt_stream(stream, addr, ConnectionSide::Responder).await {
537 tcp.connecting.lock().remove(&addr);
538 error!(parent: tcp.span(), "Failed to connect with {addr}: {e}");
539 }
540 });
541 }
542
543 fn is_self_connect(&self, addr: SocketAddr) -> bool {
545 let listening_addr = self.listening_addr().unwrap();
547
548 match listening_addr.ip().is_loopback() {
549 true => listening_addr.port() == addr.port(),
552 false => listening_addr.ip() == addr.ip(),
554 }
555 }
556
557 fn can_add_connection(&self, addr: SocketAddr) -> Result<(), ConnectError> {
562 let num_connected = self.num_connected();
564 let limit = self.config.max_connections as usize;
566
567 if num_connected >= limit {
568 warn!(parent: self.span(), "Maximum number of active connections ({limit}) reached");
569 return Err(ConnectError::MaximumConnectionsReached { limit: self.config.max_connections });
570 } else if num_connected + self.num_connecting() >= limit {
571 warn!(parent: self.span(), "Maximum number of active & pending connections ({limit}) reached");
572 return Err(ConnectError::MaximumConnectionsReached { limit: self.config.max_connections });
573 }
574
575 let num_with_ip = self.num_connections_with_ip(addr);
577 let ip_limit = self.config.max_connections_per_ip as usize;
579
580 if num_with_ip >= ip_limit {
581 warn!(
582 parent: self.span(),
583 "Maximum number of connections ({ip_limit}) with IP '{}' reached", addr.ip(),
584 );
585 return Err(ConnectError::MaximumConnectionsPerIpReached {
586 ip: addr.ip(),
587 limit: self.config.max_connections_per_ip,
588 });
589 }
590
591 Ok(())
592 }
593
594 fn num_connections_with_ip(&self, addr: SocketAddr) -> usize {
600 let ip = canonical_ip(addr);
601
602 let num_connected = self.connections.num_with_ip(addr);
605 let num_connecting = self.connecting.lock().iter().filter(|addr| canonical_ip(**addr) == ip).count();
606
607 num_connected.saturating_add(num_connecting)
608 }
609
610 async fn adapt_stream(&self, stream: TcpStream, peer_addr: SocketAddr, own_side: ConnectionSide) -> io::Result<()> {
612 if own_side == ConnectionSide::Initiator {
614 if let Ok(addr) = stream.local_addr() {
615 debug!(
616 parent: self.span(), "establishing connection with {}; the peer is connected on port {}",
617 peer_addr, addr.port()
618 );
619 } else {
620 warn!(parent: self.span(), "couldn't determine the peer's port");
621 }
622 }
623
624 let conn_span = create_connection_span(peer_addr, self.span());
625 let connection = Connection::new(peer_addr, stream, !own_side, conn_span);
626
627 let mut connection = self.enable_protocols(connection).await?;
629
630 let conn_ready_tx = connection.readiness_notifier.take();
632
633 self.connections.add(connection);
634 self.connecting.lock().remove(&peer_addr);
635
636 if let Some(tx) = conn_ready_tx {
638 let _ = tx.send(());
639 }
640
641 if let Some(handler) = self.protocols.on_connect.get() {
643 let (sender, receiver) = oneshot::channel();
644 handler.trigger((peer_addr, sender)).await;
645 if let Ok(handle) = receiver.await {
647 if let Some(conn) = self.connections.0.write().get_mut(&peer_addr) {
649 conn.tasks.push(handle);
650 } else {
651 handle.abort();
653 }
654 }
655 }
656
657 Ok(())
658 }
659
660 async fn enable_protocols(&self, conn: Connection) -> io::Result<Connection> {
662 macro_rules! enable_protocol {
664 ($handler_type: ident, $node:expr, $conn: expr) => {
665 if let Some(handler) = $node.protocols.$handler_type.get() {
666 let (conn_returner, conn_retriever) = oneshot::channel();
667
668 handler.trigger(($conn, conn_returner)).await;
669
670 match conn_retriever.await {
671 Ok(Ok(conn)) => conn,
672 Err(_) => return Err(io::ErrorKind::BrokenPipe.into()),
673 Ok(e) => return e,
674 }
675 } else {
676 $conn
677 }
678 };
679 }
680
681 let mut conn = enable_protocol!(handshake, self, conn);
682
683 if let Some(stream) = conn.stream.take() {
685 let (reader, writer) = split(stream);
686 conn.reader = Some(Box::new(reader));
687 conn.writer = Some(Box::new(writer));
688 }
689
690 let conn = enable_protocol!(reading, self, conn);
691 let conn = enable_protocol!(writing, self, conn);
692
693 Ok(conn)
694 }
695}
696
697impl fmt::Debug for Tcp {
698 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
699 write!(f, "The TCP stack config: {:?}", self.config)
700 }
701}
702
703#[cfg(test)]
704mod tests {
705 use super::*;
706
707 use std::{
708 net::{IpAddr, Ipv4Addr},
709 str::FromStr,
710 };
711
712 #[tokio::test]
713 async fn test_new() {
714 let tcp = Tcp::new(Config {
715 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
716 max_connections: 200,
717 ..Default::default()
718 });
719
720 assert_eq!(tcp.config.max_connections, 200);
721 assert_eq!(tcp.config.listener_ip, Some(IpAddr::V4(Ipv4Addr::LOCALHOST)));
722 assert_eq!(tcp.enable_listener().await.unwrap().ip(), IpAddr::V4(Ipv4Addr::LOCALHOST));
723
724 assert_eq!(tcp.num_connected(), 0);
725 assert_eq!(tcp.num_connecting(), 0);
726 }
727
728 #[tokio::test]
729 async fn test_connect() {
730 let tcp = Tcp::new(Config::default());
731 let node_ip = tcp.enable_listener().await.unwrap();
732
733 let result = tcp.connect(node_ip).await;
735 assert!(matches!(result, Err(ConnectError::SelfConnect { .. })));
736
737 assert_eq!(tcp.num_connected(), 0);
738 assert_eq!(tcp.num_connecting(), 0);
739 assert!(!tcp.is_connected(node_ip));
740 assert!(!tcp.is_connecting(node_ip));
741
742 let peer = Tcp::new(Config {
744 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
745 desired_listening_port: Some(0),
746 max_connections: 1,
747 ..Default::default()
748 });
749 let peer_ip = peer.enable_listener().await.unwrap();
750
751 tcp.connect(peer_ip).await.unwrap();
753 assert_eq!(tcp.num_connected(), 1);
754 assert_eq!(tcp.num_connecting(), 0);
755 assert!(tcp.is_connected(peer_ip));
756 assert!(!tcp.is_connecting(peer_ip));
757 }
758
759 #[tokio::test]
760 async fn test_disconnect() {
761 let tcp = Tcp::new(Config::default());
762 let _node_ip = tcp.enable_listener().await.unwrap();
763
764 let peer = Tcp::new(Config {
766 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
767 desired_listening_port: Some(0),
768 max_connections: 1,
769 ..Default::default()
770 });
771 let peer_ip = peer.enable_listener().await.unwrap();
772
773 tcp.connect(peer_ip).await.unwrap();
775 assert_eq!(tcp.num_connected(), 1);
776 assert_eq!(tcp.num_connecting(), 0);
777 assert!(tcp.is_connected(peer_ip));
778 assert!(!tcp.is_connecting(peer_ip));
779
780 let has_disconnected = tcp.disconnect(peer_ip).await;
782 assert!(has_disconnected);
783 assert_eq!(tcp.num_connected(), 0);
784 assert_eq!(tcp.num_connecting(), 0);
785 assert!(!tcp.is_connected(peer_ip));
786 assert!(!tcp.is_connecting(peer_ip));
787
788 let has_disconnected = tcp.disconnect(peer_ip).await;
790 assert!(!has_disconnected);
791 assert_eq!(tcp.num_connected(), 0);
792 assert_eq!(tcp.num_connecting(), 0);
793 assert!(!tcp.is_connected(peer_ip));
794 assert!(!tcp.is_connecting(peer_ip));
795 }
796
797 #[tokio::test]
798 async fn test_can_add_connection() {
799 let tcp = Tcp::new(Config { max_connections: 1, ..Default::default() });
800
801 let peer = Tcp::new(Config {
803 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
804 desired_listening_port: Some(0),
805 max_connections: 1,
806 ..Default::default()
807 });
808 let peer_ip = peer.enable_listener().await.unwrap();
809
810 assert!(tcp.can_add_connection(peer_ip).is_ok());
811
812 let stream = TcpStream::connect(peer_ip).await.unwrap();
814 tcp.connections.add(Connection::new(peer_ip, stream, ConnectionSide::Initiator, Span::none()));
815 assert!(tcp.can_add_connection(peer_ip).is_err());
816
817 let another_ip = SocketAddr::from_str("1.2.3.4:4242").unwrap();
820 let result = tcp.connect(another_ip).await;
821 assert!(matches!(result, Err(ConnectError::MaximumConnectionsReached { .. })));
822
823 tcp.connections.remove(peer_ip);
825 assert!(tcp.can_add_connection(peer_ip).is_ok());
826
827 tcp.connecting.lock().insert(peer_ip);
829 assert!(tcp.can_add_connection(peer_ip).is_err());
830
831 let another_ip = SocketAddr::from_str("1.2.3.4:4242").unwrap();
833 let result = tcp.connect(another_ip).await;
834 assert!(matches!(result, Err(ConnectError::MaximumConnectionsReached { .. })));
835
836 tcp.connecting.lock().remove(&peer_ip);
838 assert!(tcp.can_add_connection(peer_ip).is_ok());
839
840 let stream = TcpStream::connect(peer_ip).await.unwrap();
842 tcp.connections.add(Connection::new(peer_ip, stream, ConnectionSide::Responder, Span::none()));
843 tcp.connecting.lock().insert(peer_ip);
844 assert!(tcp.can_add_connection(peer_ip).is_err());
845
846 tcp.connections.remove(peer_ip);
848 tcp.connecting.lock().remove(&peer_ip);
849 assert!(tcp.can_add_connection(peer_ip).is_ok());
850 }
851
852 #[tokio::test]
853 async fn test_max_connections_per_ip() {
854 let tcp = Tcp::new(Config { max_connections: 10, max_connections_per_ip: 2, ..Default::default() });
855
856 let peer = Tcp::new(Config {
859 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
860 desired_listening_port: Some(0),
861 ..Default::default()
862 });
863 let peer_ip = peer.enable_listener().await.unwrap();
864
865 let first = SocketAddr::from_str("1.2.3.4:1111").unwrap();
867 let second = SocketAddr::from_str("1.2.3.4:2222").unwrap();
868 let third = SocketAddr::from_str("1.2.3.4:3333").unwrap();
869 let other_ip = SocketAddr::from_str("5.6.7.8:1111").unwrap();
870
871 for addr in [first, second] {
873 assert!(tcp.can_add_connection(addr).is_ok());
874 let stream = TcpStream::connect(peer_ip).await.unwrap();
875 tcp.connections.add(Connection::new(addr, stream, ConnectionSide::Initiator, Span::none()));
876 }
877 assert_eq!(tcp.num_connections_with_ip(first), 2);
878 assert!(matches!(
879 tcp.can_add_connection(third),
880 Err(ConnectError::MaximumConnectionsPerIpReached { limit: 2, .. })
881 ));
882
883 assert!(tcp.can_add_connection(other_ip).is_ok());
885
886 tcp.connections.remove(second);
888 assert!(tcp.can_add_connection(third).is_ok());
889 tcp.connecting.lock().insert(second);
890 assert_eq!(tcp.num_connections_with_ip(first), 2);
891 assert!(matches!(
892 tcp.can_add_connection(third),
893 Err(ConnectError::MaximumConnectionsPerIpReached { limit: 2, .. })
894 ));
895 }
896
897 #[tokio::test]
898 async fn test_max_connections_per_ip_canonicalizes_mapped_addresses() {
899 let tcp = Tcp::new(Config { max_connections: 10, max_connections_per_ip: 2, ..Default::default() });
900
901 let peer = Tcp::new(Config {
902 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
903 desired_listening_port: Some(0),
904 ..Default::default()
905 });
906 let peer_ip = peer.enable_listener().await.unwrap();
907
908 let native = SocketAddr::from_str("1.2.3.4:1111").unwrap();
910 let mapped = SocketAddr::from_str("[::ffff:1.2.3.4]:2222").unwrap();
911 let also_mapped = SocketAddr::from_str("[::ffff:1.2.3.4]:3333").unwrap();
912
913 assert_ne!(native.ip(), mapped.ip());
915
916 for addr in [native, mapped] {
917 assert!(tcp.can_add_connection(addr).is_ok());
918 let stream = TcpStream::connect(peer_ip).await.unwrap();
919 tcp.connections.add(Connection::new(addr, stream, ConnectionSide::Initiator, Span::none()));
920 }
921
922 assert_eq!(tcp.num_connections_with_ip(native), 2);
924 assert_eq!(tcp.num_connections_with_ip(mapped), 2);
925 for addr in [native, also_mapped] {
926 assert!(matches!(
927 tcp.can_add_connection(addr),
928 Err(ConnectError::MaximumConnectionsPerIpReached { limit: 2, .. })
929 ));
930 }
931 }
932
933 #[tokio::test]
934 async fn test_handle_connection() {
935 let tcp = Tcp::new(Config {
936 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
937 max_connections: 1,
938 ..Default::default()
939 });
940
941 let peer1 = Tcp::new(Config {
943 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
944 desired_listening_port: Some(0),
945 max_connections: 1,
946 ..Default::default()
947 });
948 let peer1_ip = peer1.enable_listener().await.unwrap();
949
950 let stream = TcpStream::connect(peer1_ip).await.unwrap();
952 tcp.connections.add(Connection::new(peer1_ip, stream, ConnectionSide::Responder, Span::none()));
953 assert!(tcp.can_add_connection(peer1_ip).is_err());
954 assert_eq!(tcp.num_connected(), 1);
955 assert_eq!(tcp.num_connecting(), 0);
956 assert!(tcp.is_connected(peer1_ip));
957 assert!(!tcp.is_connecting(peer1_ip));
958
959 let peer2 = Tcp::new(Config {
961 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
962 desired_listening_port: Some(0),
963 max_connections: 1,
964 ..Default::default()
965 });
966 let peer2_ip = peer2.enable_listener().await.unwrap();
967
968 let stream = TcpStream::connect(peer2_ip).await.unwrap();
970 let inbound_permits = Arc::new(Semaphore::new(1));
971 let permit = inbound_permits.clone().acquire_owned().await.unwrap();
972 tcp.handle_connection(stream, peer2_ip, permit);
973 assert!(tcp.can_add_connection(peer1_ip).is_err());
974 assert_eq!(tcp.num_connected(), 1);
975 assert_eq!(tcp.num_connecting(), 0);
976 assert!(tcp.is_connected(peer1_ip));
977 assert!(!tcp.is_connected(peer2_ip));
978 assert!(!tcp.is_connecting(peer1_ip));
979 assert!(!tcp.is_connecting(peer2_ip));
980 }
981
982 #[tokio::test]
983 async fn test_adapt_stream() {
984 let tcp = Tcp::new(Config { max_connections: 1, ..Default::default() });
985
986 let peer = Tcp::new(Config {
988 listener_ip: Some(IpAddr::V4(Ipv4Addr::LOCALHOST)),
989 desired_listening_port: Some(0),
990 max_connections: 1,
991 ..Default::default()
992 });
993 let peer_ip = peer.enable_listener().await.unwrap();
994
995 tcp.connecting.lock().insert(peer_ip);
997 assert_eq!(tcp.num_connected(), 0);
998 assert_eq!(tcp.num_connecting(), 1);
999 assert!(!tcp.is_connected(peer_ip));
1000 assert!(tcp.is_connecting(peer_ip));
1001
1002 let stream = TcpStream::connect(peer_ip).await.unwrap();
1004 tcp.adapt_stream(stream, peer_ip, ConnectionSide::Responder).await.unwrap();
1005 assert_eq!(tcp.num_connected(), 1);
1006 assert_eq!(tcp.num_connecting(), 0);
1007 assert!(tcp.is_connected(peer_ip));
1008 assert!(!tcp.is_connecting(peer_ip));
1009 }
1010}