diff --git a/src/common.rs b/src/common.rs index 80fc73a..df1a7df 100644 --- a/src/common.rs +++ b/src/common.rs @@ -7,7 +7,7 @@ use sodiumoxide::crypto::sign; use std::{ io::prelude::*, io::Read, - net::{IpAddr, SocketAddr}, + net::{IpAddr, Ipv4Addr, SocketAddr}, time::{Instant, SystemTime}, }; @@ -34,6 +34,41 @@ pub async fn listen_tcp( } } +pub fn console_addr(bind_addr: Option, port: u16) -> Option { + let bind_addr = bind_addr?; + if bind_addr.is_unspecified() || bind_addr == IpAddr::V4(Ipv4Addr::LOCALHOST) { + return None; + } + Some(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), port)) +} + +// The runtime console (check_cmd) is reached via 127.0.0.1, so when the bind +// address does not already accept connections to 127.0.0.1 (it is neither the +// any-address nor 127.0.0.1 itself), the console gets a dedicated listener +// there; it is never bound to the external bind address. +pub async fn listen_console( + bind_addr: Option, + port: u16, +) -> ResultType> { + match console_addr(bind_addr, port) { + Some(addr) => { + let listener = hbb_common::tcp::new_listener(addr, true).await?; + log::info!("Listening on tcp {} for the console", addr); + Ok(Some(listener)) + } + None => Ok(None), + } +} + +pub async fn accept_or_pending( + listener: Option<&hbb_common::tokio::net::TcpListener>, +) -> std::io::Result<(hbb_common::tokio::net::TcpStream, SocketAddr)> { + match listener { + Some(listener) => listener.accept().await, + None => std::future::pending().await, + } +} + #[allow(dead_code)] pub(crate) fn get_expired_time() -> Instant { let now = Instant::now(); @@ -325,4 +360,43 @@ mod tests { let listener = listen_tcp(Some(bind_addr), 0).await.unwrap(); assert_eq!(listener.local_addr().unwrap().ip(), bind_addr); } + + #[test] + fn console_addr_only_when_bind_does_not_cover_ipv4_localhost() { + for bind_addr in [ + None, + Some(IpAddr::V4(Ipv4Addr::UNSPECIFIED)), + Some(IpAddr::V6(Ipv6Addr::UNSPECIFIED)), + Some(IpAddr::V4(Ipv4Addr::LOCALHOST)), + ] { + assert_eq!(console_addr(bind_addr, 21117), None); + } + for bind_addr in [ + Some(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1))), + Some("2001:db8::1".parse().unwrap()), + Some(IpAddr::V6(Ipv6Addr::LOCALHOST)), + ] { + assert_eq!( + console_addr(bind_addr, 21117), + Some(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 21117)) + ); + } + } + + #[hbb_common::tokio::test] + async fn console_listener_binds_ipv4_localhost() { + let listener = listen_console(Some(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 1))), 0) + .await + .unwrap() + .unwrap(); + assert_eq!( + listener.local_addr().unwrap().ip(), + IpAddr::V4(Ipv4Addr::LOCALHOST) + ); + assert!(listen_console(None, 0).await.unwrap().is_none()); + assert!(listen_console(Some(IpAddr::V4(Ipv4Addr::LOCALHOST)), 0) + .await + .unwrap() + .is_none()); + } } diff --git a/src/relay_server.rs b/src/relay_server.rs index 70c3d3c..de1a7ea 100644 --- a/src/relay_server.rs +++ b/src/relay_server.rs @@ -92,6 +92,7 @@ pub async fn start_with_bind( io_loop( crate::common::listen_tcp(bind_addr, port).await?, crate::common::listen_tcp(bind_addr, port2).await?, + crate::common::listen_console(bind_addr, port).await?, &key, ) .await; @@ -332,7 +333,12 @@ async fn check_cmd(cmd: &str, limiter: Limiter) -> String { res } -async fn io_loop(listener: TcpListener, listener2: TcpListener, key: &str) { +async fn io_loop( + listener: TcpListener, + listener2: TcpListener, + listener_console: Option, + key: &str, +) { check_params(); let limiter = ::new(TOTAL_BANDWIDTH.load(Ordering::SeqCst) as _); loop { @@ -361,6 +367,18 @@ async fn io_loop(listener: TcpListener, listener2: TcpListener, key: &str) { } } } + res = crate::common::accept_or_pending(listener_console.as_ref()) => { + match res { + Ok((stream, addr)) => { + stream.set_nodelay(true).ok(); + handle_connection(stream, addr, &limiter, key, false).await; + } + Err(err) => { + log::error!("console listener.accept failed: {}", err); + break; + } + } + } } } } diff --git a/src/rendezvous_server.rs b/src/rendezvous_server.rs index b194ae0..eaf7190 100644 --- a/src/rendezvous_server.rs +++ b/src/rendezvous_server.rs @@ -95,6 +95,7 @@ enum LoopFailure { Listener3, Listener2, Listener, + ConsoleListener, } impl RendezvousServer { @@ -157,6 +158,7 @@ impl RendezvousServer { let mut listener = create_tcp_listener(bind_addr, port).await?; let mut listener2 = create_tcp_listener(bind_addr, nat_port).await?; let mut listener3 = create_tcp_listener(bind_addr, ws_port).await?; + let mut listener_console = listen_console(bind_addr, nat_port as _).await?; log::info!("Listening on tcp/udp {}", listener.local_addr()?); log::info!( "Listening on tcp {}, extra port for NAT test", @@ -206,6 +208,7 @@ impl RendezvousServer { &mut listener, &mut listener2, &mut listener3, + &mut listener_console, &mut socket, &key, ) @@ -223,6 +226,10 @@ impl RendezvousServer { drop(listener2); listener2 = create_tcp_listener(bind_addr, nat_port).await?; } + LoopFailure::ConsoleListener => { + drop(listener_console.take()); + listener_console = listen_console(bind_addr, nat_port as _).await?; + } LoopFailure::Listener3 => { drop(listener3); listener3 = create_tcp_listener(bind_addr, ws_port).await?; @@ -243,6 +250,7 @@ impl RendezvousServer { listener: &mut TcpListener, listener2: &mut TcpListener, listener3: &mut TcpListener, + listener_console: &mut Option, socket: &mut FramedSocket, key: &str, ) -> LoopFailure { @@ -294,6 +302,18 @@ impl RendezvousServer { } } } + res = accept_or_pending(listener_console.as_ref()) => { + match res { + Ok((stream, addr)) => { + stream.set_nodelay(true).ok(); + self.handle_listener2(stream, addr).await; + } + Err(err) => { + log::error!("console listener.accept failed: {}", err); + return LoopFailure::ConsoleListener; + } + } + } res = listener3.accept() => { match res { Ok((stream, addr)) => {