From d049afa6d07e01fecfcc9d94c2a53561eaafd967 Mon Sep 17 00:00:00 2001 From: lxien Date: Sat, 5 Sep 2026 21:54:04 +0800 Subject: [PATCH] refactor: reap JoinSet tasks on client/server to avoid growth and shutdown self-join --- client/src/control/session.rs | 31 ++++++++++++- server/src/control/session/data_pool.rs | 8 +++- server/src/control/session/mod.rs | 58 ++++++++++++++++++++----- 3 files changed, 82 insertions(+), 15 deletions(-) diff --git a/client/src/control/session.rs b/client/src/control/session.rs index ec1c9c9..3b59e0e 100644 --- a/client/src/control/session.rs +++ b/client/src/control/session.rs @@ -381,6 +381,24 @@ impl Control { } async fn reader_loop(self: Arc) -> Result { + let mut data_tasks = { + let mut slot = self.data_tasks.lock().await; + std::mem::take(&mut *slot) + }; + + let end = self.reader_loop_inner(&mut data_tasks).await; + + { + let mut slot = self.data_tasks.lock().await; + *slot = data_tasks; + } + end + } + + async fn reader_loop_inner( + self: &Arc, + data_tasks: &mut JoinSet<()>, + ) -> Result { loop { if self.cancel.is_cancelled() { return Ok(ReaderEnd::Closed); @@ -390,6 +408,14 @@ impl Control { _ = self.cancel.cancelled() => { return Ok(ReaderEnd::Closed); } + joined = data_tasks.join_next(), if !data_tasks.is_empty() => { + if let Some(Err(e)) = joined { + if !e.is_cancelled() { + tracing::debug!(error = %e, "data task join error"); + } + } + continue; + } msg = async { let mut reader = self.reader.lock().await; msg::read_msg(&mut *reader).await @@ -407,9 +433,10 @@ impl Control { return Ok(ReaderEnd::Kicked(k.reason)); } Message::ReqDataConn(_) => { - let ctl = Arc::clone(&self); + while data_tasks.try_join_next().is_some() {} + let ctl = Arc::clone(self); let cancel = self.cancel.clone(); - self.data_tasks.lock().await.spawn(async move { + data_tasks.spawn(async move { tokio::select! { _ = cancel.cancelled() => {} res = ctl.handle_req_data_conn() => { diff --git a/server/src/control/session/data_pool.rs b/server/src/control/session/data_pool.rs index c5b812c..dfcc8a7 100644 --- a/server/src/control/session/data_pool.rs +++ b/server/src/control/session/data_pool.rs @@ -19,13 +19,17 @@ impl Control { } async fn spawn_refill(self: &Arc) { + if self.closed.load(Ordering::SeqCst) { + return; + } let ctl = Arc::clone(self); - self.bg_tasks.lock().await.spawn(async move { + self.spawn_bg(async move { if ctl.closed.load(Ordering::SeqCst) { return; } let _ = ctl.request_data_conn().await; - }); + }) + .await; } pub async fn get_data_conn(self: &Arc) -> Result { diff --git a/server/src/control/session/mod.rs b/server/src/control/session/mod.rs index 650ba48..b248ade 100644 --- a/server/src/control/session/mod.rs +++ b/server/src/control/session/mod.rs @@ -46,6 +46,7 @@ pub struct Control { udp_ports: Arc, bg_tasks: Mutex>, closed: AtomicBool, + cleaning: AtomicBool, finished: watch::Sender, activated: AtomicBool, pool_count: usize, @@ -105,6 +106,7 @@ impl Control { udp_ports, bg_tasks: Mutex::new(JoinSet::new()), closed: AtomicBool::new(false), + cleaning: AtomicBool::new(false), finished, activated: AtomicBool::new(false), pool_count: pool_count.max(1), @@ -223,7 +225,7 @@ impl Control { let timeout = self.effective_ping_timeout(); if timeout > 0 { let this = Arc::clone(&self); - self.bg_tasks.lock().await.spawn(async move { + self.spawn_bg(async move { loop { if this.closed.load(Ordering::SeqCst) { break; @@ -241,11 +243,12 @@ impl Control { timeout_secs = timeout, "heartbeat timeout" ); - this.shutdown().await; + this.signal_close(); break; } } - }); + }) + .await; } } @@ -253,6 +256,7 @@ impl Control { if self.closed.load(Ordering::SeqCst) { break; } + self.reap_bg_tasks().await; let msg = tokio::select! { _ = self.shutdown_notify.notified() => { break; @@ -285,15 +289,35 @@ impl Control { Ok(()) } + pub(super) fn signal_close(&self) { + self.closed.store(true, Ordering::SeqCst); + self.shutdown_notify.notify_waiters(); + self.data_notify.notify_waiters(); + } + + pub(super) async fn spawn_bg(&self, fut: F) + where + F: std::future::Future + Send + 'static, + { + let mut bg = self.bg_tasks.lock().await; + if self.cleaning.load(Ordering::SeqCst) || self.closed.load(Ordering::SeqCst) { + return; + } + reap_join_set(&mut bg); + bg.spawn(fut); + } + + async fn reap_bg_tasks(&self) { + let mut bg = self.bg_tasks.lock().await; + reap_join_set(&mut bg); + } + pub async fn shutdown(&self) { - if self.closed.swap(true, Ordering::SeqCst) { - self.shutdown_notify.notify_waiters(); - self.data_notify.notify_waiters(); + self.signal_close(); + if self.cleaning.swap(true, Ordering::SeqCst) { self.wait_finished().await; return; } - self.shutdown_notify.notify_waiters(); - self.data_notify.notify_waiters(); { let mut tm = self.tunnels.lock().await; for (name, detached) in tm.close_all().await { @@ -305,9 +329,11 @@ impl Control { let mut writer = self.writer.lock().await; let _ = writer.shutdown().await; } - let mut bg = self.bg_tasks.lock().await; - bg.abort_all(); - while bg.join_next().await.is_some() {} + { + let mut bg = self.bg_tasks.lock().await; + bg.abort_all(); + while bg.join_next().await.is_some() {} + } self.mark_finished(); } @@ -360,6 +386,16 @@ impl Control { } } +fn reap_join_set(bg: &mut JoinSet<()>) { + while let Some(res) = bg.try_join_next() { + if let Err(e) = res { + if !e.is_cancelled() { + tracing::debug!(error = %e, "bg task join error"); + } + } + } +} + impl Drop for Control { fn drop(&mut self) { if *self.finished.borrow() {