From 2ab9103e5bb6720b396cb6b3f0114bb09d2773d7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=A1=A5=E4=B8=8B=E7=BA=A2=E8=8D=AF?= Date: Thu, 12 Mar 2026 21:43:18 +0800 Subject: [PATCH] =?UTF-8?q?=E6=80=A7=E8=83=BD=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.toml | 2 +- examples/main.rs | 1 - src/error.rs | 15 ---- src/lib.rs | 2 +- src/metadata.rs | 25 ++++--- src/scheduler.rs | 190 ++++++++++++++++++++--------------------------- src/server.rs | 85 +++++++++++---------- src/sharded.rs | 45 ----------- src/types.rs | 3 - 9 files changed, 145 insertions(+), 223 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 12d105b..5a75305 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -28,10 +28,10 @@ thiserror = "1.0" socket2 = { version = "0.5", features = ["all"] } rbit = "0.2" bytes = "1.0" -bloomfilter = "1.0" ahash = "0.8" serde_bytes = "0.11.19" metrics = { version = "0.24", optional = true } +async-channel = "2.5.0" [dev-dependencies] tracing = "0.1" diff --git a/examples/main.rs b/examples/main.rs index dc7848a..d2056f2 100644 --- a/examples/main.rs +++ b/examples/main.rs @@ -32,7 +32,6 @@ async fn main() -> Result<()> { let options = DHTOptions { port: 12313, - auto_metadata: true, metadata_timeout: 3, // ✅ 快速超时,快速失败 max_metadata_queue_size: 100000, // ✅ 大缓冲区(防止饱和) max_metadata_worker_count: 1000, // ✅ 激进并发(最大化吞吐) diff --git a/src/error.rs b/src/error.rs index 0a8fc49..4599c8e 100644 --- a/src/error.rs +++ b/src/error.rs @@ -5,21 +5,6 @@ pub enum DHTError { #[error("网络错误: {0}")] Network(#[from] std::io::Error), - #[error("Bencode 解析错误: {0}")] - Bencode(String), - - #[error("协议错误: {0}")] - Protocol(String), - - #[error("超时")] - Timeout, - - #[error("未找到 Peer")] - NoPeersAvailable, - - #[error("元数据验证失败")] - InvalidMetadata, - #[error("{0}")] Other(String), } diff --git a/src/lib.rs b/src/lib.rs index 11baee9..b2303df 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -9,7 +9,7 @@ pub mod types; pub use error::{DHTError, Result}; pub use scheduler::MetadataScheduler; pub use server::{DHTServer, HashDiscovered}; -pub use sharded::{NodeTuple, ShardedBloom, ShardedNodeQueue}; +pub use sharded::{NodeTuple, ShardedNodeQueue}; pub use types::{DHTOptions, FileInfo, NetMode, TorrentInfo}; pub mod prelude { diff --git a/src/metadata.rs b/src/metadata.rs index 251be52..d3af2b6 100644 --- a/src/metadata.rs +++ b/src/metadata.rs @@ -29,7 +29,7 @@ impl RbitFetcher { &self, info_hash: &[u8; 20], peer_addr: SocketAddr, - ) -> Option<(String, u64, Vec)> { + ) -> Option<(String, u64, Vec, u64)> { #[cfg(feature = "metrics")] counter!("dht_metadata_fetch_attempts_total").increment(1); @@ -138,18 +138,21 @@ impl RbitFetcher { } if success { let info_hash_copy = *info_hash; - let full_data_clone = full_data.clone(); - let is_valid = tokio::task::spawn_blocking(move || { + let validated = tokio::task::spawn_blocking(move || { let mut hasher = Sha1::new(); - hasher.update(&full_data_clone); + hasher.update(&full_data); let digest: [u8; 20] = hasher.finalize().into(); - digest == info_hash_copy - }).await.unwrap_or(false); + if digest == info_hash_copy { + Some(full_data) + } else { + None + } + }).await.unwrap_or(None); - if is_valid { + if validated.is_some() { #[cfg(feature = "metrics")] counter!("dht_metadata_handshake_result_total", "result" => "success").increment(1); - return Some(full_data); + return validated; } #[cfg(feature = "metrics")] counter!("dht_metadata_fetch_fail_total", "reason" => "sha1_mismatch").increment(1); @@ -172,6 +175,10 @@ impl RbitFetcher { .and_then(|v| v.as_str()) .unwrap_or("Unknown") .to_string(); + let piece_length = dict + .get(&b"piece length"[..]) + .and_then(|v| v.as_integer()) + .unwrap_or(0) as u64; let mut total_size = 0; let mut file_list = Vec::new(); if let Some(files) = dict.get(&b"files"[..]).and_then(|v| v.as_list()) { @@ -212,7 +219,7 @@ impl RbitFetcher { counter!("dht_metadata_fetch_success_total").increment(1); histogram!("dht_metadata_size_bytes").record(total_size as f64); } - return Some((name, total_size, file_list)); + return Some((name, total_size, file_list, piece_length)); } } #[cfg(feature = "metrics")] diff --git a/src/scheduler.rs b/src/scheduler.rs index d56cb26..f2d9b3b 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -3,9 +3,7 @@ use crate::server::HashDiscovered; use crate::types::TorrentInfo; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; use std::sync::{Arc, RwLock}; -#[cfg(debug_assertions)] -use std::time::Duration; -use tokio::sync::{Mutex, mpsc}; +use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; type TorrentCallback = Arc; @@ -69,8 +67,7 @@ impl MetadataScheduler { } pub async fn run(mut self) { - let (task_tx, task_rx) = mpsc::channel::(self.max_queue_size); - let task_rx = Arc::new(Mutex::new(task_rx)); + let (task_tx, task_rx) = async_channel::bounded::(self.max_queue_size); let shutdown = self.shutdown.clone(); #[cfg_attr(not(debug_assertions), allow(unused_variables))] @@ -94,20 +91,13 @@ impl MetadataScheduler { log::trace!("Worker {} 收到关闭信号,退出", worker_id); break; } - result = async { - let mut rx = task_rx.lock().await; - rx.recv().await - } => { - let hash = { - if result.is_some() { + result = task_rx.recv() => { + let hash = match result { + Ok(h) => { queue_len.fetch_sub(1, Ordering::Relaxed); + h } - result - }; - - let hash = match hash { - Some(h) => h, - None => break, + Err(_) => break, }; total_dispatched.fetch_add(1, Ordering::Relaxed); @@ -127,74 +117,52 @@ impl MetadataScheduler { }); } - #[cfg(debug_assertions)] - let mut stats_interval = tokio::time::interval(Duration::from_secs(60)); - #[cfg(debug_assertions)] - stats_interval.tick().await; + let mut stats_interval = if cfg!(debug_assertions) { + Some(tokio::time::interval(std::time::Duration::from_secs(60))) + } else { + None + }; + if let Some(ref mut interval) = stats_interval { + 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); - - 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, - } - } - - _ = stats_interval.tick() => { - self.print_stats(&task_tx); - } - - else => break, + 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); - #[cfg(not(debug_assertions))] - { - 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, + match task_tx.try_send(hash) { + Ok(_) => { + self.queue_len.fetch_add(1, Ordering::Relaxed); } + Err(async_channel::TrySendError::Full(_)) => { + self.total_dropped.fetch_add(1, Ordering::Relaxed); + } + Err(_) => break, } - None => break, } + None => break, } } + _ = async { + match stats_interval.as_mut() { + Some(interval) => interval.tick().await, + None => std::future::pending().await, + } + } => { + self.print_stats_inline(); + } } } - // 显式关闭 task_tx,让所有 worker 任务能够退出 drop(task_tx); #[cfg(debug_assertions)] log::trace!("MetadataScheduler 主循环退出,等待 worker 任务完成"); @@ -231,7 +199,9 @@ impl MetadataScheduler { _ => return, }; - if let Some((name, total_size, files)) = fetcher.fetch(&info_hash_bytes, peer_addr).await { + if let Some((name, total_size, files, piece_length)) = + fetcher.fetch(&info_hash_bytes, peer_addr).await + { let metadata = TorrentInfo { info_hash, name, @@ -239,10 +209,10 @@ impl MetadataScheduler { files, magnet_link: format!("magnet:?xt=urn:btih:{}", hash.info_hash), peers: vec![peer_addr.to_string()], - piece_length: 0, + piece_length, timestamp: std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_secs(), }; @@ -259,43 +229,45 @@ impl MetadataScheduler { } } - #[cfg(debug_assertions)] - fn print_stats(&self, task_tx: &mpsc::Sender) { - let received = self.total_received.load(Ordering::Relaxed); - let dropped = self.total_dropped.load(Ordering::Relaxed); - let dispatched = self.total_dispatched.load(Ordering::Relaxed); + fn print_stats_inline(&self) { + #[cfg(debug_assertions)] + { + let received = self.total_received.load(Ordering::Relaxed); + let dropped = self.total_dropped.load(Ordering::Relaxed); + let dispatched = self.total_dispatched.load(Ordering::Relaxed); - let drop_rate = if received > 0 { - dropped as f64 / received as f64 * 100.0 - } else { - 0.0 - }; + let drop_rate = if received > 0 { + dropped as f64 / received as f64 * 100.0 + } else { + 0.0 + }; - let queue_size = self.max_queue_size - task_tx.capacity(); - let queue_pressure = (queue_size as f64 / self.max_queue_size as f64) * 100.0; + let queue_len = self.queue_len.load(Ordering::Relaxed); + let queue_pressure = (queue_len as f64 / self.max_queue_size as f64) * 100.0; - if queue_pressure > 80.0 { - log::warn!( - "⚠️ Metadata 队列高压:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", - queue_size, - self.max_queue_size, - queue_pressure, - received, - dispatched, - dropped, - drop_rate - ); - } else { - log::info!( - "📊 Metadata 调度器统计:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", - queue_size, - self.max_queue_size, - queue_pressure, - received, - dispatched, - dropped, - drop_rate - ); + if queue_pressure > 80.0 { + log::warn!( + "Metadata 队列高压:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", + queue_len, + self.max_queue_size, + queue_pressure, + received, + dispatched, + dropped, + drop_rate + ); + } else { + log::info!( + "Metadata 调度器统计:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", + queue_len, + self.max_queue_size, + queue_pressure, + received, + dispatched, + dropped, + drop_rate + ); + } } } } diff --git a/src/server.rs b/src/server.rs index da8be22..f3939df 100644 --- a/src/server.rs +++ b/src/server.rs @@ -45,9 +45,9 @@ type FilterCallback = Arc bool + Send + Sync>; pub struct DHTServer { #[allow(dead_code)] options: DHTOptions, - node_id: Vec, + node_id: [u8; 20], socket_providers: Arc>>, - token_secret: Vec, + token_secret: [u8; 10], callback: Arc>>, filter: Arc>>, on_metadata_fetch: Arc>>, @@ -111,8 +111,8 @@ impl DHTServer { }; let node_id = generate_random_id(); - let mut rng = rand::thread_rng(); - let token_secret: Vec = (0..10).map(|_| rng.r#gen::()).collect(); + let mut token_secret = [0u8; 10]; + rand::thread_rng().fill(&mut token_secret); let node_queue = ShardedNodeQueue::new(options.node_queue_capacity); @@ -146,7 +146,7 @@ impl DHTServer { let max_metadata_queue_size = options.max_metadata_queue_size; let server = Self { options, - node_id: node_id.clone(), + node_id, socket_providers: Arc::new(socket_providers), token_secret, callback, @@ -264,13 +264,15 @@ impl DHTServer { } if let Some(nodes) = nodes_batch { - let node_id = server.node_id.clone(); + let node_id = server.node_id; for node in nodes { - let permit = semaphore.clone().acquire_owned().await.unwrap(); - let node_id_clone = node_id.clone(); - // Pick a random avaliable socket - let socket = match server.socket_providers.values().next().cloned() { + let permit = match semaphore.clone().acquire_owned().await { + Ok(p) => p, + Err(_) => break, + }; + let node_id_clone = node_id; + let socket = match server.socket_for_addr(&node.addr) { Some(sock) => sock, None => { log::warn!("未绑定任何地址"); @@ -599,12 +601,15 @@ impl DHTServer { let my_id = if let Some(target) = reference_id { generate_neighbor_target(target, &self.node_id) } else { - self.node_id.clone() + self.node_id.to_vec() }; r_dict.insert(b"id".to_vec(), serde_bencode::value::Value::Bytes(my_id)); let token = self.generate_token(remote_addr); - r_dict.insert(b"token".to_vec(), serde_bencode::value::Value::Bytes(token)); + r_dict.insert( + b"token".to_vec(), + serde_bencode::value::Value::Bytes(token.to_vec()), + ); if query_type == "get_peers" || query_type == "find_node" { let requestor_is_ipv6 = remote_addr.is_ipv6(); @@ -698,13 +703,20 @@ impl DHTServer { } } + fn socket_for_addr(&self, addr: &SocketAddr) -> Option> { + self.socket_providers + .iter() + .find(|(bind_addr, _)| bind_addr.is_ipv4() == addr.is_ipv4()) + .map(|(_, sock)| sock.clone()) + } + async fn send_find_node(&self, target_addr: &SocketAddr, target: &[u8], sender_id: &[u8]) { - if let Some(sock) = self.socket_providers.values().next().cloned() { + if let Some(sock) = self.socket_for_addr(target_addr) { send_find_node_impl(target_addr, target, sender_id, sock).await } } - fn generate_token(&self, addr: SocketAddr) -> Vec { + fn generate_token(&self, addr: SocketAddr) -> [u8; 8] { let mut hasher = ahash::AHasher::default(); match addr.ip() { @@ -714,8 +726,7 @@ impl DHTServer { self.token_secret.hash(&mut hasher); - let hash = hasher.finish(); - hash.to_le_bytes().to_vec() + hasher.finish().to_le_bytes() } fn validate_token(&self, token: &[u8], addr: SocketAddr) -> bool { @@ -723,7 +734,7 @@ impl DHTServer { return false; } let expected = self.generate_token(addr); - token == expected.as_slice() + token == expected } } @@ -811,9 +822,10 @@ async fn send_find_node_impl( } } -fn generate_random_id() -> Vec { - let mut rng = rand::thread_rng(); - (0..20).map(|_| rng.r#gen::()).collect() +fn generate_random_id() -> [u8; 20] { + let mut id = [0u8; 20]; + rand::thread_rng().fill(&mut id); + id } /// 生成邻居目标节点 ID @@ -945,7 +957,8 @@ fn process_udp_packet( return Err(ProcessUdpPacketError::InvalidPacket); } let mut data = Some(buffer[..size].to_owned().into_boxed_slice()); - let mut choked_count = 0; + let mut attempts = 0; + let max_attempts = workers.len(); while let Some(packet) = data.take() { let worker = &workers[*worker_index]; @@ -956,34 +969,28 @@ fn process_udp_packet( break; } Err(mpsc::error::TrySendError::Full((packet, _, _))) => { - choked_count += 1; - if *worker_index == 0 { - if choked_count >= workers.len() { - // all workers choked - #[cfg(feature = "metrics")] - counter!("dht_udp_packets_received_total", "status" => "queue_full") - .increment(1); + attempts += 1; + if attempts >= max_attempts { + #[cfg(feature = "metrics")] + counter!("dht_udp_packets_received_total", "status" => "queue_full") + .increment(1); - #[cfg(debug_assertions)] - log::trace!("UDP worker queue full, dropping packet"); - return Err(ProcessUdpPacketError::ChokedWorkers); - } - choked_count = 0 + #[cfg(debug_assertions)] + log::trace!("UDP worker queue full, dropping packet"); + return Err(ProcessUdpPacketError::ChokedWorkers); } - let _ = data.insert(packet); // choose the next worker + let _ = data.insert(packet); } Err(mpsc::error::TrySendError::Closed((packet, _, _))) => { log::warn!("UDP worker dropped."); - workers.swap_remove(*worker_index); // remove the dead worker. faster but does not retain ordering. - let _ = data.insert(packet); // choose the next worker + workers.swap_remove(*worker_index); + let _ = data.insert(packet); } } if workers.is_empty() { - // no live workers return Err(ProcessUdpPacketError::NoLiveWorkers); } - // dispatch messages to workers in round-robin style - *worker_index = (*worker_index + 1) % workers.len(); // note: a dead worker may be removed + *worker_index = (*worker_index + 1) % workers.len(); } Ok(()) diff --git a/src/sharded.rs b/src/sharded.rs index 8359b0c..0caf00c 100644 --- a/src/sharded.rs +++ b/src/sharded.rs @@ -1,54 +1,9 @@ -use bloomfilter::Bloom; use std::collections::{HashSet, VecDeque}; use std::net::SocketAddr; use std::sync::Mutex; -use std::sync::atomic::{AtomicUsize, Ordering}; -const BLOOM_SHARD_COUNT: usize = 32; const QUEUE_SHARD_COUNT: usize = 16; -pub struct ShardedBloom { - shards: Vec>>, - count: AtomicUsize, -} - -impl ShardedBloom { - pub fn new_for_fp_rate(expected_items: usize, fp_rate: f64) -> Self { - #[allow(clippy::manual_div_ceil)] - let items_per_shard = (expected_items + BLOOM_SHARD_COUNT - 1) / BLOOM_SHARD_COUNT; - - let shards = (0..BLOOM_SHARD_COUNT) - .map(|_| Mutex::new(Bloom::new_for_fp_rate(items_per_shard, fp_rate))) - .collect(); - - Self { - shards, - count: AtomicUsize::new(0), - } - } - - pub fn check_and_set(&self, hash: &[u8; 20]) -> bool { - let shard_idx = self.hash_to_shard(hash); - let mut shard = self.shards[shard_idx].lock().unwrap(); - let present = shard.check_and_set(hash); - - if !present { - self.count.fetch_add(1, Ordering::Relaxed); - } - present - } - - pub fn number_of_bits(&self) -> u64 { - self.count.load(Ordering::Relaxed) as u64 - } - - #[inline] - fn hash_to_shard(&self, hash: &[u8; 20]) -> usize { - let idx = (hash[0] as usize) | ((hash[1] as usize) << 8); - idx % BLOOM_SHARD_COUNT - } -} - #[derive(Debug, Clone)] pub struct NodeTuple { pub id: Vec, diff --git a/src/types.rs b/src/types.rs index caeb1ea..3e02e4c 100644 --- a/src/types.rs +++ b/src/types.rs @@ -55,8 +55,6 @@ fn format_bytes(bytes: u64) -> String { pub struct DHTOptions { pub port: u16, - pub auto_metadata: bool, - pub metadata_timeout: u64, pub max_metadata_queue_size: usize, @@ -74,7 +72,6 @@ impl Default for DHTOptions { fn default() -> Self { Self { port: 6881, - auto_metadata: true, metadata_timeout: 3, max_metadata_queue_size: 100000, max_metadata_worker_count: 1000,