diff --git a/.gitignore b/.gitignore index 03b6f9f..1d9756f 100644 --- a/.gitignore +++ b/.gitignore @@ -15,3 +15,6 @@ torrents/ # 日志 *.log + +# 个人脚本(不提交到仓库) +scripts/ diff --git a/examples/main.rs b/examples/main.rs index 8527ac7..dc7848a 100644 --- a/examples/main.rs +++ b/examples/main.rs @@ -3,18 +3,16 @@ static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc; use dht_crawler::prelude::*; -use std::sync::Arc; -use tracing_subscriber::EnvFilter; -use std::sync::atomic::{AtomicUsize, Ordering}; #[cfg(feature = "metrics")] use metrics_exporter_prometheus::PrometheusBuilder; #[cfg(feature = "metrics")] use std::net::SocketAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; +use tracing_subscriber::EnvFilter; #[tokio::main] async fn main() -> Result<()> { - - let filter = EnvFilter::try_from_default_env() - .unwrap_or_else(|_| EnvFilter::new("info")); + let filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info")); tracing_subscriber::fmt() .with_env_filter(filter) @@ -35,11 +33,11 @@ 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, // ✅ 激进并发(最大化吞吐) - netmode: NetMode::Ipv4Only, // 网络模式:Ipv4Only(仅IPv4)、Ipv6Only(仅IPv6)、DualStack(双栈,默认) - ..Default::default() // 使用默认值填充其他字段(节点队列容量等) + metadata_timeout: 3, // ✅ 快速超时,快速失败 + max_metadata_queue_size: 100000, // ✅ 大缓冲区(防止饱和) + max_metadata_worker_count: 1000, // ✅ 激进并发(最大化吞吐) + netmode: NetMode::Ipv4Only, // 网络模式:Ipv4Only(仅IPv4)、Ipv6Only(仅IPv6)、DualStack(双栈,默认) + ..Default::default() // 使用默认值填充其他字段(节点队列容量等) }; // 统计计数器 @@ -55,7 +53,7 @@ async fn main() -> Result<()> { // 设置 torrent 回调 server.on_torrent(move |_torrent| { let _count = torrent_count_clone.fetch_add(1, Ordering::Relaxed) + 1; - + // 🔇 取消打印 torrent 信息,减少日志输出 // let total_size: u64 = torrent.files.iter().map(|f| f.size).sum(); // let files_display = if torrent.files.len() <= 3 { @@ -77,9 +75,7 @@ async fn main() -> Result<()> { }); // 设置元数据获取前的检查回调 - server.on_metadata_fetch(|_hash| async move { - true - }); + server.on_metadata_fetch(|_hash| async move { true }); // 启动监控任务 let count_monitor = torrent_count.clone(); @@ -95,7 +91,8 @@ async fn main() -> Result<()> { // ✅ 监控:爬虫运行状态 log::info!( "📊 [监控] 时长: {}s | 成功抓取: ✨ {}", - uptime, success_fetch + uptime, + success_fetch ); if uptime > 0 && success_fetch > 0 { @@ -108,4 +105,3 @@ async fn main() -> Result<()> { server.start().await?; Ok(()) } - diff --git a/src/error.rs b/src/error.rs index 7bcbfb7..0a8fc49 100644 --- a/src/error.rs +++ b/src/error.rs @@ -4,22 +4,22 @@ use thiserror::Error; 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 332244d..11baee9 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,20 +1,20 @@ mod error; -mod server; -pub mod protocol; -pub mod types; pub mod metadata; -mod sharded; +pub mod protocol; pub mod scheduler; +mod server; +mod sharded; +pub mod types; pub use error::{DHTError, Result}; -pub use server::{DHTServer, HashDiscovered}; -pub use types::{DHTOptions, FileInfo, TorrentInfo, NetMode}; -pub use sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple}; pub use scheduler::MetadataScheduler; +pub use server::{DHTServer, HashDiscovered}; +pub use sharded::{NodeTuple, ShardedBloom, ShardedNodeQueue}; +pub use types::{DHTOptions, FileInfo, NetMode, TorrentInfo}; pub mod prelude { pub use crate::error::{DHTError, Result}; - pub use crate::server::DHTServer; - pub use crate::types::{DHTOptions, FileInfo, TorrentInfo, NetMode}; pub use crate::scheduler::MetadataScheduler; + pub use crate::server::DHTServer; + pub use crate::types::{DHTOptions, FileInfo, NetMode, TorrentInfo}; } diff --git a/src/metadata.rs b/src/metadata.rs index 4367e52..34fa541 100644 --- a/src/metadata.rs +++ b/src/metadata.rs @@ -1,17 +1,17 @@ -use std::collections::BTreeMap; -use std::net::SocketAddr; -use std::time::Duration; +use crate::types::FileInfo; use bytes::Bytes; #[cfg(feature = "metrics")] use metrics::{counter, histogram}; -use sha1::{Digest, Sha1}; -use tokio::time::timeout; -use rbit::{ - metadata_piece_count, ExtensionHandshake, Message, MetadataMessage, - MetadataMessageType, PeerConnection, PeerId, -}; use rbit::peer::ExtensionMessage; -use crate::types::FileInfo; +use rbit::{ + ExtensionHandshake, Message, MetadataMessage, MetadataMessageType, PeerConnection, PeerId, + metadata_piece_count, +}; +use sha1::{Digest, Sha1}; +use std::collections::BTreeMap; +use std::net::SocketAddr; +use std::time::Duration; +use tokio::time::timeout; #[derive(Clone)] pub struct RbitFetcher { @@ -38,27 +38,32 @@ impl RbitFetcher { let mut conn = match timeout( Duration::from_secs(3), PeerConnection::connect(peer_addr, *info_hash, *peer_id.as_bytes()), - ).await { + ) + .await + { Ok(Ok(c)) => { #[cfg(feature = "metrics")] - counter!("dht_metadata_connection_result_total", "result" => "success").increment(1); + counter!("dht_metadata_connection_result_total", "result" => "success") + .increment(1); c - }, + } Ok(Err(_)) => { #[cfg(feature = "metrics")] counter!("dht_metadata_connection_result_total", "result" => "failed").increment(1); return None; - }, + } Err(_) => { #[cfg(feature = "metrics")] - counter!("dht_metadata_connection_result_total", "result" => "timeout").increment(1); + counter!("dht_metadata_connection_result_total", "result" => "timeout") + .increment(1); return None; - }, + } }; if !conn.supports_extension { #[cfg(feature = "metrics")] - counter!("dht_metadata_handshake_result_total", "result" => "no_extension_support").increment(1); + counter!("dht_metadata_handshake_result_total", "result" => "no_extension_support") + .increment(1); return None; } @@ -66,7 +71,12 @@ impl RbitFetcher { let handshake = ExtensionHandshake::with_extensions(&[("ut_metadata", my_ut_metadata_id)]); if let Ok(handshake_bytes) = handshake.encode() { - let _ = conn.send(Message::Extended { id: 0, payload: handshake_bytes }).await; + let _ = conn + .send(Message::Extended { + id: 0, + payload: handshake_bytes, + }) + .await; } else { return None; } @@ -79,9 +89,8 @@ impl RbitFetcher { let result = timeout(self.timeout, async { loop { let msg = conn.receive().await.ok()?; - match msg { - Message::Extended { id, payload } => { - if id == 0 { + if let Message::Extended { id, payload } = msg { + if id == 0 { if let Ok(ExtensionMessage::Handshake(remote_hs)) = ExtensionMessage::decode(id, &payload) { if let Some(size) = remote_hs.metadata_size { metadata_size = size as u32; @@ -91,10 +100,10 @@ impl RbitFetcher { } } if metadata_size > 0 && remote_ut_metadata_id > 0 && !request_sent { - if metadata_size > 10 * 1024 * 1024 { + if metadata_size > 10 * 1024 * 1024 { #[cfg(feature = "metrics")] counter!("dht_metadata_fetch_fail_total", "reason" => "size_limit").increment(1); - return None; + return None; } let count = metadata_piece_count(metadata_size as usize); @@ -107,14 +116,12 @@ impl RbitFetcher { request_sent = true; } } else if id == my_ut_metadata_id { - if let Ok(meta_msg) = MetadataMessage::decode(&payload) { - if meta_msg.msg_type == MetadataMessageType::Data { - if let Some(data) = meta_msg.data { - #[cfg(feature = "metrics")] - counter!("dht_metadata_bytes_downloaded_total").increment(data.len() as u64); - pieces.insert(meta_msg.piece, data); - } - } + if let Ok(meta_msg) = MetadataMessage::decode(&payload) + && meta_msg.msg_type == MetadataMessageType::Data + && let Some(data) = meta_msg.data { + #[cfg(feature = "metrics")] + counter!("dht_metadata_bytes_downloaded_total").increment(data.len() as u64); + pieces.insert(meta_msg.piece, data); } if metadata_size > 0 { let total_received: usize = pieces.values().map(|p| p.len()).sum(); @@ -152,48 +159,60 @@ impl RbitFetcher { } } } - _ => {} - } } }).await; match result { Ok(Some(info_bytes)) => { - if let Ok(value) = rbit::decode(&info_bytes) { - if let Some(dict) = value.as_dict() { - let name = dict.get(&b"name"[..]).and_then(|v| v.as_str()).unwrap_or("Unknown").to_string(); + if let Ok(value) = rbit::decode(&info_bytes) + && let Some(dict) = value.as_dict() { + let name = dict + .get(&b"name"[..]) + .and_then(|v| v.as_str()) + .unwrap_or("Unknown") + .to_string(); 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()) { for file in files { - if let Some(f_dict) = file.as_dict() { - if let Some(len) = f_dict.get(&b"length"[..]).and_then(|v| v.as_integer()) { - let len = len as u64; - total_size += len; - let mut path_parts = Vec::new(); - if let Some(path_list) = f_dict.get(&b"path"[..]).and_then(|v| v.as_list()) { - for p in path_list { - if let Some(p_str) = p.as_str() { path_parts.push(p_str); } + if let Some(f_dict) = file.as_dict() + && let Some(len) = f_dict.get(&b"length"[..]).and_then(|v| v.as_integer()) { + let len = len as u64; + total_size += len; + let mut path_parts = Vec::new(); + if let Some(path_list) = + f_dict.get(&b"path"[..]).and_then(|v| v.as_list()) + { + for p in path_list { + if let Some(p_str) = p.as_str() { + path_parts.push(p_str); } } - file_list.push(FileInfo { path: path_parts.join("/"), size: len }); } + file_list.push(FileInfo { + path: path_parts.join("/"), + size: len, + }); } } - } else if let Some(len) = dict.get(&b"length"[..]).and_then(|v| v.as_integer()) { + } else if let Some(len) = + dict.get(&b"length"[..]).and_then(|v| v.as_integer()) + { total_size = len as u64; - file_list.push(FileInfo { path: name.clone(), size: total_size }); + file_list.push(FileInfo { + path: name.clone(), + size: total_size, + }); } - if total_size > 0 { + if total_size > 0 { #[cfg(feature = "metrics")] { 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)); } } - } #[cfg(feature = "metrics")] counter!("dht_metadata_fetch_fail_total", "reason" => "parse_error").increment(1); None @@ -202,7 +221,7 @@ impl RbitFetcher { #[cfg(feature = "metrics")] counter!("dht_metadata_fetch_fail_total", "reason" => "timeout").increment(1); None - }, + } } } -} \ No newline at end of file +} diff --git a/src/protocol.rs b/src/protocol.rs index 6ce934e..c9b712c 100644 --- a/src/protocol.rs +++ b/src/protocol.rs @@ -32,4 +32,3 @@ pub struct DhtResponse { #[serde(default)] pub nodes6: Option, } - diff --git a/src/scheduler.rs b/src/scheduler.rs index cee0886..88e65a8 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -1,15 +1,19 @@ +use crate::metadata::RbitFetcher; use crate::server::HashDiscovered; use crate::types::TorrentInfo; -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; +use std::sync::{Arc, RwLock}; #[cfg(debug_assertions)] use std::time::Duration; +use tokio::sync::{Mutex, mpsc}; +use tokio_util::sync::CancellationToken; type TorrentCallback = Arc; -type MetadataFetchCallback = Arc std::pin::Pin + Send>> + Send + Sync>; +type MetadataFetchCallback = Arc< + dyn Fn(String) -> std::pin::Pin + Send>> + + Send + + Sync, +>; pub struct MetadataScheduler { hash_rx: mpsc::Receiver, @@ -26,6 +30,7 @@ pub struct MetadataScheduler { } impl MetadataScheduler { + #[allow(clippy::too_many_arguments)] pub fn new( hash_rx: mpsc::Receiver, fetcher: Arc, @@ -50,23 +55,23 @@ impl MetadataScheduler { shutdown, } } - + pub fn set_callback(&mut self, callback: TorrentCallback) { if let Ok(mut guard) = self.callback.try_write() { *guard = Some(callback); } } - + pub fn set_metadata_fetch_callback(&mut self, callback: MetadataFetchCallback) { if let Ok(mut guard) = self.on_metadata_fetch.try_write() { *guard = Some(callback); } } - - pub async fn run(mut self) { + + 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 shutdown = self.shutdown.clone(); #[cfg_attr(not(debug_assertions), allow(unused_variables))] for worker_id in 0..self.max_concurrent { @@ -77,11 +82,11 @@ impl MetadataScheduler { 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 { tokio::select! { _ = shutdown_worker.cancelled() => { @@ -99,14 +104,14 @@ impl MetadataScheduler { } result }; - + let hash = match hash { Some(h) => h, None => break, }; - + total_dispatched.fetch_add(1, Ordering::Relaxed); - + Self::process_hash( hash, &fetcher, @@ -116,17 +121,17 @@ impl MetadataScheduler { } } } - + #[cfg(debug_assertions)] log::trace!("Worker {} 退出", worker_id); }); } - + #[cfg(debug_assertions)] let mut stats_interval = tokio::time::interval(Duration::from_secs(60)); #[cfg(debug_assertions)] stats_interval.tick().await; - + let shutdown = self.shutdown.clone(); loop { #[cfg(debug_assertions)] @@ -139,7 +144,7 @@ impl MetadataScheduler { } 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); @@ -150,15 +155,15 @@ impl MetadataScheduler { Err(_) => break, } } - + _ = stats_interval.tick() => { self.print_stats(&task_tx); } - + else => break, } } - + #[cfg(not(debug_assertions))] { tokio::select! { @@ -171,7 +176,7 @@ impl MetadataScheduler { 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); @@ -188,13 +193,13 @@ impl MetadataScheduler { } } } - + // 显式关闭 task_tx,让所有 worker 任务能够退出 drop(task_tx); #[cfg(debug_assertions)] log::trace!("MetadataScheduler 主循环退出,等待 worker 任务完成"); } - + async fn process_hash( hash: HashDiscovered, fetcher: &Arc, @@ -203,7 +208,7 @@ impl MetadataScheduler { ) { let info_hash = hash.info_hash.clone(); let peer_addr = hash.peer_addr; - + let maybe_check_fn = { match on_metadata_fetch.read() { Ok(guard) => guard.clone(), @@ -211,12 +216,11 @@ impl MetadataScheduler { } }; - if let Some(f) = maybe_check_fn { - if !f(info_hash.clone()).await { - return; - } + if let Some(f) = maybe_check_fn + && !f(info_hash.clone()).await { + return; } - + let info_hash_bytes: [u8; 20] = match hex::decode(&info_hash) { Ok(bytes) if bytes.len() == 20 => { let mut arr = [0u8; 20]; @@ -225,7 +229,7 @@ impl MetadataScheduler { } _ => return, }; - + if let Some((name, total_size, files)) = fetcher.fetch(&info_hash_bytes, peer_addr).await { let metadata = TorrentInfo { info_hash, @@ -240,35 +244,35 @@ impl MetadataScheduler { .unwrap() .as_secs(), }; - + let maybe_torrent_cb = { match callback.read() { Ok(guard) => guard.clone(), Err(_) => return, } }; - + if let Some(cb) = maybe_torrent_cb { cb(metadata); } } } - + #[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); - + 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; - + if queue_pressure > 80.0 { log::warn!( "⚠️ Metadata 队列高压:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", diff --git a/src/server.rs b/src/server.rs index 27dc401..3f85ca2 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,24 +1,24 @@ use crate::error::Result; use crate::metadata::RbitFetcher; -use crate::protocol::{DhtMessage, DhtArgs, DhtResponse}; +use crate::protocol::{DhtArgs, DhtMessage, DhtResponse}; use crate::scheduler::MetadataScheduler; -use crate::types::{DHTOptions, TorrentInfo, NetMode}; -use crate::sharded::{ShardedNodeQueue, NodeTuple}; -use rand::Rng; +use crate::sharded::{NodeTuple, ShardedNodeQueue}; +use crate::types::{DHTOptions, NetMode, TorrentInfo}; use ahash::AHasher; -use std::hash::{Hash, Hasher}; -use std::net::{IpAddr, Ipv6Addr, SocketAddr}; -use std::sync::{Arc, RwLock}; -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}; #[cfg(feature = "metrics")] use metrics::{counter, gauge}; -use std::pin::Pin; +use rand::Rng; +use socket2::{Domain, Protocol, Socket, Type}; use std::future::Future; +use std::hash::{Hash, Hasher}; +use std::net::{IpAddr, Ipv6Addr, SocketAddr}; +use std::pin::Pin; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, RwLock}; +use std::time::Duration; +use tokio::net::UdpSocket; +use tokio::sync::{Semaphore, mpsc}; +use tokio_util::sync::CancellationToken; const BOOTSTRAP_NODES: &[&str] = &[ "router.bittorrent.com:6881", @@ -64,37 +64,45 @@ impl DHTServer { NetMode::Ipv4Only => { let sock = Socket::new(Domain::IPV4, Type::DGRAM, Some(Protocol::UDP))?; #[cfg(not(windows))] - { let _ = sock.set_reuse_port(true); } + { + let _ = sock.set_reuse_port(true); + } let _ = sock.set_reuse_address(true); sock.set_nonblocking(true)?; - + let _ = sock.set_recv_buffer_size(32 * 1024 * 1024); let _ = sock.set_send_buffer_size(8 * 1024 * 1024); let addr: SocketAddr = format!("0.0.0.0:{}", options.port).parse().unwrap(); sock.bind(&addr.into())?; (Arc::new(UdpSocket::from_std(sock.into())?), None) - }, + } NetMode::Ipv6Only => { let sock = Socket::new(Domain::IPV6, Type::DGRAM, Some(Protocol::UDP))?; #[cfg(not(windows))] - { let _ = sock.set_reuse_port(true); } + { + let _ = sock.set_reuse_port(true); + } let _ = sock.set_reuse_address(true); #[cfg(not(windows))] - { let _ = sock.set_only_v6(true); } + { + let _ = sock.set_only_v6(true); + } sock.set_nonblocking(true)?; - + let _ = sock.set_recv_buffer_size(32 * 1024 * 1024); let _ = sock.set_send_buffer_size(8 * 1024 * 1024); let addr: SocketAddr = format!("[::]:{}", options.port).parse().unwrap(); sock.bind(&addr.into())?; (Arc::new(UdpSocket::from_std(sock.into())?), None) - }, + } NetMode::DualStack => { let sock_v4 = Socket::new(Domain::IPV4, Type::DGRAM, Some(Protocol::UDP))?; #[cfg(not(windows))] - { let _ = sock_v4.set_reuse_port(true); } + { + let _ = sock_v4.set_reuse_port(true); + } let _ = sock_v4.set_reuse_address(true); sock_v4.set_nonblocking(true)?; let _ = sock_v4.set_recv_buffer_size(32 * 1024 * 1024); @@ -105,10 +113,14 @@ impl DHTServer { let sock_v6 = Socket::new(Domain::IPV6, Type::DGRAM, Some(Protocol::UDP))?; #[cfg(not(windows))] - { let _ = sock_v6.set_reuse_port(true); } + { + let _ = sock_v6.set_reuse_port(true); + } let _ = sock_v6.set_reuse_address(true); #[cfg(not(windows))] - { let _ = sock_v6.set_only_v6(true); } + { + let _ = sock_v6.set_only_v6(true); + } sock_v6.set_nonblocking(true)?; let _ = sock_v6.set_recv_buffer_size(32 * 1024 * 1024); let _ = sock_v6.set_send_buffer_size(8 * 1024 * 1024); @@ -117,7 +129,7 @@ impl DHTServer { let socket_v6 = Some(Arc::new(UdpSocket::from_std(sock_v6.into())?)); (socket, socket_v6) - }, + } }; let node_id = generate_random_id(); @@ -129,10 +141,10 @@ impl DHTServer { let (hash_tx, hash_rx) = mpsc::channel::(10000); let fetcher = Arc::new(RbitFetcher::new(options.metadata_timeout)); - + let callback = Arc::new(RwLock::new(None)); let on_metadata_fetch = Arc::new(RwLock::new(None)); - + let metadata_queue_len = Arc::new(AtomicUsize::new(0)); let shutdown = CancellationToken::new(); @@ -187,19 +199,15 @@ impl DHTServer { fn select_socket(&self, addr: &SocketAddr) -> &Arc { match self.options.netmode { - NetMode::Ipv4Only => { - &self.socket - }, - NetMode::Ipv6Only => { - &self.socket - }, + NetMode::Ipv4Only => &self.socket, + NetMode::Ipv6Only => &self.socket, NetMode::DualStack => { if addr.is_ipv6() { self.socket_v6.as_ref().unwrap_or(&self.socket) } else { &self.socket } - }, + } } } @@ -208,20 +216,24 @@ impl DHTServer { F: Fn(String) -> Fut + Send + Sync + 'static, Fut: Future + Send + 'static, { - *self.on_metadata_fetch.write().unwrap() = Some(Arc::new(move |hash| { - Box::pin(callback(hash)) - })); + *self.on_metadata_fetch.write().unwrap() = + Some(Arc::new(move |hash| Box::pin(callback(hash)))); } - pub fn on_torrent(&self, callback: F) where F: Fn(TorrentInfo) + Send + Sync + 'static { + pub fn on_torrent(&self, callback: F) + where + F: Fn(TorrentInfo) + Send + Sync + 'static, + { *self.callback.write().unwrap() = Some(Arc::new(callback)); } - - pub fn set_filter(&self, filter: F) where F: Fn(&str) -> bool + Send + Sync + 'static { + + pub fn set_filter(&self, filter: F) + where + F: Fn(&str) -> bool + Send + Sync + 'static, + { *self.filter.write().unwrap() = Some(Arc::new(filter)); } - pub fn get_node_pool_size(&self) -> usize { self.node_queue.len() } @@ -253,14 +265,14 @@ impl DHTServer { let queue_len = server.metadata_queue_len.load(Ordering::Relaxed); let queue_pressure = queue_len as f64 / server.max_metadata_queue_size as f64; - + #[cfg(feature = "metrics")] { gauge!("dht_metadata_queue_size").set(queue_len as f64); gauge!("dht_metadata_worker_pressure").set(queue_pressure); gauge!("dht_node_queue_size").set(server.node_queue.len() as f64); } - + let (batch_size, sleep_duration) = if queue_pressure < 0.8 { (200, Duration::from_millis(10)) } else if queue_pressure < 0.95 { @@ -274,9 +286,9 @@ impl DHTServer { NetMode::Ipv6Only => Some(true), NetMode::DualStack => None, }; - + let queue_empty = server.node_queue.is_empty_for(filter_ipv6); - + let nodes_batch = { if queue_empty || batch_size == 0 { None @@ -302,7 +314,7 @@ impl DHTServer { let socket = server.socket.clone(); let socket_v6 = server.socket_v6.clone(); let netmode = server.options.netmode; - + for node in nodes { let permit = semaphore.clone().acquire_owned().await.unwrap(); let node_id_clone = node_id.clone(); @@ -310,9 +322,10 @@ impl DHTServer { let socket_v6_clone = socket_v6.clone(); let node_addr = node.addr; let node_id_for_target = node.id; - + tokio::spawn(async move { - let neighbor_id = generate_neighbor_target(&node_id_for_target, &node_id_clone); + let neighbor_id = + generate_neighbor_target(&node_id_for_target, &node_id_clone); let random_target = generate_random_id(); let _ = send_find_node_impl( node_addr, @@ -321,7 +334,8 @@ impl DHTServer { &socket_clone, socket_v6_clone.as_ref(), netmode, - ).await; + ) + .await; drop(permit); }); } @@ -396,7 +410,7 @@ impl DHTServer { shutdown: CancellationToken, ) { let num_workers = senders.len(); - + tokio::spawn(async move { let mut buf = [0u8; 65536]; let mut next_worker_idx = 0; @@ -422,7 +436,7 @@ impl DHTServer { log::trace!("⚠️ 拒绝异常大的 UDP 包: {} 字节 from {}", size, addr); continue; } - + if size == 0 || buf[0] != b'd' { #[cfg(feature = "metrics")] counter!("dht_udp_packets_received_total", "status" => "dropped_magic").increment(1); @@ -465,7 +479,11 @@ impl DHTServer { async fn handle_message(&self, data: &[u8], addr: SocketAddr) -> Result<()> { if !self.is_addr_allowed(&addr) { #[cfg(debug_assertions)] - log::trace!("⚠️ 拒绝不匹配的地址类型: {} (当前模式: {:?})", addr, self.options.netmode); + log::trace!( + "⚠️ 拒绝不匹配的地址类型: {} (当前模式: {:?})", + addr, + self.options.netmode + ); return Ok(()); } @@ -475,7 +493,7 @@ impl DHTServer { #[cfg(feature = "metrics")] counter!("dht_messages_parse_error_total").increment(1); return Ok(()); - }, + } }; #[cfg(feature = "metrics")] @@ -485,7 +503,7 @@ impl DHTServer { "q" => "q", "r" => "r", "e" => "e", - _ => "unknown", // 将所有非法/未知类型归一化 + _ => "unknown", // 将所有非法/未知类型归一化 }; counter!("dht_messages_processed_total", "type" => label).increment(1); } @@ -506,7 +524,12 @@ impl DHTServer { Ok(()) } - async fn handle_query(&self, msg: &DhtMessage, query_type: &[u8], addr: SocketAddr) -> Result<()> { + async fn handle_query( + &self, + msg: &DhtMessage, + query_type: &[u8], + addr: SocketAddr, + ) -> Result<()> { let args = match &msg.a { Some(a) => a, None => return Ok(()), @@ -514,12 +537,14 @@ impl DHTServer { let transaction_id = &msg.t; let sender_id: Option<&[u8]> = args.id.as_deref().map(|v| v.as_slice()); - let target_id_fallback: Option<&[u8]> = args.target.as_deref() + let target_id_fallback: Option<&[u8]> = args + .target + .as_deref() .or(args.info_hash.as_deref()) .map(|v| v.as_slice()); let q_str = std::str::from_utf8(query_type).unwrap_or(""); - + #[cfg(feature = "metrics")] { let label = match q_str { @@ -527,7 +552,7 @@ impl DHTServer { "find_node" => "find_node", "get_peers" => "get_peers", "announce_peer" => "announce_peer", - "vote" => "vote", + "vote" => "vote", _ => "other_or_invalid", }; counter!("dht_queries_total", "q" => label).increment(1); @@ -537,16 +562,18 @@ impl DHTServer { self.handle_announce_peer(args, addr).await?; } - self.send_response(transaction_id, addr, q_str, sender_id, target_id_fallback).await?; + self.send_response(transaction_id, addr, q_str, sender_id, target_id_fallback) + .await?; Ok(()) } async fn handle_announce_peer(&self, args: &DhtArgs, addr: SocketAddr) -> Result<()> { if let Some(token) = &args.token { - if !self.validate_token(token, addr) { + if !self.validate_token(token, addr) { #[cfg(feature = "metrics")] - counter!("dht_announce_peer_blocked_total", "reason" => "invalid_token").increment(1); - return Ok(()); + counter!("dht_announce_peer_blocked_total", "reason" => "invalid_token") + .increment(1); + return Ok(()); } } else { return Ok(()); @@ -554,17 +581,18 @@ impl DHTServer { if let Some(info_hash) = &args.info_hash { let info_hash_arr: [u8; 20] = match info_hash.as_ref().try_into() { - Ok(arr) => arr, Err(_) => return Ok(()), + Ok(arr) => arr, + Err(_) => return Ok(()), }; let hash_hex = hex::encode(info_hash_arr); let filter_cb = self.filter.read().unwrap().clone(); - if let Some(f) = filter_cb { - if !f(&hash_hex) { - #[cfg(feature = "metrics")] - counter!("dht_announce_peer_blocked_total", "reason" => "filtered").increment(1); - return Ok(()); - } + if let Some(f) = filter_cb + && !f(&hash_hex) { + #[cfg(feature = "metrics")] + counter!("dht_announce_peer_blocked_total", "reason" => "filtered") + .increment(1); + return Ok(()); } #[cfg(feature = "metrics")] @@ -574,7 +602,11 @@ impl DHTServer { log::debug!("🔥 新 Hash: {} 来自 {}", hash_hex, addr); let port = if let Some(implied) = args.implied_port { - if implied != 0 { addr.port() } else { args.port.unwrap_or(0) } + if implied != 0 { + addr.port() + } else { + args.port.unwrap_or(0) + } } else { args.port.unwrap_or(addr.port()) }; @@ -586,7 +618,7 @@ impl DHTServer { discovered_at: std::time::Instant::now(), }; - if let Err(_) = self.hash_tx.try_send(event) { + if self.hash_tx.try_send(event).is_err() { #[cfg(debug_assertions)] log::debug!("⚠️ Hash 队列满,丢弃 hash"); } @@ -610,15 +642,18 @@ impl DHTServer { return; } - if nodes_bytes.len() % 26 != 0 { return; } + #[allow(clippy::manual_is_multiple_of)] + if nodes_bytes.len() % 26 != 0 { + return; + } for chunk in nodes_bytes.chunks(26) { let id = chunk[0..20].to_vec(); let port = u16::from_be_bytes([chunk[24], chunk[25]]); - + let ip = std::net::Ipv4Addr::new(chunk[20], chunk[21], chunk[22], chunk[23]); let addr = SocketAddr::new(std::net::IpAddr::V4(ip), port); - + #[cfg(feature = "metrics")] counter!("dht_nodes_discovered_total", "ip_version" => "v4").increment(1); @@ -631,7 +666,10 @@ impl DHTServer { return; } - if nodes_bytes.len() % 38 != 0 { return; } + #[allow(clippy::manual_is_multiple_of)] + if nodes_bytes.len() % 38 != 0 { + return; + } for chunk in nodes_bytes.chunks(38) { let id = chunk[0..20].to_vec(); let port = u16::from_be_bytes([chunk[36], chunk[37]]); @@ -679,9 +717,9 @@ impl DHTServer { NetMode::Ipv6Only => Some(true), NetMode::DualStack => Some(requestor_is_ipv6), }; - + let nodes = self.node_queue.get_random_nodes(8, filter_ipv6); - + let mut nodes_data = Vec::new(); let mut nodes6_data = Vec::new(); @@ -691,29 +729,40 @@ impl DHTServer { nodes_data.extend_from_slice(&node.id); nodes_data.extend_from_slice(&ip.octets()); nodes_data.extend_from_slice(&node.addr.port().to_be_bytes()); - }, + } IpAddr::V6(ip) => { nodes6_data.extend_from_slice(&node.id); nodes6_data.extend_from_slice(&ip.octets()); nodes6_data.extend_from_slice(&node.addr.port().to_be_bytes()); - }, + } } } - + if requestor_is_ipv6 { if !nodes6_data.is_empty() { - r_dict.insert(b"nodes6".to_vec(), serde_bencode::value::Value::Bytes(nodes6_data)); - } - } else { - if !nodes_data.is_empty() { - r_dict.insert(b"nodes".to_vec(), serde_bencode::value::Value::Bytes(nodes_data)); + r_dict.insert( + b"nodes6".to_vec(), + serde_bencode::value::Value::Bytes(nodes6_data), + ); } + } else if !nodes_data.is_empty() { + r_dict.insert( + b"nodes".to_vec(), + serde_bencode::value::Value::Bytes(nodes_data), + ); } } - let mut response: std::collections::HashMap = std::collections::HashMap::new(); - response.insert("t".to_string(), serde_bencode::value::Value::Bytes(tid.to_vec())); - response.insert("y".to_string(), serde_bencode::value::Value::Bytes(b"r".to_vec())); + let mut response: std::collections::HashMap = + std::collections::HashMap::new(); + response.insert( + "t".to_string(), + serde_bencode::value::Value::Bytes(tid.to_vec()), + ); + response.insert( + "y".to_string(), + serde_bencode::value::Value::Bytes(b"r".to_vec()), + ); response.insert("r".to_string(), serde_bencode::value::Value::Dict(r_dict)); if let Ok(encoded) = serde_bencode::to_bytes(&response) { @@ -730,28 +779,33 @@ impl DHTServer { async fn bootstrap(&self) { let target = generate_random_id(); for node in BOOTSTRAP_NODES { - match tokio::net::lookup_host(node).await { - Ok(addrs) => { - for addr in addrs { - match self.options.netmode { - NetMode::Ipv4Only => { - if addr.is_ipv6() { continue; } - }, - NetMode::Ipv6Only => { - if addr.is_ipv4() { continue; } - }, - NetMode::DualStack => { - }, + if let Ok(addrs) = tokio::net::lookup_host(node).await { + for addr in addrs { + match self.options.netmode { + NetMode::Ipv4Only => { + if addr.is_ipv6() { + continue; + } } - let _ = self.send_find_node(addr, &target, &self.node_id).await; + NetMode::Ipv6Only => { + if addr.is_ipv4() { + continue; + } + } + NetMode::DualStack => {} } + let _ = self.send_find_node(addr, &target, &self.node_id).await; } - Err(_) => {} } } } - async fn send_find_node(&self, addr: SocketAddr, target: &[u8], sender_id: &[u8]) -> Result<()> { + async fn send_find_node( + &self, + addr: SocketAddr, + target: &[u8], + sender_id: &[u8], + ) -> Result<()> { send_find_node_impl( addr, target, @@ -759,24 +813,24 @@ impl DHTServer { &self.socket, self.socket_v6.as_ref(), self.options.netmode, - ).await + ) + .await } fn generate_token(&self, addr: SocketAddr) -> Vec { - let mut hasher = AHasher::default(); - + match addr.ip() { IpAddr::V4(ip) => ip.octets().hash(&mut hasher), IpAddr::V6(ip) => ip.octets().hash(&mut hasher), } - + self.token_secret.hash(&mut hasher); - + let hash = hasher.finish(); hash.to_le_bytes().to_vec() } - + fn validate_token(&self, token: &[u8], addr: SocketAddr) -> bool { if token.len() != 8 { return false; @@ -787,25 +841,25 @@ impl DHTServer { } /// 发送 DHT find_node 查询消息 -/// +/// /// 这是 DHT 协议中的核心操作之一,用于向指定节点查询包含目标 ID 的节点信息。 /// 该方法构建符合 BEP5 (BitTorrent DHT Protocol) 规范的消息并异步发送。 -/// +/// /// # 参数 -/// +/// /// * `addr` - 目标节点的 Socket 地址 /// * `target` - 要查找的目标节点 ID (20 字节) /// * `sender_id` - 发送者的节点 ID (20 字节),用于标识自己 /// * `socket` - IPv4 UDP socket 的引用 /// * `socket_v6` - IPv6 UDP socket 的可选引用(仅在双栈模式下需要) /// * `netmode` - 网络模式:仅 IPv4、仅 IPv6 或双栈模式 -/// +/// /// # 返回值 -/// +/// /// 返回 `Result<()>`,成功时返回 `Ok(())`,失败时返回错误信息 -/// +/// /// # 消息格式 -/// +/// /// 构建的 DHT 消息格式如下: /// ```bencode /// { @@ -818,9 +872,9 @@ impl DHTServer { /// } /// } /// ``` -/// +/// /// # 网络模式处理 -/// +/// /// * `Ipv4Only`: 始终使用 IPv4 socket /// * `Ipv6Only`: 始终使用 IPv4 socket(IPv6 模式下 socket 实际是 IPv6) /// * `DualStack`: 根据目标地址类型自动选择 IPv4 或 IPv6 socket @@ -834,14 +888,30 @@ async fn send_find_node_impl( ) -> Result<()> { // 构建查询参数 let mut args = std::collections::HashMap::new(); - args.insert(b"id".to_vec(), serde_bencode::value::Value::Bytes(sender_id.to_vec())); - args.insert(b"target".to_vec(), serde_bencode::value::Value::Bytes(target.to_vec())); + args.insert( + b"id".to_vec(), + serde_bencode::value::Value::Bytes(sender_id.to_vec()), + ); + args.insert( + b"target".to_vec(), + serde_bencode::value::Value::Bytes(target.to_vec()), + ); // 构建完整的 DHT 消息 - let mut msg: std::collections::HashMap = std::collections::HashMap::new(); - msg.insert("t".to_string(), serde_bencode::value::Value::Bytes(vec![0, 1])); // 事务 ID - msg.insert("y".to_string(), serde_bencode::value::Value::Bytes(b"q".to_vec())); // 消息类型:查询 - msg.insert("q".to_string(), serde_bencode::value::Value::Bytes(b"find_node".to_vec())); // 查询类型 + let mut msg: std::collections::HashMap = + std::collections::HashMap::new(); + msg.insert( + "t".to_string(), + serde_bencode::value::Value::Bytes(vec![0, 1]), + ); // 事务 ID + msg.insert( + "y".to_string(), + serde_bencode::value::Value::Bytes(b"q".to_vec()), + ); // 消息类型:查询 + msg.insert( + "q".to_string(), + serde_bencode::value::Value::Bytes(b"find_node".to_vec()), + ); // 查询类型 msg.insert("a".to_string(), serde_bencode::value::Value::Dict(args)); // 参数字典 // 将消息编码为 bencode 格式并发送 @@ -856,7 +926,7 @@ async fn send_find_node_impl( } else { socket } - }, + } }; // 异步发送 UDP 数据包 #[cfg(feature = "metrics")] @@ -875,37 +945,37 @@ fn generate_random_id() -> Vec { } /// 生成邻居目标节点 ID -/// +/// /// 该方法用于生成一个"看起来像"远程节点 ID 但实际基于本地节点 ID 的邻居节点 ID。 /// 这是 DHT 协议中的一个重要优化策略,用于提高查询成功率和保护节点 ID 隐私。 -/// +/// /// # 工作原理 -/// +/// /// 1. 取远程节点 ID 的前 6 个字节作为前缀(如果远程 ID 长度足够) /// 2. 用本地节点 ID 的剩余部分填充 /// 3. 如果本地 ID 不够长,用随机字节填充到 20 字节(标准 DHT 节点 ID 长度) -/// +/// /// 这样生成的 ID 在 ID 空间中既接近远程节点(前 6 字节相同),又基于本地节点 /// (后续字节来自本地 ID),从而在 DHT 路由时更容易获得相关响应。 -/// +/// /// # 参数 -/// +/// /// * `remote_id` - 远程节点的 ID(通常是查询目标节点或请求方的 ID) /// * `local_id` - 本地节点的 ID(通常是自己真实的节点 ID) -/// +/// /// # 返回值 -/// +/// /// 返回一个 20 字节的节点 ID Vec,其前 6 字节来自 `remote_id`,后续字节来自 `local_id` -/// +/// /// # 使用场景 -/// +/// /// 1. **发送查询时**:使用邻居 ID 作为发送者 ID,让远程节点认为查询来自一个接近目标 ID 的节点, /// 从而返回更相关的节点列表 /// 2. **发送响应时**:使用邻居 ID 作为响应中的节点 ID,保护真实本地 ID 的隐私, /// 同时提高返回节点的相关性 -/// +/// /// # 示例 -/// +/// /// ``` /// // 假设: /// // remote_id = [0x01, 0x02, 0x03, 0x04, 0x05, 0x06, ...] diff --git a/src/sharded.rs b/src/sharded.rs index a5709f2..574e6bd 100644 --- a/src/sharded.rs +++ b/src/sharded.rs @@ -14,33 +14,34 @@ pub struct ShardedBloom { 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 { + + 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); @@ -68,26 +69,25 @@ impl NodeQueueShard { capacity, } } - + fn push(&mut self, node: NodeTuple) { if self.index.contains(&node.addr) { return; } - if self.queue.len() >= self.capacity { - if let Some(removed) = self.queue.pop_front() { - self.index.remove(&removed.addr); - } + if self.queue.len() >= self.capacity + && let Some(removed) = self.queue.pop_front() { + self.index.remove(&removed.addr); } self.index.insert(node.addr); self.queue.push_back(node); } - + fn pop_batch(&mut self, count: usize) -> Vec { let actual_count = count.min(self.queue.len()); let mut nodes = Vec::with_capacity(actual_count); - + for _ in 0..actual_count { if let Some(node) = self.queue.pop_front() { self.index.remove(&node.addr); @@ -96,11 +96,11 @@ impl NodeQueueShard { } nodes } - + fn len(&self) -> usize { self.queue.len() } - + fn is_empty(&self) -> bool { self.queue.is_empty() } @@ -113,22 +113,26 @@ pub struct ShardedNodeQueue { impl ShardedNodeQueue { pub fn new(total_capacity: usize) -> Self { + #[allow(clippy::manual_div_ceil)] let capacity_per_shard = (total_capacity + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - + let shards_v4 = (0..QUEUE_SHARD_COUNT) .map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard))) .collect(); - + let shards_v6 = (0..QUEUE_SHARD_COUNT) .map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard))) .collect(); - - Self { shards_v4, shards_v6 } + + Self { + shards_v4, + shards_v6, + } } - + pub fn push(&self, node: NodeTuple) { let shard_idx = self.addr_to_shard(&node.addr); - + if node.addr.is_ipv6() { let mut shard = self.shards_v6[shard_idx].lock().unwrap(); shard.push(node); @@ -137,11 +141,12 @@ impl ShardedNodeQueue { shard.push(node); } } - + pub fn pop_batch(&self, count: usize, filter_ipv6: Option) -> Vec { let mut result = Vec::with_capacity(count); + #[allow(clippy::manual_div_ceil)] let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - + match filter_ipv6 { Some(true) => { for shard in &self.shards_v6 { @@ -152,7 +157,7 @@ impl ShardedNodeQueue { let nodes = s.pop_batch(per_shard); result.extend(nodes); } - }, + } Some(false) => { for shard in &self.shards_v4 { if result.len() >= count { @@ -162,101 +167,102 @@ impl ShardedNodeQueue { let nodes = s.pop_batch(per_shard); result.extend(nodes); } - }, + } None => { for i in 0..QUEUE_SHARD_COUNT { if result.len() >= count { break; } - + let mut s4 = self.shards_v4[i].lock().unwrap(); let nodes4 = s4.pop_batch(per_shard / 2); result.extend(nodes4); drop(s4); - + if result.len() >= count { break; } - + let mut s6 = self.shards_v6[i].lock().unwrap(); let nodes6 = s6.pop_batch(per_shard / 2); result.extend(nodes6); drop(s6); } - }, + } } - + result } - + pub fn get_random_nodes(&self, count: usize, filter_ipv6: Option) -> Vec { match filter_ipv6 { - Some(true) => { - self.get_random_nodes_from_shards(&self.shards_v6, count) - }, - Some(false) => { - self.get_random_nodes_from_shards(&self.shards_v4, count) - }, + Some(true) => self.get_random_nodes_from_shards(&self.shards_v6, count), + Some(false) => self.get_random_nodes_from_shards(&self.shards_v4, count), None => { let count_v4 = count / 2; let count_v6 = count - count_v4; let mut result = Vec::with_capacity(count); - + result.extend(self.get_random_nodes_from_shards(&self.shards_v4, count_v4)); result.extend(self.get_random_nodes_from_shards(&self.shards_v6, count_v6)); - + result - }, + } } } - - fn get_random_nodes_from_shards(&self, shards: &[Mutex], count: usize) -> Vec { + + fn get_random_nodes_from_shards( + &self, + shards: &[Mutex], + count: usize, + ) -> Vec { use rand::Rng; let mut rng = rand::thread_rng(); - + if count <= 16 { let mut result = Vec::with_capacity(count); + #[allow(clippy::manual_div_ceil)] let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - + for shard in shards { if result.len() >= count { break; } - + let s = shard.lock().unwrap(); let shard_len = s.queue.len(); - + if shard_len == 0 { continue; } - + let to_take = per_shard.min(shard_len).min(count - result.len()); - + let mut indices: Vec = (0..shard_len).collect(); - + for i in 0..to_take { let j = rng.gen_range(i..shard_len); indices.swap(i, j); } - - for i in 0..to_take { - if let Some(node) = s.queue.get(indices[i]) { + + for &idx in indices.iter().take(to_take) { + if let Some(node) = s.queue.get(idx) { result.push(node.clone()); } } } - + result } else { let mut result = Vec::with_capacity(count); let mut seen = 0usize; - + for shard in shards { let s = shard.lock().unwrap(); - + for node in s.queue.iter() { seen += 1; - + if result.len() < count { result.push(node.clone()); } else { @@ -267,49 +273,51 @@ impl ShardedNodeQueue { } } } - + result } } - + pub fn len(&self) -> usize { - let len_v4: usize = self.shards_v4 + let len_v4: usize = self + .shards_v4 .iter() .map(|shard| shard.lock().unwrap().len()) .sum(); - let len_v6: usize = self.shards_v6 + let len_v6: usize = self + .shards_v6 .iter() .map(|shard| shard.lock().unwrap().len()) .sum(); len_v4 + len_v6 } - + pub fn is_empty(&self) -> bool { - let empty_v4 = self.shards_v4 + let empty_v4 = self + .shards_v4 .iter() .all(|shard| shard.lock().unwrap().is_empty()); - let empty_v6 = self.shards_v6 + let empty_v6 = self + .shards_v6 .iter() .all(|shard| shard.lock().unwrap().is_empty()); empty_v4 && empty_v6 } - + pub fn is_empty_for(&self, filter_ipv6: Option) -> bool { match filter_ipv6 { - Some(true) => { - self.shards_v6 - .iter() - .all(|shard| shard.lock().unwrap().is_empty()) - }, - Some(false) => { - self.shards_v4 - .iter() - .all(|shard| shard.lock().unwrap().is_empty()) - }, + Some(true) => self + .shards_v6 + .iter() + .all(|shard| shard.lock().unwrap().is_empty()), + Some(false) => self + .shards_v4 + .iter() + .all(|shard| shard.lock().unwrap().is_empty()), None => self.is_empty(), } } - + #[inline] fn addr_to_shard(&self, addr: &SocketAddr) -> usize { let hash = match addr.ip() { @@ -325,4 +333,3 @@ impl ShardedNodeQueue { hash % QUEUE_SHARD_COUNT } } - diff --git a/src/types.rs b/src/types.rs index 81dbc11..6e15b55 100644 --- a/src/types.rs +++ b/src/types.rs @@ -1,18 +1,13 @@ use serde::{Deserialize, Serialize}; -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] pub enum NetMode { Ipv4Only, Ipv6Only, + #[default] DualStack, } -impl Default for NetMode { - fn default() -> Self { - Self::DualStack - } -} - #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TorrentInfo { pub info_hash: String, @@ -85,4 +80,4 @@ impl Default for DHTOptions { node_queue_capacity: 100000, } } -} \ No newline at end of file +}