修复优雅关闭

This commit is contained in:
桥下红药
2026-01-17 19:57:21 +08:00
parent f41c45a36b
commit 2b7542c09d
4 changed files with 175 additions and 101 deletions
+1
View File
@@ -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"
+3
View File
@@ -19,6 +19,9 @@ pub enum DHTError {
#[error("元数据验证失败")]
InvalidMetadata,
#[error("{0}")]
Other(String),
}
pub type Result<T> = std::result::Result<T, DHTError>;
+66 -33
View File
@@ -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<AtomicU64>,
total_dispatched: Arc<AtomicU64>,
queue_len: Arc<AtomicUsize>,
shutdown: CancellationToken,
}
impl MetadataScheduler {
@@ -32,6 +34,7 @@ impl MetadataScheduler {
callback: Arc<RwLock<Option<TorrentCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
queue_len: Arc<AtomicUsize>,
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::<HashDiscovered>(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,
}
}
}
+105 -68
View File
@@ -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<HashDiscovered>,
metadata_queue_len: Arc<AtomicUsize>,
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<UdpSocket>,
senders: Vec<mpsc::Sender<(Vec<u8>, 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<()> {