diff --git a/Cargo.toml b/Cargo.toml index 17dedd3..c81eb39 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -29,7 +29,7 @@ bytes = "1.0" bloomfilter = "1.0" ahash = "0.8" serde_bytes = "0.11.19" -metrics = "0.24" +metrics = { version = "0.24", optional = true } [dev-dependencies] tracing = "0.1" @@ -38,6 +38,7 @@ mimalloc = "0.1" [features] default = [] +metrics = ["dep:metrics"] mimalloc = [] [[example]] diff --git a/src/lib.rs b/src/lib.rs index 7747cd1..332244d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -12,7 +12,6 @@ pub use types::{DHTOptions, FileInfo, TorrentInfo, NetMode}; pub use sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple}; pub use scheduler::MetadataScheduler; -// 重新导出常用类型 pub mod prelude { pub use crate::error::{DHTError, Result}; pub use crate::server::DHTServer; diff --git a/src/metadata.rs b/src/metadata.rs index 40e9b9f..8f03274 100644 --- a/src/metadata.rs +++ b/src/metadata.rs @@ -28,12 +28,8 @@ impl RbitFetcher { info_hash: &[u8; 20], peer_addr: SocketAddr, ) -> Option<(String, u64, Vec)> { - let info_hash_hex = hex::encode(info_hash); - log::debug!("[Metadata] 开始获取: {} @ {}", info_hash_hex, peer_addr); - let peer_id = PeerId::generate(); - // DHT 网络很不稳定,如果 3 秒连不上,基本就是连不上了,不要浪费时间 let mut conn = match timeout( Duration::from_secs(3), PeerConnection::connect(peer_addr, *info_hash, *peer_id.as_bytes()), diff --git a/src/protocol.rs b/src/protocol.rs index 7ef342b..6ce934e 100644 --- a/src/protocol.rs +++ b/src/protocol.rs @@ -4,9 +4,9 @@ use serde::Deserialize; #[allow(dead_code)] pub struct DhtMessage { pub t: serde_bytes::ByteBuf, - #[allow(dead_code)] // 用于快速预检查,不在反序列化后使用 + #[allow(dead_code)] pub y: String, - #[allow(dead_code)] // 用于快速预检查,不在反序列化后使用 + #[allow(dead_code)] pub q: Option, pub a: Option, pub r: Option, diff --git a/src/scheduler.rs b/src/scheduler.rs index 40ede20..91e6548 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -10,29 +10,16 @@ use std::time::Duration; type TorrentCallback = Arc; type MetadataFetchCallback = Arc std::pin::Pin + Send>> + Send + Sync>; -/// 元数据调度器(优雅版:Worker 池 + Channel) -/// 负责管理元数据获取队列和任务调度 pub struct MetadataScheduler { - // 输入通道 hash_rx: mpsc::Receiver, - - // 配置 max_queue_size: usize, max_concurrent: usize, - - // 元数据获取器 fetcher: Arc, - - // 回调 callback: Arc>>, on_metadata_fetch: Arc>>, - - // 统计(使用 Atomic 支持多线程访问) total_received: Arc, total_dropped: Arc, total_dispatched: Arc, - - // 共享的队列长度计数器(用于向 Server 反馈背压) queue_len: Arc, } @@ -44,7 +31,7 @@ impl MetadataScheduler { max_concurrent: usize, callback: Arc>>, on_metadata_fetch: Arc>>, - queue_len: Arc, // 新增参数 + queue_len: Arc, ) -> Self { Self { hash_rx, @@ -60,44 +47,37 @@ impl MetadataScheduler { } } - /// 设置 torrent 回调 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) { - // 创建任务队列(channel 自带背压) let (task_tx, task_rx) = mpsc::channel::(self.max_queue_size); let task_rx = Arc::new(Mutex::new(task_rx)); - // 启动 Worker 池 for worker_id in 0..self.max_concurrent { let task_rx = task_rx.clone(); let fetcher = self.fetcher.clone(); let callback = self.callback.clone(); let on_metadata_fetch = self.on_metadata_fetch.clone(); let total_dispatched = self.total_dispatched.clone(); - let queue_len = self.queue_len.clone(); // 传递计数器 + let queue_len = self.queue_len.clone(); tokio::spawn(async move { log::trace!("Worker {} 启动", worker_id); loop { - // Worker 从队列取任务(阻塞等待,零延迟) 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); } @@ -106,12 +86,11 @@ impl MetadataScheduler { let hash = match hash { Some(h) => h, - None => break, // Channel 关闭,退出 + None => break, }; total_dispatched.fetch_add(1, Ordering::Relaxed); - // 执行任务 Self::process_hash( hash, &fetcher, @@ -124,7 +103,6 @@ impl MetadataScheduler { }); } - // 主循环:只负责接收 hash 并转发到 worker 队列 #[cfg(debug_assertions)] let mut stats_interval = tokio::time::interval(Duration::from_secs(60)); #[cfg(debug_assertions)] @@ -137,17 +115,14 @@ impl MetadataScheduler { Some(hash) = self.hash_rx.recv() => { self.total_received.fetch_add(1, Ordering::Relaxed); - // 尝试发送到 worker 队列 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, // Channel 关闭 + Err(_) => break, } } @@ -165,26 +140,22 @@ impl MetadataScheduler { Some(hash) => { self.total_received.fetch_add(1, Ordering::Relaxed); - // 尝试发送到 worker 队列 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, // Channel 关闭 + Err(_) => break, } } - None => break, // Channel 关闭 + None => break, } } } } - /// 处理单个 hash(Worker 调用) async fn process_hash( hash: HashDiscovered, fetcher: &Arc, @@ -194,11 +165,10 @@ 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(), - Err(_) => return, // 锁中毒 + Err(_) => return, } }; @@ -208,7 +178,6 @@ impl MetadataScheduler { } } - // 解码 info_hash let info_hash_bytes: [u8; 20] = match hex::decode(&info_hash) { Ok(bytes) if bytes.len() == 20 => { let mut arr = [0u8; 20]; @@ -218,7 +187,6 @@ impl MetadataScheduler { _ => return, }; - // 获取元数据 if let Some((name, total_size, files)) = fetcher.fetch(&info_hash_bytes, peer_addr).await { let metadata = TorrentInfo { info_hash, @@ -234,11 +202,10 @@ impl MetadataScheduler { .as_secs(), }; - // 获取回调快照并释放锁 let maybe_torrent_cb = { match callback.read() { Ok(guard) => guard.clone(), - Err(_) => return, // 锁中毒 + Err(_) => return, } }; @@ -248,7 +215,6 @@ impl MetadataScheduler { } } - /// 输出统计信息(仅在 debug 模式下编译) #[cfg(debug_assertions)] fn print_stats(&self, task_tx: &mpsc::Sender) { let received = self.total_received.load(Ordering::Relaxed); @@ -264,7 +230,6 @@ impl MetadataScheduler { 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 d81f609..6c6ab1d 100644 --- a/src/server.rs +++ b/src/server.rs @@ -24,11 +24,9 @@ const BOOTSTRAP_NODES: &[&str] = &[ "dht.aelitis.com:6881", ]; -// 类型定义 pub type BoxedBoolFuture = Pin + Send>>; pub type MetadataFetchCallback = Arc BoxedBoolFuture + Send + Sync>; -// Hash 发现事件 #[derive(Debug, Clone)] pub struct HashDiscovered { pub info_hash: String, @@ -36,8 +34,6 @@ pub struct HashDiscovered { pub discovered_at: std::time::Instant, } -// --------------------------------------------------------------- - type TorrentCallback = Arc; type FilterCallback = Arc bool + Send + Sync>; @@ -49,18 +45,11 @@ pub struct DHTServer { socket: Arc, socket_v6: Option>, token_secret: Vec, - callback: Arc>>, filter: Arc>>, on_metadata_fetch: Arc>>, - - // 使用分片锁,大幅减少竞争 node_queue: Arc, - - // 发送 hash 发现事件 hash_tx: mpsc::Sender, - - // Metadata 队列长度(用于自适应爬取速度) metadata_queue_len: Arc, max_metadata_queue_size: usize, } @@ -75,9 +64,8 @@ impl DHTServer { let _ = sock.set_reuse_address(true); sock.set_nonblocking(true)?; - // 增加网络缓冲区以应对高QPS - let _ = sock.set_recv_buffer_size(32 * 1024 * 1024); // 32MB - let _ = sock.set_send_buffer_size(8 * 1024 * 1024); // 8MB + 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())?; @@ -88,7 +76,6 @@ impl DHTServer { #[cfg(not(windows))] { let _ = sock.set_reuse_port(true); } let _ = sock.set_reuse_address(true); - // 设置仅IPv6模式(Windows默认是仅IPv6,Linux/Unix需要设置) #[cfg(not(windows))] { let _ = sock.set_only_v6(true); } sock.set_nonblocking(true)?; @@ -101,7 +88,6 @@ impl DHTServer { (Arc::new(UdpSocket::from_std(sock.into())?), None) }, NetMode::DualStack => { - // IPv4 socket let sock_v4 = Socket::new(Domain::IPV4, Type::DGRAM, Some(Protocol::UDP))?; #[cfg(not(windows))] { let _ = sock_v4.set_reuse_port(true); } @@ -113,13 +99,12 @@ impl DHTServer { sock_v4.bind(&addr_v4.into())?; let socket = Arc::new(UdpSocket::from_std(sock_v4.into())?); - // IPv6 socket 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_address(true); #[cfg(not(windows))] - { let _ = sock_v6.set_only_v6(true); } // 仅IPv6,避免与IPv4冲突 + { 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); @@ -135,22 +120,15 @@ impl DHTServer { let mut rng = rand::thread_rng(); let token_secret: Vec = (0..10).map(|_| rng.r#gen::()).collect(); - // 使用分片队列 - // 队列容量:从配置获取 let node_queue = ShardedNodeQueue::new(options.node_queue_capacity); - // ----------------------------------------------------------- - // 内部初始化 MetadataScheduler - // ----------------------------------------------------------- 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 scheduler = MetadataScheduler::new( @@ -163,13 +141,13 @@ impl DHTServer { metadata_queue_len.clone(), ); - // 启动 Scheduler tokio::spawn(async move { scheduler.run().await; }); + let max_metadata_queue_size = options.max_metadata_queue_size; let server = Self { - options: options.clone(), + options, node_id: node_id.clone(), socket, socket_v6, @@ -180,7 +158,7 @@ impl DHTServer { filter: Arc::new(RwLock::new(None)), hash_tx, metadata_queue_len, - max_metadata_queue_size: options.max_metadata_queue_size, + max_metadata_queue_size, }; Ok(server) @@ -190,32 +168,23 @@ impl DHTServer { Ok(self.socket.local_addr()?) } - /// 验证地址类型是否与当前 netmode 配置匹配 - /// - /// 防御性编程:虽然 socket 层面理论上不应该接收到不匹配的数据包, - /// 但在某些特殊情况下(如系统配置、双栈模式切换等)可能会有问题。 - /// 此方法确保在应用层也进行验证,避免处理不匹配的地址类型。 fn is_addr_allowed(&self, addr: &SocketAddr) -> bool { match self.options.netmode { NetMode::Ipv4Only => addr.is_ipv4(), NetMode::Ipv6Only => addr.is_ipv6(), - NetMode::DualStack => true, // 双栈模式接受所有地址类型 + NetMode::DualStack => true, } } - /// 根据目标地址选择合适的socket fn select_socket(&self, addr: &SocketAddr) -> &Arc { match self.options.netmode { NetMode::Ipv4Only => { - // IPv4Only 模式:只有 IPv4 socket &self.socket }, NetMode::Ipv6Only => { - // IPv6Only 模式:只有 IPv6 socket &self.socket }, NetMode::DualStack => { - // 双栈模式:根据地址类型选择 if addr.is_ipv6() { self.socket_v6.as_ref().unwrap_or(&self.socket) } else { @@ -225,25 +194,6 @@ impl DHTServer { } } - /// 设置元数据获取前的检查回调 - /// - /// 此回调在发现新的 info_hash 后,但在实际连接对等端获取元数据之前执行。 - /// 你可以在这里进行去重检查(如查询数据库),返回 `true` 表示继续获取,`false` 表示跳过。 - /// - /// # 注意事项 - /// - 回调是在 `MetadataScheduler` 的 Worker 线程中异步执行的(通过 `.await`)。 - /// - 支持耗时操作(如数据库查询),但请注意 Worker 数量限制(默认 500)。 - /// - 如果回调执行过慢,可能会导致任务队列堆积。 - /// - /// # 示例 - /// ```rust,ignore - /// server.on_metadata_fetch(|hash| async move { - /// // 检查数据库是否存在 - /// // let exists = db.has(hash).await; - /// // !exists - /// true - /// }); - /// ``` pub fn on_metadata_fetch(&self, callback: F) where F: Fn(String) -> Fut + Send + Sync + 'static, @@ -254,41 +204,10 @@ impl DHTServer { })); } - /// 设置成功获取到种子信息的回调 - /// - /// 当成功从对等端下载并解析出种子元数据(Metadata)后调用。 - /// - /// # 注意事项 - /// - 此回调是在 Worker 线程中同步执行的。 - /// - 如果包含耗时操作(如写入大量数据或复杂计算),**必须**在回调内部手动使用 `tokio::spawn`。 - /// - 否则会阻塞当前的元数据获取 Worker,降低系统吞吐量。 - /// - /// # 示例 - /// ```rust,ignore - /// server.on_torrent(|info| { - /// // 简单操作可以直接做 - /// println!("Got torrent: {}", info.name); - /// - /// // 耗时操作建议 spawn - /// tokio::spawn(async move { - /// save_to_db(info).await; - /// }); - /// }); - /// ``` pub fn on_torrent(&self, callback: F) where F: Fn(TorrentInfo) + Send + Sync + 'static { *self.callback.write().unwrap() = Some(Arc::new(callback)); } - /// 设置 Hash 过滤器 - /// - /// 在处理 `announce_peer` 消息时,用于快速判断是否应该处理该 Hash。 - /// 这通常用于布隆过滤器之前的黑名单或白名单机制。 - /// - /// # 注意事项 - /// - 此回调是在 UDP 处理线程中**同步执行**的。 - /// - **绝对禁止**执行任何耗时操作(如 IO、数据库查询、锁等待)。 - /// - 任何延迟都会直接阻塞网络包的接收,导致丢包。 - /// - 应仅进行纯内存的快速判断。 pub fn set_filter(&self, filter: F) where F: Fn(&str) -> bool + Send + Sync + 'static { *self.filter.write().unwrap() = Some(Arc::new(filter)); } @@ -310,30 +229,23 @@ impl DHTServer { let mut loop_tick = 0; loop { - // 🚀 自适应爬取速度:根据 Metadata 队列负载调整爬取策略 let queue_len = server.metadata_queue_len.load(Ordering::Relaxed); let queue_pressure = queue_len as f64 / server.max_metadata_queue_size as f64; - // 动态计算批次大小和休眠时间 let (batch_size, sleep_duration) = if queue_pressure < 0.8 { - // 🟡 黄区:队列有压力,适度减速 (200, Duration::from_millis(10)) } else if queue_pressure < 0.95 { - // 🟠 橙区:队列高压,大幅减速 (20, Duration::from_millis(500)) } else { - // 🔴 红区:队列爆满,暂停主动爬取 (0, Duration::from_millis(1000)) }; - // 根据配置决定从哪个队列获取节点 let filter_ipv6 = match server.options.netmode { NetMode::Ipv4Only => Some(false), NetMode::Ipv6Only => Some(true), NetMode::DualStack => None, }; - // 检查对应队列是否为空 let queue_empty = server.node_queue.is_empty_for(filter_ipv6); let nodes_batch = { @@ -354,13 +266,30 @@ impl DHTServer { } if let Some(nodes) = nodes_batch { + let node_id = server.node_id.clone(); + 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 server_clone = server.clone(); + let node_id_clone = node_id.clone(); + let socket_clone = socket.clone(); + 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, &server_clone.node_id); + let neighbor_id = generate_neighbor_target(&node_id_for_target, &node_id_clone); let random_target = generate_random_id(); - let _ = server_clone.send_find_node(node.addr, &random_target, &neighbor_id).await; + let _ = send_find_node_impl( + node_addr, + &random_target, + &neighbor_id, + &socket_clone, + socket_v6_clone.as_ref(), + netmode, + ).await; drop(permit); }); } @@ -407,15 +336,12 @@ impl DHTServer { loop { match socket.recv_from(&mut buf).await { Ok((size, addr)) => { - // 🛡️ 安全检查1:拒绝异常大的包(DHT 消息通常 < 2KB) if size > 8192 { #[cfg(debug_assertions)] log::trace!("⚠️ 拒绝异常大的 UDP 包: {} 字节 from {}", size, addr); continue; } - // 🛡️ 安全检查2:快速检查是否是有效的 Bencode 字典 - // DHT KRPC 消息(BEP-5)必须是字典,首字符必须是 'd' if size == 0 || buf[0] != b'd' { continue; } @@ -441,7 +367,6 @@ impl DHTServer { } }); - // IPv6 接收任务 if let Some(socket_v6) = socket_v6 { let senders_v6 = senders_for_v6; tokio::spawn(async move { @@ -478,10 +403,8 @@ impl DHTServer { } async fn handle_message(&self, data: &[u8], addr: SocketAddr) -> Result<()> { - // 🛡️ 验证地址类型是否与当前 netmode 配置匹配 - // 防御性编程:虽然 socket 层面理论上不应该接收到不匹配的数据包, - // 但在某些特殊情况下(如系统配置、双栈模式切换等)可能会有问题 if !self.is_addr_allowed(&addr) { + #[cfg(debug_assertions)] log::trace!("⚠️ 拒绝不匹配的地址类型: {} (当前模式: {:?})", addr, self.options.netmode); return Ok(()); } @@ -550,7 +473,6 @@ impl DHTServer { #[cfg(debug_assertions)] log::debug!("🔥 新 Hash: {} 来自 {}", hash_hex, addr); - // 解耦:发送 hash 发现事件 let port = if let Some(implied) = args.implied_port { if implied != 0 { addr.port() } else { args.port.unwrap_or(0) } } else { @@ -564,7 +486,6 @@ impl DHTServer { discovered_at: std::time::Instant::now(), }; - // 使用 try_send,队列满时直接丢弃(背压) if let Err(_) = self.hash_tx.try_send(event) { #[cfg(debug_assertions)] log::debug!("⚠️ Hash 队列满,丢弃 hash"); @@ -575,11 +496,9 @@ impl DHTServer { } async fn handle_response(&self, response: &DhtResponse) -> Result<()> { - // 处理 IPv4 节点 if let Some(nodes_bytes) = &response.nodes { self.process_compact_nodes(nodes_bytes); } - // 处理 IPv6 节点 if let Some(nodes6_bytes) = &response.nodes6 { self.process_compact_nodes_v6(nodes6_bytes); } @@ -587,14 +506,12 @@ impl DHTServer { } fn process_compact_nodes(&self, nodes_bytes: &[u8]) { - // 根据配置决定是否处理IPv4节点 if self.options.netmode == NetMode::Ipv6Only { return; } 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]]); @@ -607,7 +524,6 @@ impl DHTServer { } fn process_compact_nodes_v6(&self, nodes_bytes: &[u8]) { - // 根据配置决定是否处理IPv6节点 if self.options.netmode == NetMode::Ipv4Only { return; } @@ -618,10 +534,9 @@ impl DHTServer { let port = u16::from_be_bytes([chunk[36], chunk[37]]); let ip_bytes: [u8; 16] = match chunk[20..36].try_into() { Ok(b) => b, - Err(_) => continue, // 如果转换失败(理论上不会),跳过该节点 + Err(_) => continue, }; let ip = Ipv6Addr::from(ip_bytes); - // 过滤掉不可用地址 (组播, 未指定等) if !ip.is_unspecified() && !ip.is_multicast() { let addr = SocketAddr::new(IpAddr::V6(ip), port); self.node_queue.push(NodeTuple { id, addr }); @@ -651,15 +566,13 @@ impl DHTServer { r_dict.insert(b"token".to_vec(), serde_bencode::value::Value::Bytes(token)); if query_type == "get_peers" || query_type == "find_node" { - // 根据配置和请求方IP类型决定需要获取的节点类型 let requestor_is_ipv6 = addr.is_ipv6(); let filter_ipv6 = match self.options.netmode { - NetMode::Ipv4Only => Some(false), // 只要 IPv4 - NetMode::Ipv6Only => Some(true), // 只要 IPv6 - NetMode::DualStack => Some(requestor_is_ipv6), // 双栈模式:根据请求方IP类型返回对应类型的节点 + NetMode::Ipv4Only => Some(false), + 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(); @@ -667,13 +580,11 @@ impl DHTServer { for node in nodes { match node.addr.ip() { - // IPv4 节点 IpAddr::V4(ip) => { 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()); }, - // IPv6 节点 IpAddr::V6(ip) => { nodes6_data.extend_from_slice(&node.id); nodes6_data.extend_from_slice(&ip.octets()); @@ -682,15 +593,11 @@ impl DHTServer { } } - // 根据请求方IP类型返回对应类型的节点 - // 在单栈模式下,get_random_nodes 已经过滤了节点类型,所以这里直接根据请求方类型返回即可 if requestor_is_ipv6 { - // 请求方是IPv6:返回IPv6节点 if !nodes6_data.is_empty() { r_dict.insert(b"nodes6".to_vec(), serde_bencode::value::Value::Bytes(nodes6_data)); } } else { - // 请求方是IPv4:返回IPv4节点 if !nodes_data.is_empty() { r_dict.insert(b"nodes".to_vec(), serde_bencode::value::Value::Bytes(nodes_data)); } @@ -714,7 +621,6 @@ impl DHTServer { match tokio::net::lookup_host(node).await { Ok(addrs) => { for addr in addrs { - // 根据配置过滤地址 match self.options.netmode { NetMode::Ipv4Only => { if addr.is_ipv6() { continue; } @@ -723,7 +629,6 @@ impl DHTServer { if addr.is_ipv4() { continue; } }, NetMode::DualStack => { - // 双栈模式,接受所有地址 }, } let _ = self.send_find_node(addr, &target, &self.node_id).await; @@ -735,36 +640,27 @@ impl DHTServer { } async fn send_find_node(&self, addr: SocketAddr, target: &[u8], sender_id: &[u8]) -> 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())); - - let mut msg: std::collections::HashMap = std::collections::HashMap::new(); - msg.insert("t".to_string(), serde_bencode::value::Value::Bytes(vec![0, 1])); - 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)); - - if let Ok(encoded) = serde_bencode::to_bytes(&msg) { - let _ = self.select_socket(&addr).send_to(&encoded, addr).await; - } - Ok(()) + send_find_node_impl( + addr, + target, + sender_id, + &self.socket, + self.socket_v6.as_ref(), + self.options.netmode, + ).await } fn generate_token(&self, addr: SocketAddr) -> Vec { let mut hasher = AHasher::default(); - // Hash IP地址 match addr.ip() { IpAddr::V4(ip) => ip.octets().hash(&mut hasher), IpAddr::V6(ip) => ip.octets().hash(&mut hasher), } - // Hash 密钥 self.token_secret.hash(&mut hasher); - // 返回 8 字节 token let hash = hasher.finish(); hash.to_le_bytes().to_vec() } @@ -778,6 +674,41 @@ impl DHTServer { } } +async fn send_find_node_impl( + addr: SocketAddr, + target: &[u8], + sender_id: &[u8], + socket: &Arc, + socket_v6: Option<&Arc>, + netmode: NetMode, +) -> 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())); + + let mut msg: std::collections::HashMap = std::collections::HashMap::new(); + msg.insert("t".to_string(), serde_bencode::value::Value::Bytes(vec![0, 1])); + 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)); + + if let Ok(encoded) = serde_bencode::to_bytes(&msg) { + let selected_socket = match netmode { + NetMode::Ipv4Only => socket, + NetMode::Ipv6Only => socket, + NetMode::DualStack => { + if addr.is_ipv6() { + socket_v6.unwrap_or(socket) + } else { + socket + } + }, + }; + let _ = selected_socket.send_to(&encoded, addr).await; + } + Ok(()) +} + fn generate_random_id() -> Vec { let mut rng = rand::thread_rng(); (0..20).map(|_| rng.r#gen::()).collect() diff --git a/src/sharded.rs b/src/sharded.rs index 8409583..a5709f2 100644 --- a/src/sharded.rs +++ b/src/sharded.rs @@ -1,29 +1,18 @@ -// 分片锁实现 - 大幅减少锁竞争,提升并发性能 -// -// 核心思想:1个大锁 → N个小锁 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; // 32个布隆过滤器分片 -const QUEUE_SHARD_COUNT: usize = 16; // 16个队列分片 +const BLOOM_SHARD_COUNT: usize = 32; +const QUEUE_SHARD_COUNT: usize = 16; -// ==================== 分片布隆过滤器 ==================== - -/// 分片布隆过滤器 - 减少锁竞争 -/// -/// 将单个布隆过滤器拆分为32个分片,每个分片独立锁 -/// 不同的hash会落到不同的分片上,大幅减少竞争 pub struct ShardedBloom { shards: Vec>>, count: AtomicUsize, } impl ShardedBloom { - /// 创建新的分片布隆过滤器 pub fn new_for_fp_rate(expected_items: usize, fp_rate: f64) -> Self { let items_per_shard = (expected_items + BLOOM_SHARD_COUNT - 1) / BLOOM_SHARD_COUNT; @@ -37,43 +26,34 @@ impl ShardedBloom { } } - /// 检查并设置元素(原子操作) 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 } - /// 获取实际发现的唯一 InfoHash 数量 pub fn number_of_bits(&self) -> u64 { self.count.load(Ordering::Relaxed) as u64 } - /// 根据hash计算分片索引 #[inline] fn hash_to_shard(&self, hash: &[u8; 20]) -> usize { - // 使用hash的前两个字节计算分片 let idx = (hash[0] as usize) | ((hash[1] as usize) << 8); idx % BLOOM_SHARD_COUNT } } -// ==================== 分片节点队列 ==================== - -/// 节点信息 #[derive(Debug, Clone)] pub struct NodeTuple { pub id: Vec, pub addr: SocketAddr, } -/// 单个队列分片 struct NodeQueueShard { queue: VecDeque, index: HashSet, @@ -94,7 +74,6 @@ impl NodeQueueShard { return; } - // 如果满了,移除最早的一个(保持流动性,优胜劣汰) if self.queue.len() >= self.capacity { if let Some(removed) = self.queue.pop_front() { self.index.remove(&removed.addr); @@ -127,14 +106,12 @@ impl NodeQueueShard { } } -/// 分片节点队列 - 支持高并发,IPv4 和 IPv6 节点分开存储 pub struct ShardedNodeQueue { - shards_v4: Vec>, // IPv4 节点分片 - shards_v6: Vec>, // IPv6 节点分片 + shards_v4: Vec>, + shards_v6: Vec>, } impl ShardedNodeQueue { - /// 创建新的分片队列 pub fn new(total_capacity: usize) -> Self { let capacity_per_shard = (total_capacity + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; @@ -149,7 +126,6 @@ impl ShardedNodeQueue { Self { shards_v4, shards_v6 } } - /// 添加节点(根据地址类型自动存入对应队列) pub fn push(&self, node: NodeTuple) { let shard_idx = self.addr_to_shard(&node.addr); @@ -162,18 +138,12 @@ impl ShardedNodeQueue { } } - /// 批量弹出节点 - /// - /// # Arguments - /// * `count` - 需要获取的节点数量 - /// * `filter_ipv6` - 如果为 `Some(true)`,只从 IPv6 队列获取;如果为 `Some(false)`,只从 IPv4 队列获取;如果为 `None`,从两个队列混合获取 pub fn pop_batch(&self, count: usize, filter_ipv6: Option) -> Vec { let mut result = Vec::with_capacity(count); let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; match filter_ipv6 { Some(true) => { - // 只从 IPv6 队列获取 for shard in &self.shards_v6 { if result.len() >= count { break; @@ -184,7 +154,6 @@ impl ShardedNodeQueue { } }, Some(false) => { - // 只从 IPv4 队列获取 for shard in &self.shards_v4 { if result.len() >= count { break; @@ -195,13 +164,11 @@ impl ShardedNodeQueue { } }, None => { - // 混合模式:从两个队列交替获取 for i in 0..QUEUE_SHARD_COUNT { if result.len() >= count { break; } - // 从 IPv4 分片获取 let mut s4 = self.shards_v4[i].lock().unwrap(); let nodes4 = s4.pop_batch(per_shard / 2); result.extend(nodes4); @@ -211,7 +178,6 @@ impl ShardedNodeQueue { break; } - // 从 IPv6 分片获取 let mut s6 = self.shards_v6[i].lock().unwrap(); let nodes6 = s6.pop_batch(per_shard / 2); result.extend(nodes6); @@ -223,22 +189,15 @@ impl ShardedNodeQueue { result } - /// 获取随机节点(用于DHT响应) - /// # Arguments - /// * `count` - 需要获取的节点数量 - /// * `filter_ipv6` - 如果为 `Some(true)`,只返回 IPv6 节点;如果为 `Some(false)`,只返回 IPv4 节点;如果为 `None`,返回所有节点(混合) pub fn get_random_nodes(&self, count: usize, filter_ipv6: Option) -> Vec { match filter_ipv6 { Some(true) => { - // 只要 IPv6 节点 self.get_random_nodes_from_shards(&self.shards_v6, count) }, Some(false) => { - // 只要 IPv4 节点 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); @@ -251,12 +210,10 @@ impl ShardedNodeQueue { } } - /// 从指定的分片组中获取随机节点 fn get_random_nodes_from_shards(&self, shards: &[Mutex], count: usize) -> Vec { use rand::Rng; let mut rng = rand::thread_rng(); - // 🚀 策略1:小规模请求用快速路径(最常见:8个节点) if count <= 16 { let mut result = Vec::with_capacity(count); let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; @@ -273,19 +230,15 @@ impl ShardedNodeQueue { continue; } - // 从当前分片随机选择最多 per_shard 个节点 let to_take = per_shard.min(shard_len).min(count - result.len()); - // 生成随机索引(不重复) let mut indices: Vec = (0..shard_len).collect(); - // 只 shuffle 前 to_take 个(部分 shuffle,Fisher-Yates 优化) for i in 0..to_take { let j = rng.gen_range(i..shard_len); indices.swap(i, j); } - // 取前 to_take 个索引对应的节点 for i in 0..to_take { if let Some(node) = s.queue.get(indices[i]) { result.push(node.clone()); @@ -295,11 +248,9 @@ impl ShardedNodeQueue { result } else { - // 🚀 策略2:大规模请求用储层采样 let mut result = Vec::with_capacity(count); let mut seen = 0usize; - // 储层采样算法 for shard in shards { let s = shard.lock().unwrap(); @@ -307,10 +258,8 @@ impl ShardedNodeQueue { seen += 1; if result.len() < count { - // 前 count 个直接加入 result.push(node.clone()); } else { - // 后续以 count/seen 的概率替换 let j = rng.gen_range(0..seen); if j < count { result[j] = node.clone(); @@ -323,8 +272,6 @@ impl ShardedNodeQueue { } } - - /// 获取总长度(IPv4 + IPv6) pub fn len(&self) -> usize { let len_v4: usize = self.shards_v4 .iter() @@ -337,7 +284,6 @@ impl ShardedNodeQueue { len_v4 + len_v6 } - /// 检查是否为空 pub fn is_empty(&self) -> bool { let empty_v4 = self.shards_v4 .iter() @@ -348,20 +294,14 @@ impl ShardedNodeQueue { empty_v4 && empty_v6 } - /// 检查指定地址族的队列是否为空 - /// - /// # Arguments - /// * `filter_ipv6` - 如果为 `Some(true)`,检查 IPv6 队列;如果为 `Some(false)`,检查 IPv4 队列;如果为 `None`,检查两个队列 pub fn is_empty_for(&self, filter_ipv6: Option) -> bool { match filter_ipv6 { Some(true) => { - // 检查 IPv6 队列 self.shards_v6 .iter() .all(|shard| shard.lock().unwrap().is_empty()) }, Some(false) => { - // 检查 IPv4 队列 self.shards_v4 .iter() .all(|shard| shard.lock().unwrap().is_empty()) @@ -370,10 +310,8 @@ impl ShardedNodeQueue { } } - /// 根据地址计算分片索引 #[inline] fn addr_to_shard(&self, addr: &SocketAddr) -> usize { - // 使用端口和IP最后一个字节 let hash = match addr.ip() { std::net::IpAddr::V4(ip) => { let octets = ip.octets(); diff --git a/src/types.rs b/src/types.rs index 074fa4c..81dbc11 100644 --- a/src/types.rs +++ b/src/types.rs @@ -1,13 +1,9 @@ use serde::{Deserialize, Serialize}; -/// 网络模式配置 #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum NetMode { - /// 仅使用 IPv4 Ipv4Only, - /// 仅使用 IPv6 Ipv6Only, - /// 双栈(同时支持 IPv4 和 IPv6) DualStack, } @@ -17,7 +13,6 @@ impl Default for NetMode { } } -/// 完整的种子信息(包含元数据) #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TorrentInfo { pub info_hash: String, @@ -30,7 +25,6 @@ pub struct TorrentInfo { pub timestamp: u64, } -/// 文件信息 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct FileInfo { pub path: String, @@ -62,45 +56,32 @@ fn format_bytes(bytes: u64) -> String { format!("{:.2} {}", size, UNITS[unit_index]) } -/// DHT 服务器配置 #[derive(Debug, Clone)] pub struct DHTOptions { - /// DHT 端口 pub port: u16, - /// 是否自动获取元数据 pub auto_metadata: bool, - /// 元数据获取超时(秒) pub metadata_timeout: u64, - /// 元数据获取队列大小(背压限制) pub max_metadata_queue_size: usize, - /// 并发元数据获取工作线程数 pub max_metadata_worker_count: usize, - /// 网络模式配置(仅IPv4、仅IPv6、或双栈) pub netmode: NetMode, - /// 节点队列容量(默认 100000) pub node_queue_capacity: usize, } impl Default for DHTOptions { fn default() -> Self { Self { - port: 6881, // BitTorrent DHT 默认端口 + port: 6881, auto_metadata: true, - // 缩短超时,快速失败,不等待慢节点 metadata_timeout: 3, - // 加大队列,防止流量高峰丢包 max_metadata_queue_size: 100000, - // 提高并发,模拟 Node.js 的高并发 IO max_metadata_worker_count: 1000, - // 默认双栈 netmode: NetMode::Ipv4Only, - // 节点队列容量:100000 个节点(扩容以适应 DHT 网络裂变速度) node_queue_capacity: 100000, } }