diff --git a/src/lib.rs b/src/lib.rs index e1a2c38..16f907a 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -8,7 +8,7 @@ pub mod scheduler; // 元数据调度器 pub use error::{DHTError, Result}; pub use server::{DHTServer, HashDiscovered}; -pub use types::{DHTOptions, FileInfo, TorrentInfo}; +pub use types::{DHTOptions, FileInfo, TorrentInfo, NetMode}; pub use sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple}; pub use scheduler::MetadataScheduler; @@ -16,6 +16,6 @@ pub use scheduler::MetadataScheduler; pub mod prelude { pub use crate::error::{DHTError, Result}; pub use crate::server::DHTServer; - pub use crate::types::{DHTOptions, FileInfo, TorrentInfo}; + pub use crate::types::{DHTOptions, FileInfo, TorrentInfo, NetMode}; pub use crate::scheduler::MetadataScheduler; } diff --git a/src/main.rs b/src/main.rs index ce47962..55c2348 100644 --- a/src/main.rs +++ b/src/main.rs @@ -21,11 +21,12 @@ async fn main() -> Result<()> { .init(); let options = DHTOptions { - port: 45452, + port: 12313, auto_metadata: true, metadata_timeout: 3, // ✅ 快速超时,快速失败 max_metadata_queue_size: 100000, // ✅ 大缓冲区(防止饱和) max_metadata_worker_count: 1000, // ✅ 激进并发(最大化吞吐) + netmode: NetMode::DualStack, // 网络模式:Ipv4Only(仅IPv4)、Ipv6Only(仅IPv6)、DualStack(双栈,默认) }; // 统计计数器 diff --git a/src/protocol.rs b/src/protocol.rs index f0fa4a1..7ef342b 100644 --- a/src/protocol.rs +++ b/src/protocol.rs @@ -29,5 +29,7 @@ pub struct DhtResponse { pub id: Option, #[serde(default)] pub nodes: Option, + #[serde(default)] + pub nodes6: Option, } diff --git a/src/server.rs b/src/server.rs index 761d073..a1a8297 100644 --- a/src/server.rs +++ b/src/server.rs @@ -2,12 +2,12 @@ use crate::error::Result; use crate::metadata::RbitFetcher; use crate::protocol::{DhtMessage, DhtArgs, DhtResponse}; use crate::scheduler::MetadataScheduler; -use crate::types::{DHTOptions, TorrentInfo}; +use crate::types::{DHTOptions, TorrentInfo, NetMode}; use crate::sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple}; use rand::Rng; use ahash::AHasher; use std::hash::{Hash, Hasher}; -use std::net::{IpAddr, SocketAddr}; +use std::net::{IpAddr, Ipv6Addr, SocketAddr}; use std::sync::{Arc, RwLock}; use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; @@ -49,6 +49,7 @@ pub struct DHTServer { options: DHTOptions, node_id: Vec, socket: Arc, + socket_v6: Option>, token_secret: Vec, // 这些回调现在与 MetadataScheduler 共享 @@ -71,19 +72,69 @@ pub struct DHTServer { impl DHTServer { pub async fn new(options: DHTOptions) -> Result { - let socket = Socket::new(Domain::IPV4, Type::DGRAM, Some(Protocol::UDP))?; - #[cfg(not(windows))] - { let _ = socket.set_reuse_port(true); } - let _ = socket.set_reuse_address(true); - socket.set_nonblocking(true)?; - - // 增加网络缓冲区以应对高QPS - let _ = socket.set_recv_buffer_size(32 * 1024 * 1024); // 32MB(原16MB) - let _ = socket.set_send_buffer_size(8 * 1024 * 1024); // 8MB(原4MB) + let (socket, socket_v6) = match options.netmode { + 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_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 addr: SocketAddr = format!("0.0.0.0:{}", options.port).parse().unwrap(); - socket.bind(&addr.into())?; - let socket = UdpSocket::from_std(socket.into())?; + 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_address(true); + // 设置仅IPv6模式(Windows默认是仅IPv6,Linux/Unix需要设置) + #[cfg(not(windows))] + { 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 => { + // IPv4 socket + 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_address(true); + sock_v4.set_nonblocking(true)?; + let _ = sock_v4.set_recv_buffer_size(32 * 1024 * 1024); + let _ = sock_v4.set_send_buffer_size(8 * 1024 * 1024); + let addr_v4: SocketAddr = format!("0.0.0.0:{}", options.port).parse().unwrap(); + 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冲突 + 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); + let addr_v6: SocketAddr = format!("[::]:{}", options.port).parse().unwrap(); + sock_v6.bind(&addr_v6.into())?; + let socket_v6 = Some(Arc::new(UdpSocket::from_std(sock_v6.into())?)); + + (socket, socket_v6) + }, + }; let node_id = generate_random_id(); let mut rng = rand::thread_rng(); @@ -129,7 +180,8 @@ impl DHTServer { let server = Self { options: options.clone(), node_id: node_id.clone(), - socket: Arc::new(socket), + socket, + socket_v6, token_secret, callback, on_metadata_fetch, @@ -149,6 +201,41 @@ 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, // 双栈模式接受所有地址类型 + } + } + + /// 根据目标地址选择合适的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 { + &self.socket + } + }, + } + } + /// 设置元数据获取前的检查回调 /// /// 此回调在发现新的 info_hash 后,但在实际连接对等端获取元数据之前执行。 @@ -160,7 +247,7 @@ impl DHTServer { /// - 如果回调执行过慢,可能会导致任务队列堆积。 /// /// # 示例 - /// ```rust,ignore + /// ```rust /// server.on_metadata_fetch(|hash| async move { /// // 检查数据库是否存在 /// // let exists = db.has(hash).await; @@ -188,7 +275,7 @@ impl DHTServer { /// - 否则会阻塞当前的元数据获取 Worker,降低系统吞吐量。 /// /// # 示例 - /// ```rust,ignore + /// ```rust /// server.on_torrent(|info| { /// // 简单操作可以直接做 /// println!("Got torrent: {}", info.name); @@ -250,17 +337,14 @@ impl DHTServer { let mut loop_tick = 0; loop { - // 自适应爬取速度:根据 Metadata 队列负载调整爬取策略 + // 🚀 自适应爬取速度:根据 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.5 { - // 🟢 绿区:队列空闲,全速爬取 - (200, Duration::from_millis(10)) - } else if queue_pressure < 0.8 { + let (batch_size, sleep_duration) = if queue_pressure < 0.8 { // 🟡 黄区:队列有压力,适度减速 - (200, Duration::from_millis(20)) + (200, Duration::from_millis(10)) } else if queue_pressure < 0.95 { // 🟠 橙区:队列高压,大幅减速 (20, Duration::from_millis(500)) @@ -269,11 +353,21 @@ impl DHTServer { (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 = { - if server.node_queue.is_empty() || batch_size == 0 { + if queue_empty || batch_size == 0 { None } else { - Some(server.node_queue.pop_batch(batch_size)) + Some(server.node_queue.pop_batch(batch_size, filter_ipv6)) } }; @@ -309,6 +403,7 @@ impl DHTServer { fn start_receiver(&self) { let socket = self.socket.clone(); + let socket_v6 = self.socket_v6.clone(); let server = self.clone(); let num_workers = std::thread::available_parallelism() @@ -331,6 +426,7 @@ impl DHTServer { }); } + let senders_for_v6 = senders.clone(); tokio::spawn(async move { let mut buf = [0u8; 65536]; let mut next_worker_idx = 0; @@ -340,7 +436,6 @@ impl DHTServer { Ok((size, addr)) => { // 🛡️ 安全检查1:拒绝异常大的包(DHT 消息通常 < 2KB) if size > 8192 { - #[cfg(debug_assertions)] log::trace!("⚠️ 拒绝异常大的 UDP 包: {} 字节 from {}", size, addr); continue; } @@ -371,9 +466,52 @@ impl DHTServer { } } }); + + // IPv6 接收任务 + 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<()> { + // 🛡️ 验证地址类型是否与当前 netmode 配置匹配 + // 防御性编程:虽然 socket 层面理论上不应该接收到不匹配的数据包, + // 但在某些特殊情况下(如系统配置、双栈模式切换等)可能会有问题 + if !self.is_addr_allowed(&addr) { + log::trace!("⚠️ 拒绝不匹配的地址类型: {} (当前模式: {:?})", addr, self.options.netmode); + return Ok(()); + } + let msg: DhtMessage = match serde_bencode::from_bytes(data) { Ok(m) => m, Err(_) => return Ok(()), @@ -469,7 +607,7 @@ impl DHTServer { // 使用 try_send,队列满时直接丢弃(背压) if let Err(_) = self.hash_tx.try_send(event) { #[cfg(debug_assertions)] - log::trace!("⚠️ Hash 队列满,丢弃 hash"); + log::debug!("⚠️ Hash 队列满,丢弃 hash"); } } } @@ -477,13 +615,23 @@ 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); + } Ok(()) } fn process_compact_nodes(&self, nodes_bytes: &[u8]) { + // 根据配置决定是否处理IPv4节点 + if self.options.netmode == NetMode::Ipv6Only { + return; + } + if nodes_bytes.len() % 26 != 0 { return; } // 使用分片队列,直接并发插入(无锁竞争) @@ -498,6 +646,29 @@ impl DHTServer { } } + fn process_compact_nodes_v6(&self, nodes_bytes: &[u8]) { + // 根据配置决定是否处理IPv6节点 + if self.options.netmode == NetMode::Ipv4Only { + return; + } + + 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]]); + let ip_bytes: [u8; 16] = match chunk[20..36].try_into() { + Ok(b) => b, + 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 }); + } + } + } + async fn send_response( &self, tid: &[u8], @@ -520,20 +691,50 @@ impl DHTServer { r_dict.insert(b"token".to_vec(), serde_bencode::value::Value::Bytes(token)); if query_type == "get_peers" || query_type == "find_node" { - // 使用分片队列获取随机节点(无锁竞争) - let nodes = self.node_queue.get_random_nodes(8); + // 根据配置和请求方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类型返回对应类型的节点 + }; + + // 使用分片队列获取随机节点(无锁竞争,带地址族过滤) + let nodes = self.node_queue.get_random_nodes(8, filter_ipv6); let mut nodes_data = Vec::new(); + let mut nodes6_data = Vec::new(); + for node in nodes { - nodes_data.extend_from_slice(&node.id); match node.addr.ip() { - IpAddr::V4(ip) => nodes_data.extend_from_slice(&ip.octets()), - _ => continue, + // 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()); + nodes6_data.extend_from_slice(&node.addr.port().to_be_bytes()); + }, } - nodes_data.extend_from_slice(&node.addr.port().to_be_bytes()); } - r_dict.insert(b"nodes".to_vec(), serde_bencode::value::Value::Bytes(nodes_data)); + // 根据请求方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)); + } + } } let mut response: std::collections::HashMap = std::collections::HashMap::new(); @@ -542,7 +743,7 @@ impl DHTServer { response.insert("r".to_string(), serde_bencode::value::Value::Dict(r_dict)); if let Ok(encoded) = serde_bencode::to_bytes(&response) { - let _ = self.socket.send_to(&encoded, addr).await; + let _ = self.select_socket(&addr).send_to(&encoded, addr).await; } Ok(()) } @@ -553,7 +754,18 @@ impl DHTServer { match tokio::net::lookup_host(node).await { Ok(addrs) => { for addr in addrs { - if addr.is_ipv6() { continue; } + // 根据配置过滤地址 + match self.options.netmode { + NetMode::Ipv4Only => { + if addr.is_ipv6() { continue; } + }, + NetMode::Ipv6Only => { + if addr.is_ipv4() { continue; } + }, + NetMode::DualStack => { + // 双栈模式,接受所有地址 + }, + } let _ = self.send_find_node(addr, &target, &self.node_id).await; } } @@ -574,7 +786,7 @@ impl DHTServer { msg.insert("a".to_string(), serde_bencode::value::Value::Dict(args)); if let Ok(encoded) = serde_bencode::to_bytes(&msg) { - let _ = self.socket.send_to(&encoded, addr).await; + let _ = self.select_socket(&addr).send_to(&encoded, addr).await; } Ok(()) } diff --git a/src/sharded.rs b/src/sharded.rs index 5d38a42..e7f80e7 100644 --- a/src/sharded.rs +++ b/src/sharded.rs @@ -129,9 +129,10 @@ impl NodeQueueShard { } } -/// 分片节点队列 - 支持高并发 +/// 分片节点队列 - 支持高并发,IPv4 和 IPv6 节点分开存储 pub struct ShardedNodeQueue { - shards: Vec>, + shards_v4: Vec>, // IPv4 节点分片 + shards_v6: Vec>, // IPv6 节点分片 } impl ShardedNodeQueue { @@ -139,134 +140,238 @@ impl ShardedNodeQueue { pub fn new(total_capacity: usize) -> Self { let capacity_per_shard = (total_capacity + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - let shards = (0..QUEUE_SHARD_COUNT) + let shards_v4 = (0..QUEUE_SHARD_COUNT) .map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard))) .collect(); - Self { shards } + let shards_v6 = (0..QUEUE_SHARD_COUNT) + .map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard))) + .collect(); + + Self { shards_v4, shards_v6 } } - /// 添加节点 + /// 添加节点(根据地址类型自动存入对应队列) pub fn push(&self, node: NodeTuple) { let shard_idx = self.addr_to_shard(&node.addr); - let mut shard = self.shards[shard_idx].lock().unwrap(); - shard.push(node); + + if node.addr.is_ipv6() { + let mut shard = self.shards_v6[shard_idx].lock().unwrap(); + shard.push(node); + } else { + let mut shard = self.shards_v4[shard_idx].lock().unwrap(); + shard.push(node); + } } /// 批量弹出节点 - pub fn pop_batch(&self, count: usize) -> Vec { + /// + /// # 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; - // 从所有分片获取 - for shard in &self.shards { - if result.len() >= count { - break; - } - - let mut s = shard.lock().unwrap(); - let nodes = s.pop_batch(per_shard); - result.extend(nodes); + match filter_ipv6 { + Some(true) => { + // 只从 IPv6 队列获取 + for shard in &self.shards_v6 { + if result.len() >= count { + break; + } + let mut s = shard.lock().unwrap(); + let nodes = s.pop_batch(per_shard); + result.extend(nodes); + } + }, + Some(false) => { + // 只从 IPv4 队列获取 + for shard in &self.shards_v4 { + if result.len() >= count { + break; + } + let mut s = shard.lock().unwrap(); + let nodes = s.pop_batch(per_shard); + result.extend(nodes); + } + }, + 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); + drop(s4); + + if result.len() >= count { + break; + } + + // 从 IPv6 分片获取 + let mut s6 = self.shards_v6[i].lock().unwrap(); + let nodes6 = s6.pop_batch(per_shard / 2); + result.extend(nodes6); + drop(s6); + } + }, } result } /// 获取随机节点(用于DHT响应) - /// 🚀 优化:使用储层采样算法,O(n)时间,无需clone全部节点 - pub fn get_random_nodes(&self, count: usize) -> Vec { + /// 🚀 优化:IPv4 和 IPv6 分开存储,直接从对应队列获取,无需过滤 + /// + /// # 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); + + 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 { use rand::Rng; let mut rng = rand::thread_rng(); // 🚀 策略1:小规模请求用快速路径(最常见:8个节点) if count <= 16 { - return self.get_random_nodes_fast(count); - } - - // 🚀 策略2:大规模请求用储层采样 - let mut result = Vec::with_capacity(count); - let mut seen = 0usize; - - // 储层采样算法 - for shard in &self.shards { - let s = shard.lock().unwrap(); + let mut result = Vec::with_capacity(count); + let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - for node in s.queue.iter() { - seen += 1; + for shard in shards { + if result.len() >= count { + break; + } - 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(); + let s = shard.lock().unwrap(); + let shard_len = s.queue.len(); + + if shard_len == 0 { + 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()); } } } - } - - result - } - - /// 快速路径:小规模随机选择(针对常见的8节点请求) - fn get_random_nodes_fast(&self, count: usize) -> Vec { - use rand::Rng; - let mut rng = rand::thread_rng(); - let mut result = Vec::with_capacity(count); - - // 从每个分片随机选择几个节点 - let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT; - - for shard in &self.shards { - if result.len() >= count { - break; - } - let s = shard.lock().unwrap(); - let shard_len = s.queue.len(); + result + } else { + // 🚀 策略2:大规模请求用储层采样 + let mut result = Vec::with_capacity(count); + let mut seen = 0usize; - if shard_len == 0 { - 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()); + // 储层采样算法 + for shard in shards { + let s = shard.lock().unwrap(); + + for node in s.queue.iter() { + 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(); + } + } } } + + result } - - result } - /// 获取总长度 + + /// 获取总长度(IPv4 + IPv6) pub fn len(&self) -> usize { - self.shards + let len_v4: usize = self.shards_v4 .iter() .map(|shard| shard.lock().unwrap().len()) - .sum() + .sum(); + 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 { - self.shards + let empty_v4 = self.shards_v4 .iter() - .all(|shard| shard.lock().unwrap().is_empty()) + .all(|shard| shard.lock().unwrap().is_empty()); + let empty_v6 = self.shards_v6 + .iter() + .all(|shard| shard.lock().unwrap().is_empty()); + 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()) + }, + None => self.is_empty(), + } } /// 根据地址计算分片索引 diff --git a/src/types.rs b/src/types.rs index ae6df7e..08cc9f6 100644 --- a/src/types.rs +++ b/src/types.rs @@ -1,5 +1,22 @@ use serde::{Deserialize, Serialize}; +/// 网络模式配置 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum NetMode { + /// 仅使用 IPv4 + Ipv4Only, + /// 仅使用 IPv6 + Ipv6Only, + /// 双栈(同时支持 IPv4 和 IPv6) + DualStack, +} + +impl Default for NetMode { + fn default() -> Self { + Self::DualStack + } +} + /// 完整的种子信息(包含元数据) #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TorrentInfo { @@ -62,6 +79,9 @@ pub struct DHTOptions { /// 并发元数据获取工作线程数 pub max_metadata_worker_count: usize, + + /// 网络模式配置(仅IPv4、仅IPv6、或双栈) + pub netmode: NetMode, } impl Default for DHTOptions { @@ -75,6 +95,8 @@ impl Default for DHTOptions { max_metadata_queue_size: 10000, // 提高并发,模拟 Node.js 的高并发 IO max_metadata_worker_count: 1000, + // 默认双栈 + netmode: NetMode::DualStack, } } } \ No newline at end of file