diff --git a/Cargo.toml b/Cargo.toml index c81eb39..ee67da4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -16,6 +16,7 @@ path = "src/lib.rs" [dependencies] tokio = { version = "1.35", features = ["rt", "rt-multi-thread", "net", "sync", "time", "macros"] } +tokio-util = { version = "0.7" } serde = { version = "1.0", features = ["derive"] } serde_bencode = "0.2" sha1 = "0.10" diff --git a/src/error.rs b/src/error.rs index 5505397..7bcbfb7 100644 --- a/src/error.rs +++ b/src/error.rs @@ -19,6 +19,9 @@ pub enum DHTError { #[error("元数据验证失败")] InvalidMetadata, + + #[error("{0}")] + Other(String), } pub type Result = std::result::Result; diff --git a/src/scheduler.rs b/src/scheduler.rs index d83e3f5..cee0886 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -4,6 +4,7 @@ use crate::metadata::RbitFetcher; use std::sync::{Arc, RwLock}; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; use tokio::sync::{mpsc, Mutex}; +use tokio_util::sync::CancellationToken; #[cfg(debug_assertions)] use std::time::Duration; @@ -21,6 +22,7 @@ pub struct MetadataScheduler { total_dropped: Arc, total_dispatched: Arc, queue_len: Arc, + shutdown: CancellationToken, } impl MetadataScheduler { @@ -32,6 +34,7 @@ impl MetadataScheduler { callback: Arc>>, on_metadata_fetch: Arc>>, queue_len: Arc, + shutdown: CancellationToken, ) -> Self { Self { hash_rx, @@ -44,6 +47,7 @@ impl MetadataScheduler { total_dropped: Arc::new(AtomicU64::new(0)), total_dispatched: Arc::new(AtomicU64::new(0)), queue_len, + shutdown, } } @@ -63,6 +67,8 @@ impl MetadataScheduler { let (task_tx, task_rx) = mpsc::channel::(self.max_queue_size); let task_rx = Arc::new(Mutex::new(task_rx)); + let shutdown = self.shutdown.clone(); + #[cfg_attr(not(debug_assertions), allow(unused_variables))] for worker_id in 0..self.max_concurrent { let task_rx = task_rx.clone(); let fetcher = self.fetcher.clone(); @@ -70,36 +76,48 @@ impl MetadataScheduler { let on_metadata_fetch = self.on_metadata_fetch.clone(); let total_dispatched = self.total_dispatched.clone(); let queue_len = self.queue_len.clone(); + let shutdown_worker = shutdown.clone(); tokio::spawn(async move { #[cfg(debug_assertions)] log::trace!("Worker {} 启动", worker_id); loop { - let hash = { - let mut rx = task_rx.lock().await; - let h = rx.recv().await; - if h.is_some() { - queue_len.fetch_sub(1, Ordering::Relaxed); + tokio::select! { + _ = shutdown_worker.cancelled() => { + #[cfg(debug_assertions)] + log::trace!("Worker {} 收到关闭信号,退出", worker_id); + break; } - h - }; - - let hash = match hash { - Some(h) => h, - None => break, - }; - - total_dispatched.fetch_add(1, Ordering::Relaxed); - - Self::process_hash( - hash, - &fetcher, - &callback, - &on_metadata_fetch, - ).await; + result = async { + let mut rx = task_rx.lock().await; + rx.recv().await + } => { + let hash = { + if result.is_some() { + queue_len.fetch_sub(1, Ordering::Relaxed); + } + result + }; + + let hash = match hash { + Some(h) => h, + None => break, + }; + + total_dispatched.fetch_add(1, Ordering::Relaxed); + + Self::process_hash( + hash, + &fetcher, + &callback, + &on_metadata_fetch, + ).await; + } + } } + #[cfg(debug_assertions)] log::trace!("Worker {} 退出", worker_id); }); } @@ -109,10 +127,16 @@ impl MetadataScheduler { #[cfg(debug_assertions)] stats_interval.tick().await; + let shutdown = self.shutdown.clone(); loop { #[cfg(debug_assertions)] { tokio::select! { + _ = shutdown.cancelled() => { + #[cfg(debug_assertions)] + log::trace!("MetadataScheduler 主循环收到关闭信号,退出"); + break; + } Some(hash) = self.hash_rx.recv() => { self.total_received.fetch_add(1, Ordering::Relaxed); @@ -137,21 +161,30 @@ impl MetadataScheduler { #[cfg(not(debug_assertions))] { - match self.hash_rx.recv().await { - Some(hash) => { - self.total_received.fetch_add(1, Ordering::Relaxed); - - match task_tx.try_send(hash) { - Ok(_) => { - self.queue_len.fetch_add(1, Ordering::Relaxed); + tokio::select! { + _ = shutdown.cancelled() => { + #[cfg(debug_assertions)] + log::trace!("MetadataScheduler 主循环收到关闭信号,退出"); + break; + } + result = self.hash_rx.recv() => { + match result { + Some(hash) => { + self.total_received.fetch_add(1, Ordering::Relaxed); + + match task_tx.try_send(hash) { + Ok(_) => { + self.queue_len.fetch_add(1, Ordering::Relaxed); + } + Err(mpsc::error::TrySendError::Full(_)) => { + self.total_dropped.fetch_add(1, Ordering::Relaxed); + } + Err(_) => break, + } } - Err(mpsc::error::TrySendError::Full(_)) => { - self.total_dropped.fetch_add(1, Ordering::Relaxed); - } - Err(_) => break, + None => break, } } - None => break, } } } diff --git a/src/server.rs b/src/server.rs index 6c6ab1d..0fcf668 100644 --- a/src/server.rs +++ b/src/server.rs @@ -13,6 +13,7 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; use tokio::net::UdpSocket; use tokio::sync::{mpsc, Semaphore}; +use tokio_util::sync::CancellationToken; use socket2::{Socket, Domain, Type, Protocol}; use std::pin::Pin; use std::future::Future; @@ -52,6 +53,7 @@ pub struct DHTServer { hash_tx: mpsc::Sender, metadata_queue_len: Arc, max_metadata_queue_size: usize, + shutdown: CancellationToken, } impl DHTServer { @@ -131,6 +133,9 @@ impl DHTServer { let metadata_queue_len = Arc::new(AtomicUsize::new(0)); + let shutdown = CancellationToken::new(); + let shutdown_for_scheduler = shutdown.clone(); + let scheduler = MetadataScheduler::new( hash_rx, fetcher, @@ -139,6 +144,7 @@ impl DHTServer { callback.clone(), on_metadata_fetch.clone(), metadata_queue_len.clone(), + shutdown_for_scheduler, ); tokio::spawn(async move { @@ -159,6 +165,7 @@ impl DHTServer { hash_tx, metadata_queue_len, max_metadata_queue_size, + shutdown, }; Ok(server) @@ -218,17 +225,30 @@ impl DHTServer { } pub async fn start(&self) -> Result<()> { + // 检查是否已经被关闭 + if self.shutdown.is_cancelled() { + log::warn!("⚠️ 尝试启动已关闭的服务器"); + return Err(crate::error::DHTError::Other("服务器已关闭".to_string())); + } self.start_receiver(); self.bootstrap().await; let server = self.clone(); + let shutdown = self.shutdown.clone(); tokio::spawn(async move { let semaphore = Arc::new(Semaphore::new(2000)); let mut loop_tick = 0; loop { + // 检查关闭信号 + if shutdown.is_cancelled() { + #[cfg(debug_assertions)] + log::trace!("主循环收到关闭信号,退出"); + break; + } + let queue_len = server.metadata_queue_len.load(Ordering::Relaxed); let queue_pressure = queue_len as f64 / server.max_metadata_queue_size as f64; @@ -260,7 +280,10 @@ impl DHTServer { if nodes_batch.is_none() || loop_tick % 50 == 0 { server.bootstrap().await; if nodes_batch.is_none() { - tokio::time::sleep(sleep_duration).await; + tokio::select! { + _ = shutdown.cancelled() => break, + _ = tokio::time::sleep(sleep_duration) => {}, + } continue; } } @@ -295,18 +318,26 @@ impl DHTServer { } } - tokio::time::sleep(sleep_duration).await; + tokio::select! { + _ = shutdown.cancelled() => break, + _ = tokio::time::sleep(sleep_duration) => {}, + } } }); - - std::future::pending::<()>().await; + self.shutdown.cancelled().await; Ok(()) } + /// 显式关闭服务器,停止所有后台任务 + pub fn shutdown(&self) { + self.shutdown.cancel(); + } + fn start_receiver(&self) { let socket = self.socket.clone(); let socket_v6 = self.socket_v6.clone(); let server = self.clone(); + let shutdown = self.shutdown.clone(); let num_workers = std::thread::available_parallelism() .map(|n| n.get()) @@ -320,86 +351,92 @@ impl DHTServer { senders.push(tx); let server_clone = server.clone(); + let shutdown_worker = shutdown.clone(); tokio::spawn(async move { - while let Some((data, addr)) = rx.recv().await { - let _ = server_clone.handle_message(&data, addr).await; + loop { + tokio::select! { + _ = shutdown_worker.cancelled() => { + #[cfg(debug_assertions)] + log::trace!("Worker 收到关闭信号,退出"); + break; + } + msg = rx.recv() => { + match msg { + Some((data, addr)) => { + let _ = server_clone.handle_message(&data, addr).await; + } + None => break, + } + } + } } }); } - let senders_for_v6 = senders.clone(); + Self::spawn_udp_reader(socket, senders.clone(), shutdown.clone()); + + if let Some(socket_v6) = socket_v6 { + Self::spawn_udp_reader(socket_v6, senders, shutdown); + } + } + + fn spawn_udp_reader( + socket: Arc, + senders: Vec, SocketAddr)>>, + shutdown: CancellationToken, + ) { + let num_workers = senders.len(); + tokio::spawn(async move { let mut buf = [0u8; 65536]; let mut next_worker_idx = 0; loop { - match socket.recv_from(&mut buf).await { - Ok((size, addr)) => { - if size > 8192 { - #[cfg(debug_assertions)] - log::trace!("⚠️ 拒绝异常大的 UDP 包: {} 字节 from {}", size, addr); - continue; - } - - if size == 0 || buf[0] != b'd' { - continue; - } - - let data = buf[..size].to_vec(); - - let tx = &senders[next_worker_idx]; - next_worker_idx = (next_worker_idx + 1) % num_workers; - - match tx.try_send((data, addr)) { - Ok(_) => {}, - Err(mpsc::error::TrySendError::Full(_)) => { - #[cfg(debug_assertions)] - log::trace!("UDP worker queue full, dropping packet"); - }, - Err(_) => { break; } - } + tokio::select! { + _ = shutdown.cancelled() => { + #[cfg(debug_assertions)] + log::trace!("UDP 读取循环收到关闭信号,退出"); + break; } - Err(_e) => { - tokio::time::sleep(Duration::from_millis(1)).await; + result = socket.recv_from(&mut buf) => { + match result { + Ok((size, addr)) => { + if size > 8192 { + #[cfg(debug_assertions)] + log::trace!("⚠️ 拒绝异常大的 UDP 包: {} 字节 from {}", size, addr); + continue; + } + + if size == 0 || buf[0] != b'd' { + continue; + } + + let data = buf[..size].to_vec(); + + let tx = &senders[next_worker_idx]; + next_worker_idx = (next_worker_idx + 1) % num_workers; + + match tx.try_send((data, addr)) { + Ok(_) => {}, + Err(mpsc::error::TrySendError::Full(_)) => { + #[cfg(debug_assertions)] + log::trace!("UDP worker queue full, dropping packet"); + }, + Err(_) => { break; } + } + } + Err(_e) => { + tokio::select! { + _ = shutdown.cancelled() => break, + _ = tokio::time::sleep(Duration::from_millis(1)) => {}, + } + } + } } } } }); - - if let Some(socket_v6) = socket_v6 { - let senders_v6 = senders_for_v6; - tokio::spawn(async move { - let mut buf = [0u8; 65536]; - let mut next_worker_idx = 0; - - loop { - match socket_v6.recv_from(&mut buf).await { - Ok((size, addr)) => { - if size > 8192 { continue; } - if size == 0 || buf[0] != b'd' { continue; } - - let data = buf[..size].to_vec(); - - let tx = &senders_v6[next_worker_idx]; - next_worker_idx = (next_worker_idx + 1) % num_workers; - - match tx.try_send((data, addr)) { - Ok(_) => {}, - Err(mpsc::error::TrySendError::Full(_)) => { - #[cfg(debug_assertions)] - log::trace!("UDP worker queue full, dropping packet"); - }, - Err(_) => { break; } - } - } - Err(_e) => { - tokio::time::sleep(Duration::from_millis(1)).await; - } - } - } - }); - } } async fn handle_message(&self, data: &[u8], addr: SocketAddr) -> Result<()> {