修复优雅关闭
This commit is contained in:
@@ -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"
|
||||
|
||||
@@ -19,6 +19,9 @@ pub enum DHTError {
|
||||
|
||||
#[error("元数据验证失败")]
|
||||
InvalidMetadata,
|
||||
|
||||
#[error("{0}")]
|
||||
Other(String),
|
||||
}
|
||||
|
||||
pub type Result<T> = std::result::Result<T, DHTError>;
|
||||
|
||||
+66
-33
@@ -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
@@ -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<()> {
|
||||
|
||||
Reference in New Issue
Block a user