From 7cc88969184b50e4f59816c7b6032cef07cb27e6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=A1=A5=E4=B8=8B=E7=BA=A2=E8=8D=AF?= Date: Mon, 22 Dec 2025 21:32:14 +0800 Subject: [PATCH] init --- Cargo.toml | 59 +++++ src/error.rs | 24 ++ src/lib.rs | 21 ++ src/main.rs | 105 ++++++++ src/metadata.rs | 171 +++++++++++++ src/protocol.rs | 33 +++ src/scheduler.rs | 264 ++++++++++++++++++++ src/server.rs | 626 +++++++++++++++++++++++++++++++++++++++++++++++ src/sharded.rs | 289 ++++++++++++++++++++++ src/types.rs | 80 ++++++ 10 files changed, 1672 insertions(+) create mode 100644 Cargo.toml create mode 100644 src/error.rs create mode 100644 src/lib.rs create mode 100644 src/main.rs create mode 100644 src/metadata.rs create mode 100644 src/protocol.rs create mode 100644 src/scheduler.rs create mode 100644 src/server.rs create mode 100644 src/sharded.rs create mode 100644 src/types.rs diff --git a/Cargo.toml b/Cargo.toml new file mode 100644 index 0000000..9e1ebef --- /dev/null +++ b/Cargo.toml @@ -0,0 +1,59 @@ +[package] +name = "dht_crawler" +version = "3.0.0" +edition = "2021" + +[dependencies] +tokio = { version = "1.35", features = ["full"] } +serde = { version = "1.0", features = ["derive"] } +serde_bencode = "0.2" # 更成熟的 bencode 库,支持 UTF-8 +sha1 = "0.10" +hex = "0.4" +rand = "0.8" + +log = "0.4" +tracing = "0.1.43" +tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt"] } +tracing-appender = "0.2" +tracing-log = "0.2" + + +thiserror = "1.0" + +socket2 = { version = "0.5", features = ["all"] } +rbit = "0.1" +bytes = "1.0" + +mimalloc = "0.1" +bloomfilter = "1.0" +ahash = "0.8" # 快速哈希算法(比 SHA1 快 10倍+) + +# 可选:用于 Web API +actix-web = { version = "4.4", optional = true } +encoding_rs = { version = "0.8.35", optional = true } +serde_bytes = "0.11.19" + + +[features] +default = [] +web = ["actix-web"] +encoding = ["dep:encoding_rs"] + +# ==================== 性能优化配置 ==================== + +[profile.release] +opt-level = 3 # 最高优化级别 +lto = "fat" # 链接时优化(LTO)- 显著提升性能 +codegen-units = 1 # 单编译单元 - 更好的优化但编译慢 +panic = "abort" # panic时直接终止 - 减少二进制大小 +strip = true # 移除调试符号 - 减小二进制 + +# 开发时快速编译 +[profile.dev] +opt-level = 0 +debug = true + +# 性能测试配置 +[profile.bench] +inherits = "release" +debug = true # 保留符号以便性能分析 diff --git a/src/error.rs b/src/error.rs new file mode 100644 index 0000000..5505397 --- /dev/null +++ b/src/error.rs @@ -0,0 +1,24 @@ +use thiserror::Error; + +#[derive(Error, Debug)] +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, +} + +pub type Result = std::result::Result; diff --git a/src/lib.rs b/src/lib.rs new file mode 100644 index 0000000..e1a2c38 --- /dev/null +++ b/src/lib.rs @@ -0,0 +1,21 @@ +mod error; +mod server; +pub mod protocol; +pub mod types; +pub mod metadata; // 公开 metadata 模块 +mod sharded; // 分片锁模块 +pub mod scheduler; // 元数据调度器 + +pub use error::{DHTError, Result}; +pub use server::{DHTServer, HashDiscovered}; +pub use types::{DHTOptions, FileInfo, TorrentInfo}; +pub use sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple}; +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::scheduler::MetadataScheduler; +} diff --git a/src/main.rs b/src/main.rs new file mode 100644 index 0000000..ce47962 --- /dev/null +++ b/src/main.rs @@ -0,0 +1,105 @@ +#[global_allocator] +static GLOBAL: MiMalloc = MiMalloc; + +use dht_crawler::prelude::*; +use std::sync::Arc; +use mimalloc::MiMalloc; +use tracing_subscriber::EnvFilter; +use std::sync::atomic::{AtomicUsize, Ordering}; + +#[tokio::main] +async fn main() -> Result<()> { + + if std::env::var("RUST_LOG").is_err() { + std::env::set_var("RUST_LOG", "info"); + } + + // 直接输出到 stdout,避免 _guard 被 drop 导致日志丢失 + tracing_subscriber::fmt() + .with_env_filter(EnvFilter::from_default_env()) + .with_ansi(true) + .init(); + + let options = DHTOptions { + port: 45452, + auto_metadata: true, + metadata_timeout: 3, // ✅ 快速超时,快速失败 + max_metadata_queue_size: 100000, // ✅ 大缓冲区(防止饱和) + max_metadata_worker_count: 1000, // ✅ 激进并发(最大化吞吐) + }; + + // 统计计数器 + let torrent_count = Arc::new(AtomicUsize::new(0)); + let torrent_count_clone = torrent_count.clone(); + + // 🚀 初始化 DHT Server + log::info!("🔧 正在初始化 DHT Server..."); + let server = DHTServer::new(options.clone()).await?; + + log::info!("🚀 DHT Server 启动,监听端口: {}", options.port); + + // 设置 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 { + // torrent.files.iter() + // .map(|f| format!("{} ({})", f.path, format_size(f.size))) + // .collect::>() + // .join(", ") + // } else { + // format!("{}个文件", torrent.files.len()) + // }; + // + // log::info!( + // "🎉 [{}] {} ({}, {})", + // count, + // torrent.name, + // format_size(total_size), + // files_display + // ); + }); + + // 设置元数据获取前的检查回调 + server.on_metadata_fetch(|_hash| async move { + true + }); + + server.set_filter(|_hash| { + true + }); + + server.on_duplicate(|_hash| { + + }); + + // 启动监控任务 + let dht_monitor = server.clone(); + let count_monitor = torrent_count.clone(); + tokio::spawn(async move { + let mut interval = tokio::time::interval(std::time::Duration::from_secs(5)); + let start_time = std::time::Instant::now(); + + loop { + interval.tick().await; + let success_fetch = count_monitor.load(Ordering::Relaxed); + let uptime = start_time.elapsed().as_secs(); + + // ✅ 监控:布隆过滤器的位使用情况反映了爬虫的活跃度 + log::info!( + "📊 [监控] 时长: {}s | 成功抓取: ✨ {} | 活跃指纹: {}", + uptime, success_fetch, dht_monitor.get_seen_count() + ); + + if uptime > 0 && success_fetch > 0 { + let speed = (success_fetch as f64) / (uptime as f64 / 60.0); + log::info!("📈 平均抓取速度: {:.2} 种子/分钟", speed); + } + } + }); + + server.start().await?; + Ok(()) +} \ No newline at end of file diff --git a/src/metadata.rs b/src/metadata.rs new file mode 100644 index 0000000..8023ea7 --- /dev/null +++ b/src/metadata.rs @@ -0,0 +1,171 @@ +use std::collections::BTreeMap; +use std::net::SocketAddr; +use std::time::Duration; +use bytes::Bytes; +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; + +#[derive(Clone)] +pub struct RbitFetcher { + timeout: Duration, +} + +impl RbitFetcher { + pub fn new(timeout_secs: u64) -> Self { + Self { + timeout: Duration::from_secs(if timeout_secs == 0 { 15 } else { timeout_secs }), + } + } + + pub async fn fetch( + &self, + 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(); + + // 🔥 修改点:缩短连接超时到 3 秒 + // DHT 网络很不稳定,如果 3 秒连不上,基本就是连不上了,不要浪费时间 + let mut conn = match timeout( + Duration::from_secs(3), + PeerConnection::connect(peer_addr, *info_hash, *peer_id.as_bytes()), + ).await { + Ok(Ok(c)) => c, + Ok(Err(_)) => return None, + Err(_) => return None, + }; + + if !conn.supports_extension { + return None; + } + + let my_ut_metadata_id = 1; + 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; + } else { + return None; + } + + let mut metadata_size = 0; + let mut remote_ut_metadata_id = 0; + let mut pieces: BTreeMap = BTreeMap::new(); + let mut request_sent = false; + + let result = timeout(self.timeout, async { + loop { + let msg = conn.receive().await.ok()?; + match msg { + Message::Extended { id, payload } => { + 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; + } + if let Some(ext_id) = remote_hs.get_extension_id("ut_metadata") { + remote_ut_metadata_id = ext_id; + } + } + if metadata_size > 0 && remote_ut_metadata_id > 0 && !request_sent { + if metadata_size > 10 * 1024 * 1024 { return None; } + + let count = metadata_piece_count(metadata_size as usize); + for i in 0..count { + let req = MetadataMessage::request(i as u32); + if let Ok(encoded) = req.encode() { + let _ = conn.send(Message::Extended { id: remote_ut_metadata_id, payload: encoded }).await; + } + } + 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 { + pieces.insert(meta_msg.piece, data); + } + } + } + if metadata_size > 0 { + let total_received: usize = pieces.values().map(|p| p.len()).sum(); + if total_received >= metadata_size as usize { + let mut full_data = Vec::with_capacity(metadata_size as usize); + let count = metadata_piece_count(metadata_size as usize); + let mut success = true; + for i in 0..count { + if let Some(p) = pieces.get(&(i as u32)) { + full_data.extend_from_slice(p); + } else { + success = false; break; + } + } + if success { + let info_hash_copy = *info_hash; + let full_data_clone = full_data.clone(); + let is_valid = tokio::task::spawn_blocking(move || { + let mut hasher = Sha1::new(); + hasher.update(&full_data_clone); + let digest: [u8; 20] = hasher.finalize().into(); + digest == info_hash_copy + }).await.unwrap_or(false); + + if is_valid { + return Some(full_data); + } + return None; + } + } + } + } + } + _ => {} + } + } + }).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(); + 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); } + } + } + 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()) { + total_size = len as u64; + file_list.push(FileInfo { path: name.clone(), size: total_size }); + } + if total_size > 0 { return Some((name, total_size, file_list)); } + } + } + None + } + _ => None, + } + } +} \ No newline at end of file diff --git a/src/protocol.rs b/src/protocol.rs new file mode 100644 index 0000000..f0fa4a1 --- /dev/null +++ b/src/protocol.rs @@ -0,0 +1,33 @@ +use serde::Deserialize; + +#[derive(Deserialize, Debug)] +#[allow(dead_code)] +pub struct DhtMessage { + pub t: serde_bytes::ByteBuf, + #[allow(dead_code)] // 用于快速预检查,不在反序列化后使用 + pub y: String, + #[allow(dead_code)] // 用于快速预检查,不在反序列化后使用 + pub q: Option, + pub a: Option, + pub r: Option, +} + +#[derive(Deserialize, Debug)] +pub struct DhtArgs { + pub id: Option, + pub target: Option, + pub info_hash: Option, + pub token: Option, + pub port: Option, + pub implied_port: Option, +} + +#[derive(Deserialize, Debug)] +pub struct DhtResponse { + #[serde(default)] + #[allow(dead_code)] + pub id: Option, + #[serde(default)] + pub nodes: Option, +} + diff --git a/src/scheduler.rs b/src/scheduler.rs new file mode 100644 index 0000000..df0a200 --- /dev/null +++ b/src/scheduler.rs @@ -0,0 +1,264 @@ +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 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, +} + +impl MetadataScheduler { + pub fn new( + hash_rx: mpsc::Receiver, + fetcher: Arc, + max_queue_size: usize, + max_concurrent: usize, + callback: Arc>>, + on_metadata_fetch: Arc>>, + queue_len: Arc, // 新增参数 + ) -> Self { + Self { + hash_rx, + max_queue_size, + max_concurrent, + fetcher, + callback, + on_metadata_fetch, + total_received: Arc::new(AtomicU64::new(0)), + total_dropped: Arc::new(AtomicU64::new(0)), + total_dispatched: Arc::new(AtomicU64::new(0)), + queue_len, + } + } + + /// 设置 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(); // 传递计数器 + + 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); + } + h + }; + + let hash = match hash { + Some(h) => h, + None => break, // Channel 关闭,退出 + }; + + total_dispatched.fetch_add(1, Ordering::Relaxed); + + // 执行任务 + Self::process_hash( + hash, + &fetcher, + &callback, + &on_metadata_fetch, + ).await; + } + + log::trace!("Worker {} 退出", worker_id); + }); + } + + // 主循环:只负责接收 hash 并转发到 worker 队列 + let mut stats_interval = tokio::time::interval(Duration::from_secs(60)); + stats_interval.tick().await; + + loop { + tokio::select! { + 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 关闭 + } + } + + _ = stats_interval.tick() => { + self.print_stats(&task_tx); + } + + else => break, + } + } + + // log::info!("🛑 Metadata 调度器停止"); + } + + /// 处理单个 hash(Worker 调用) + async fn process_hash( + hash: HashDiscovered, + fetcher: &Arc, + callback: &Arc>>, + on_metadata_fetch: &Arc>>, + ) { + 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, // 锁中毒 + } + }; + + if let Some(f) = maybe_check_fn { + if !f(info_hash.clone()).await { + return; + } + } + + // 解码 info_hash + let info_hash_bytes: [u8; 20] = match hex::decode(&info_hash) { + Ok(bytes) if bytes.len() == 20 => { + let mut arr = [0u8; 20]; + arr.copy_from_slice(&bytes); + arr + } + _ => return, + }; + + // 获取元数据 + if let Some((name, total_size, files)) = fetcher.fetch(&info_hash_bytes, peer_addr).await { + let metadata = TorrentInfo { + info_hash, + name, + total_size, + files, + magnet_link: format!("magnet:?xt=urn:btih:{}", hash.info_hash), + peers: vec![peer_addr.to_string()], + piece_length: 0, + timestamp: std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .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); + } + } + } + + /// 输出统计信息 + 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}%)", + queue_size, + self.max_queue_size, + queue_pressure, + received, + dispatched, + dropped, + drop_rate + ); + } else { + log::info!( + "📊 Metadata 调度器统计:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)", + queue_size, + self.max_queue_size, + queue_pressure, + received, + dispatched, + dropped, + drop_rate + ); + } + } +} diff --git a/src/server.rs b/src/server.rs new file mode 100644 index 0000000..dd4bb19 --- /dev/null +++ b/src/server.rs @@ -0,0 +1,626 @@ +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::sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple}; +use rand::Rng; +use ahash::AHasher; +use std::hash::{Hash, Hasher}; +use std::net::{IpAddr, 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 socket2::{Socket, Domain, Type, Protocol}; +use std::pin::Pin; +use std::future::Future; + +const BOOTSTRAP_NODES: &[&str] = &[ + "router.bittorrent.com:6881", + "dht.transmissionbt.com:6881", + "router.utorrent.com:6881", + "dht.aelitis.com:6881", +]; + +// 类型定义 +pub type BoxedBoolFuture = Pin + Send>>; +pub type MetadataFetchCallback = Arc BoxedBoolFuture + Send + Sync>; + +// Hash 发现事件 +/// DHT Server 发现 hash 后发送此事件,由独立的 MetadataScheduler 处理 +#[derive(Debug, Clone)] +pub struct HashDiscovered { + pub info_hash: String, + pub peer_addr: SocketAddr, + pub discovered_at: std::time::Instant, +} + +// --------------------------------------------------------------- + +type TorrentCallback = Arc; +type FilterCallback = Arc bool + Send + Sync>; +type DuplicateCallback = Arc; + +#[derive(Clone)] +pub struct DHTServer { + #[allow(dead_code)] + options: DHTOptions, + node_id: Vec, + socket: Arc, + token_secret: Vec, + + // 这些回调现在与 MetadataScheduler 共享 + callback: Arc>>, + filter: Arc>>, + on_duplicate: Arc>>, + on_metadata_fetch: Arc>>, + + // 使用分片锁,大幅减少竞争 + node_queue: Arc, + seen_hashes: Arc, + + // 发送 hash 发现事件 + hash_tx: mpsc::Sender, + + // Metadata 队列长度(用于自适应爬取速度) + metadata_queue_len: Arc, + max_metadata_queue_size: usize, +} + +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 addr: SocketAddr = format!("0.0.0.0:{}", options.port).parse().unwrap(); + socket.bind(&addr.into())?; + let socket = UdpSocket::from_std(socket.into())?; + + let node_id = generate_random_id(); + let mut rng = rand::thread_rng(); + let token_secret: Vec = (0..10).map(|_| rng.gen()).collect(); + + // 使用分片队列和分片布隆过滤器 + // 队列容量:100000 个节点(扩容以适应 DHT 网络裂变速度) + let node_queue = ShardedNodeQueue::new(100000); + + // 布隆过滤器:预期500万元素,0.1%误判率 + // 内存使用:约 90MB(32分片 × 2.8MB) + let bloom = ShardedBloom::new_for_fp_rate(5_000_000, 0.001); + + // ----------------------------------------------------------- + // 内部初始化 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( + hash_rx, + fetcher, + options.max_metadata_queue_size, + options.max_metadata_worker_count, + callback.clone(), + on_metadata_fetch.clone(), + metadata_queue_len.clone(), + ); + + // 启动 Scheduler + tokio::spawn(async move { + scheduler.run().await; + }); + + let server = Self { + options: options.clone(), + node_id: node_id.clone(), + socket: Arc::new(socket), + token_secret, + callback, + on_metadata_fetch, + node_queue: Arc::new(node_queue), + seen_hashes: Arc::new(bloom), + filter: Arc::new(RwLock::new(None)), + on_duplicate: Arc::new(RwLock::new(None)), + hash_tx, + metadata_queue_len, + max_metadata_queue_size: options.max_metadata_queue_size, + }; + + Ok(server) + } + + pub fn local_addr(&self) -> Result { + Ok(self.socket.local_addr()?) + } + + /// 设置元数据获取前的检查回调 + /// + /// 此回调在发现新的 info_hash 后,但在实际连接对等端获取元数据之前执行。 + /// 你可以在这里进行去重检查(如查询数据库),返回 `true` 表示继续获取,`false` 表示跳过。 + /// + /// # 注意事项 + /// - 回调是在 `MetadataScheduler` 的 Worker 线程中异步执行的(通过 `.await`)。 + /// - 支持耗时操作(如数据库查询),但请注意 Worker 数量限制(默认 500)。 + /// - 如果回调执行过慢,可能会导致任务队列堆积。 + /// + /// # 示例 + /// ```rust + /// 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, + Fut: Future + Send + 'static, + { + *self.on_metadata_fetch.write().unwrap() = Some(Arc::new(move |hash| { + Box::pin(callback(hash)) + })); + } + + /// 设置成功获取到种子信息的回调 + /// + /// 当成功从对等端下载并解析出种子元数据(Metadata)后调用。 + /// + /// # 注意事项 + /// - 此回调是在 Worker 线程中同步执行的。 + /// - 如果包含耗时操作(如写入大量数据或复杂计算),**必须**在回调内部手动使用 `tokio::spawn`。 + /// - 否则会阻塞当前的元数据获取 Worker,降低系统吞吐量。 + /// + /// # 示例 + /// ```rust + /// 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)); + } + + /// 设置重复 Hash 发现的回调 + /// + /// 当接收到的 Hash 已经被布隆过滤器标记为“已存在”时调用。 + /// + /// # 注意事项 + /// - 库内部已经自动为每次调用包裹了 `tokio::spawn`。 + /// - 因此你可以放心地在回调中执行耗时操作(如数据库记录),而不用担心阻塞 UDP 线程。 + /// - 虽然内部有 spawn,但频繁触发仍会产生大量任务,请注意资源控制。 + pub fn on_duplicate(&self, callback: F) where F: Fn(&str) + Send + Sync + 'static { + *self.on_duplicate.write().unwrap() = Some(Arc::new(callback)); + } + + pub fn get_seen_count(&self) -> usize { + // 分片布隆过滤器的位数统计 + self.seen_hashes.number_of_bits() as usize + } + + pub fn get_node_pool_size(&self) -> usize { + self.node_queue.len() + } + + pub async fn start(&self) -> Result<()> { + + self.start_receiver(); + self.bootstrap().await; + + let server = self.clone(); + + tokio::spawn(async move { + let semaphore = Arc::new(Semaphore::new(2000)); + 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.5 { + // 🟢 绿区:队列空闲,全速爬取 + (200, Duration::from_millis(10)) + } else if queue_pressure < 0.8 { + // 🟡 黄区:队列有压力,适度减速 + (200, Duration::from_millis(20)) + } else if queue_pressure < 0.95 { + // 🟠 橙区:队列高压,大幅减速 + (20, Duration::from_millis(500)) + } else { + // 🔴 红区:队列爆满,暂停主动爬取 + (0, Duration::from_millis(1000)) + }; + + let nodes_batch = { + if server.node_queue.is_empty() || batch_size == 0 { + None + } else { + Some(server.node_queue.pop_batch(batch_size)) + } + }; + + loop_tick += 1; + if nodes_batch.is_none() || loop_tick % 50 == 0 { + server.bootstrap().await; + if nodes_batch.is_none() { + tokio::time::sleep(sleep_duration).await; + continue; + } + } + + if let Some(nodes) = nodes_batch { + for node in nodes { + let permit = semaphore.clone().acquire_owned().await.unwrap(); + let server_clone = server.clone(); + tokio::spawn(async move { + let neighbor_id = generate_neighbor_target(&node.id, &server_clone.node_id); + let random_target = generate_random_id(); + let _ = server_clone.send_find_node(node.addr, &random_target, &neighbor_id).await; + drop(permit); + }); + } + } + + tokio::time::sleep(sleep_duration).await; + } + }); + + std::future::pending::<()>().await; + Ok(()) + } + + fn start_receiver(&self) { + let socket = self.socket.clone(); + let server = self.clone(); + + let num_workers = std::thread::available_parallelism() + .map(|n| n.get()) + .unwrap_or(8); + + let queue_size = 5000; + + let mut senders = Vec::with_capacity(num_workers); + for _ in 0..num_workers { + let (tx, mut rx) = mpsc::channel::<(Vec, SocketAddr)>(queue_size); + senders.push(tx); + + let server_clone = server.clone(); + + tokio::spawn(async move { + while let Some((data, addr)) = rx.recv().await { + let _ = server_clone.handle_message(&data, addr).await; + } + }); + } + + tokio::spawn(async move { + let mut buf = [0u8; 65536]; + let mut next_worker_idx = 0; + + loop { + match socket.recv_from(&mut buf).await { + Ok((size, addr)) => { + // 🛡️ 安全检查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; + } + + let data = buf[..size].to_vec(); + + let tx = &senders[next_worker_idx]; + next_worker_idx = (next_worker_idx + 1) % num_workers; + + match tx.try_send((data, addr)) { + Ok(_) => {}, + Err(mpsc::error::TrySendError::Full(_)) => { + #[cfg(debug_assertions)] + log::trace!("UDP worker queue full, dropping packet"); + }, + Err(_) => { break; } + } + } + Err(_e) => { + tokio::time::sleep(Duration::from_millis(1)).await; + } + } + } + }); + } + + async fn handle_message(&self, data: &[u8], addr: SocketAddr) -> Result<()> { + let msg: DhtMessage = match serde_bencode::from_bytes(data) { + Ok(m) => m, + Err(_) => return Ok(()), + }; + + match msg.y.as_str() { + "q" => { + if let Some(q_type) = &msg.q { + self.handle_query(&msg, q_type.as_bytes(), addr).await?; + } + } + "r" => { + if let Some(response) = &msg.r { + self.handle_response(response).await?; + } + } + _ => {} + } + Ok(()) + } + + async fn handle_query(&self, msg: &DhtMessage, query_type: &[u8], addr: SocketAddr) -> Result<()> { + let args = match &msg.a { + Some(a) => a, + None => return Ok(()), + }; + + 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() + .or(args.info_hash.as_deref()) + .map(|v| v.as_slice()); + + let q_str = std::str::from_utf8(query_type).unwrap_or(""); + + if q_str == "announce_peer" { + self.handle_announce_peer(args, addr).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) { return Ok(()); } + } else { + return Ok(()); + } + + 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(()), + }; + let hash_hex = hex::encode(info_hash_arr); + + // 使用分片布隆过滤器进行高效去重 + let is_duplicate = self.seen_hashes.check_and_set(&info_hash_arr); + + if is_duplicate { + let dup_cb = self.on_duplicate.read().unwrap().clone(); + if let Some(cb) = dup_cb { + let hash_hex_clone = hash_hex.clone(); + tokio::spawn(async move { + cb(&hash_hex_clone); + }); + } + return Ok(()); + } + + let filter_cb = self.filter.read().unwrap().clone(); + if let Some(f) = filter_cb { + if !f(&hash_hex) { return Ok(()); } + } + + #[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 { + args.port.unwrap_or(addr.port()) + }; + + if port > 0 { + let event = HashDiscovered { + info_hash: hash_hex, + peer_addr: SocketAddr::new(addr.ip(), port), + discovered_at: std::time::Instant::now(), + }; + + // 使用 try_send,队列满时直接丢弃(背压) + if let Err(_) = self.hash_tx.try_send(event) { + #[cfg(debug_assertions)] + log::trace!("⚠️ Hash 队列满,丢弃 hash"); + } + } + } + Ok(()) + } + + async fn handle_response(&self, response: &DhtResponse) -> Result<()> { + if let Some(nodes_bytes) = &response.nodes { + self.process_compact_nodes(nodes_bytes); + } + Ok(()) + } + + fn process_compact_nodes(&self, nodes_bytes: &[u8]) { + 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); + + self.node_queue.push(NodeTuple { id, addr }); + } + } + + async fn send_response( + &self, + tid: &[u8], + addr: SocketAddr, + query_type: &str, + sender_id: Option<&[u8]>, + target_id_fallback: Option<&[u8]>, + ) -> Result<()> { + let mut r_dict = std::collections::HashMap::new(); + + let reference_id = sender_id.or(target_id_fallback); + let my_id = if let Some(target) = reference_id { + generate_neighbor_target(target, &self.node_id) + } else { + self.node_id.clone() + }; + + r_dict.insert(b"id".to_vec(), serde_bencode::value::Value::Bytes(my_id)); + let token = self.generate_token(addr); + 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); + + let mut nodes_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, + } + 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)); + } + + 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) { + let _ = self.socket.send_to(&encoded, addr).await; + } + Ok(()) + } + + 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 { + if addr.is_ipv6() { continue; } + 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<()> { + 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.socket.send_to(&encoded, addr).await; + } + Ok(()) + } + + 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() + } + + fn validate_token(&self, token: &[u8], addr: SocketAddr) -> bool { + if token.len() != 8 { + return false; + } + let expected = self.generate_token(addr); + token == expected.as_slice() + } +} + +fn generate_random_id() -> Vec { + let mut rng = rand::thread_rng(); + (0..20).map(|_| rng.gen()).collect() +} + +fn generate_neighbor_target(remote_id: &[u8], local_id: &[u8]) -> Vec { + let mut id = Vec::with_capacity(20); + let prefix_len = std::cmp::min(remote_id.len(), 6); + id.extend_from_slice(&remote_id[..prefix_len]); + if local_id.len() > prefix_len { + id.extend_from_slice(&local_id[prefix_len..]); + } else { + while id.len() < 20 { + id.push(rand::random()); + } + } + id +} diff --git a/src/sharded.rs b/src/sharded.rs new file mode 100644 index 0000000..5d38a42 --- /dev/null +++ b/src/sharded.rs @@ -0,0 +1,289 @@ +// 分片锁实现 - 大幅减少锁竞争,提升并发性能 +// +// 核心思想:1个大锁 → N个小锁 +// 性能提升:预期 3-4 倍 + +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个队列分片 + +// ==================== 分片布隆过滤器 ==================== + +/// 分片布隆过滤器 - 减少锁竞争 +/// +/// 将单个布隆过滤器拆分为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; + + let shards = (0..BLOOM_SHARD_COUNT) + .map(|_| Mutex::new(Bloom::new_for_fp_rate(items_per_shard, fp_rate))) + .collect(); + + Self { + shards, + count: AtomicUsize::new(0), + } + } + + /// 检查并设置元素(原子操作) + pub fn check_and_set(&self, hash: &[u8; 20]) -> bool { + let shard_idx = self.hash_to_shard(hash); + let mut shard = self.shards[shard_idx].lock().unwrap(); + let present = shard.check_and_set(hash); + + // 如果之前不存在,增加计数 + if !present { + self.count.fetch_add(1, Ordering::Relaxed); + } + present + } + + /// 获取实际发现的唯一 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, + capacity: usize, +} + +impl NodeQueueShard { + fn new(capacity: usize) -> Self { + Self { + queue: VecDeque::with_capacity(capacity), + index: HashSet::with_capacity(capacity), + 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); + } + } + + 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); + nodes.push(node); + } + } + nodes + } + + fn len(&self) -> usize { + self.queue.len() + } + + fn is_empty(&self) -> bool { + self.queue.is_empty() + } +} + +/// 分片节点队列 - 支持高并发 +pub struct ShardedNodeQueue { + shards: Vec>, +} + +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) + .map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard))) + .collect(); + + Self { shards } + } + + /// 添加节点 + 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); + } + + /// 批量弹出节点 + pub fn pop_batch(&self, count: usize) -> 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); + } + + result + } + + /// 获取随机节点(用于DHT响应) + /// 🚀 优化:使用储层采样算法,O(n)时间,无需clone全部节点 + pub fn get_random_nodes(&self, 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(); + + 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 + } + + /// 快速路径:小规模随机选择(针对常见的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(); + + 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 + } + + /// 获取总长度 + pub fn len(&self) -> usize { + self.shards + .iter() + .map(|shard| shard.lock().unwrap().len()) + .sum() + } + + /// 检查是否为空 + pub fn is_empty(&self) -> bool { + self.shards + .iter() + .all(|shard| shard.lock().unwrap().is_empty()) + } + + /// 根据地址计算分片索引 + #[inline] + fn addr_to_shard(&self, addr: &SocketAddr) -> usize { + // 使用端口和IP最后一个字节 + let hash = match addr.ip() { + std::net::IpAddr::V4(ip) => { + let octets = ip.octets(); + (octets[3] as usize) ^ (addr.port() as usize) + } + std::net::IpAddr::V6(ip) => { + let octets = ip.octets(); + (octets[15] as usize) ^ (addr.port() as usize) + } + }; + hash % QUEUE_SHARD_COUNT + } +} + diff --git a/src/types.rs b/src/types.rs new file mode 100644 index 0000000..ae6df7e --- /dev/null +++ b/src/types.rs @@ -0,0 +1,80 @@ +use serde::{Deserialize, Serialize}; + +/// 完整的种子信息(包含元数据) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TorrentInfo { + pub info_hash: String, + pub magnet_link: String, + pub name: String, + pub total_size: u64, + pub files: Vec, + pub piece_length: u64, + pub peers: Vec, + pub timestamp: u64, +} + +/// 文件信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FileInfo { + pub path: String, + pub size: u64, +} + +impl TorrentInfo { + pub fn format_size(&self) -> String { + format_bytes(self.total_size) + } +} + +impl FileInfo { + pub fn format_size(&self) -> String { + format_bytes(self.size) + } +} + +fn format_bytes(bytes: u64) -> String { + const UNITS: &[&str] = &["B", "KB", "MB", "GB", "TB"]; + let mut size = bytes as f64; + let mut unit_index = 0; + + while size >= 1024.0 && unit_index < UNITS.len() - 1 { + size /= 1024.0; + unit_index += 1; + } + + 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, +} + +impl Default for DHTOptions { + fn default() -> Self { + Self { + port: 0, + auto_metadata: true, + // 缩短超时,快速失败,不等待慢节点 + metadata_timeout: 10, + // 加大队列,防止流量高峰丢包 + max_metadata_queue_size: 10000, + // 提高并发,模拟 Node.js 的高并发 IO + max_metadata_worker_count: 1000, + } + } +} \ No newline at end of file