diff --git a/easytier-web/src/client_manager/listener_tests.rs b/easytier-web/src/client_manager/listener_tests.rs new file mode 100644 index 00000000..b4bac4dd --- /dev/null +++ b/easytier-web/src/client_manager/listener_tests.rs @@ -0,0 +1,164 @@ +use super::*; + +use easytier::proto::rpc::standalone::{ + RuntimeRpcListener, runtime_rpc_dialer, runtime_rpc_listener, +}; +use easytier_core::connectivity::protocol::raw::TunnelDialer; +use std::sync::atomic::AtomicUsize; +use tokio::{io::AsyncReadExt, net::TcpStream, sync::Notify, time::timeout}; + +#[derive(Debug)] +struct TestListener { + inner: RuntimeRpcListener, + accepted: Arc, + second_accept: Option>, +} + +#[async_trait::async_trait] +impl SocketListener for TestListener { + type Accepted = Box; + + async fn listen(&mut self) -> anyhow::Result<()> { + self.inner.listen().await + } + + async fn accept(&mut self) -> anyhow::Result { + let tunnel = self.inner.accept().await?; + if self.accepted.fetch_add(1, Ordering::SeqCst) == 1 + && let Some(ready) = &self.second_accept + { + // Model a listener that has accepted a socket but is still + // performing its transport upgrade when another handshake ends. + ready.notified().await; + } + Ok(tunnel) + } + + fn local_url(&self) -> url::Url { + self.inner.local_url() + } +} + +async fn wait_until(mut condition: impl FnMut() -> bool) { + timeout(Duration::from_secs(1), async { + while !condition() { + tokio::time::sleep(Duration::from_millis(1)).await; + } + }) + .await + .expect("listener did not make progress"); +} + +async fn manager() -> ClientManager { + ClientManager::new( + Db::memory_db().await, + None, + HeartbeatPolicy::default(), + Arc::new(FeatureFlags::default()), + Arc::new(crate::webhook::WebhookConfig::new( + None, None, None, None, None, + )), + ) +} + +#[tokio::test] +async fn slow_handshakes_do_not_block_other_clients() { + let mut manager = manager().await; + let url = manager + .add_listener(runtime_rpc_listener("127.0.0.1:0".parse().unwrap())) + .await + .unwrap(); + let addr = url.socket_addrs(|| None).unwrap()[0]; + let mut idle_clients = Vec::new(); + for _ in 0..4 { + idle_clients.push(TcpStream::connect(addr).await.unwrap()); + } + + // A client that sends its handshake must not wait for the preceding + // connections' three-second first-packet timeouts. + let client = runtime_rpc_dialer(url).connect().await.unwrap(); + let _client = timeout( + Duration::from_secs(1), + web_security::upgrade_client_tunnel(client), + ) + .await + .expect("healthy client was blocked by idle connections") + .unwrap(); + wait_until(|| manager.client_sessions.len() == 1).await; + + // Dropping the manager must also cancel handshakes that are still waiting. + drop(manager); + for mut client in idle_clients { + let mut byte = [0]; + assert_eq!( + timeout(Duration::from_secs(1), client.read(&mut byte)) + .await + .expect("pending handshake survived manager shutdown") + .unwrap(), + 0, + ); + } +} + +#[tokio::test] +#[ignore = "opens 4097 TCP connections; run explicitly with a sufficient file descriptor limit"] +async fn pending_handshakes_are_bounded_and_release_capacity() { + let mut manager = manager().await; + let accepted = Arc::new(AtomicUsize::new(0)); + let url = manager + .add_listener(TestListener { + inner: runtime_rpc_listener("127.0.0.1:0".parse().unwrap()), + accepted: accepted.clone(), + second_accept: None, + }) + .await + .unwrap(); + let addr = url.socket_addrs(|| None).unwrap()[0]; + let mut clients = Vec::new(); + for _ in 0..=MAX_PENDING_HANDSHAKES { + clients.push(TcpStream::connect(addr).await.unwrap()); + } + wait_until(|| accepted.load(Ordering::SeqCst) >= MAX_PENDING_HANDSHAKES).await; + tokio::time::sleep(Duration::from_millis(50)).await; + assert_eq!(accepted.load(Ordering::SeqCst), MAX_PENDING_HANDSHAKES); + + drop(clients.remove(0)); + wait_until(|| accepted.load(Ordering::SeqCst) == MAX_PENDING_HANDSHAKES + 1).await; +} + +#[tokio::test] +async fn completed_handshake_does_not_cancel_pending_accept() { + let mut manager = manager().await; + let accepted = Arc::new(AtomicUsize::new(0)); + let ready = Arc::new(Notify::new()); + let url = manager + .add_listener(TestListener { + inner: runtime_rpc_listener("127.0.0.1:0".parse().unwrap()), + accepted: accepted.clone(), + second_accept: Some(ready.clone()), + }) + .await + .unwrap(); + let dialer = runtime_rpc_dialer(url); + let first = dialer.connect().await.unwrap(); + let second = dialer.connect().await.unwrap(); + wait_until(|| accepted.load(Ordering::SeqCst) == 2).await; + + let _first = timeout( + Duration::from_secs(1), + web_security::upgrade_client_tunnel(first), + ) + .await + .unwrap() + .unwrap(); + wait_until(|| manager.client_sessions.len() == 1).await; + ready.notify_one(); + let _second = timeout( + Duration::from_secs(1), + web_security::upgrade_client_tunnel(second), + ) + .await + .expect("pending accept was cancelled when the first handshake completed") + .unwrap(); + wait_until(|| manager.client_sessions.len() == 2).await; +} diff --git a/easytier-web/src/client_manager/mod.rs b/easytier-web/src/client_manager/mod.rs index 866f3a1b..a9df1b17 100644 --- a/easytier-web/src/client_manager/mod.rs +++ b/easytier-web/src/client_manager/mod.rs @@ -1,3 +1,5 @@ +#[cfg(test)] +mod listener_tests; mod managed_config; mod runtime_reconcile; pub mod session; @@ -42,6 +44,7 @@ const MAX_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(60); const MIN_HEARTBEAT_TIMEOUT: Duration = Duration::from_secs(5); const MAX_HEARTBEAT_TIMEOUT: Duration = Duration::from_secs(120); const HEARTBEAT_TIMEOUT_MARGIN: Duration = Duration::from_secs(5); +const MAX_PENDING_HANDSHAKES: usize = 4096; #[derive(Debug, Clone, Copy)] pub(crate) struct HeartbeatPolicy { @@ -203,40 +206,57 @@ impl ClientManager { let feature_flags = self.feature_flags.clone(); let webhook_config = self.webhook_config.clone(); self.tasks.spawn(async move { - while let Ok(tunnel) = listener.accept().await { - let (tunnel, secure) = match web_security::accept_or_upgrade_server_tunnel( - tunnel, - ) - .await - { - Ok(v) => v, - Err(error) => { - tracing::warn!(%error, "failed to accept secure tunnel, dropping connection"); - continue; - } - }; - let info = tunnel.info().unwrap(); - let client_url: url::Url = info.remote_addr.unwrap().into(); - let location = Self::lookup_location(&client_url, geoip_db.clone()); - tracing::info!( - "New session from {:?}, secure: {}, location: {:?}", - client_url, - secure, - location - ); - let mut session = Session::new( - storage.clone(), - client_url.clone(), - location, - heartbeat_policy, - feature_flags.clone(), - webhook_config.clone(), - next_session_epoch.fetch_add(1, Ordering::Relaxed) + 1, - ); - session.serve(tunnel).await; - let session = Arc::new(session); - sessions.insert(client_url, session.clone()); - session.mark_route_ready(); + let mut handshakes = JoinSet::new(); + 'accept: loop { + // Some listeners include a WebSocket upgrade in accept(). Keep + // that future alive while processing completed handshakes. + let accepting = listener.accept(); + tokio::pin!(accepting); + loop { + let result = tokio::select! { + accepted = &mut accepting, if handshakes.len() < MAX_PENDING_HANDSHAKES => { + let Ok(tunnel) = accepted else { break 'accept }; + handshakes.spawn(web_security::accept_or_upgrade_server_tunnel(tunnel)); + break; + } + result = handshakes.join_next(), if !handshakes.is_empty() => { + result.unwrap() + } + }; + let (tunnel, secure) = match result { + Ok(Ok(value)) => value, + Ok(Err(error)) => { + tracing::warn!(%error, "failed to accept secure tunnel, dropping connection"); + continue; + } + Err(error) => { + tracing::warn!(%error, "secure tunnel handshake task failed"); + continue; + } + }; + let info = tunnel.info().unwrap(); + let client_url: url::Url = info.remote_addr.unwrap().into(); + let location = Self::lookup_location(&client_url, geoip_db.clone()); + tracing::info!( + "New session from {:?}, secure: {}, location: {:?}", + client_url, + secure, + location + ); + let mut session = Session::new( + storage.clone(), + client_url.clone(), + location, + heartbeat_policy, + feature_flags.clone(), + webhook_config.clone(), + next_session_epoch.fetch_add(1, Ordering::Relaxed) + 1, + ); + session.serve(tunnel).await; + let session = Arc::new(session); + sessions.insert(client_url, session.clone()); + session.mark_route_ready(); + } } listeners_cnt.fetch_sub(1, Ordering::Relaxed); });