删除注释,优化部分clone使用

This commit is contained in:
桥下红药
2026-01-17 18:22:36 +08:00
parent ab4247a6e8
commit 217a934980
8 changed files with 92 additions and 281 deletions
+2 -1
View File
@@ -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]]
-1
View File
@@ -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;
-4
View File
@@ -28,12 +28,8 @@ impl RbitFetcher {
info_hash: &[u8; 20],
peer_addr: SocketAddr,
) -> Option<(String, u64, Vec<FileInfo>)> {
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()),
+2 -2
View File
@@ -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<String>,
pub a: Option<DhtArgs>,
pub r: Option<DhtResponse>,
+8 -43
View File
@@ -10,29 +10,16 @@ use std::time::Duration;
type TorrentCallback = Arc<dyn Fn(TorrentInfo) + Send + Sync>;
type MetadataFetchCallback = Arc<dyn Fn(String) -> std::pin::Pin<Box<dyn std::future::Future<Output = bool> + Send>> + Send + Sync>;
/// 元数据调度器(优雅版:Worker 池 + Channel
/// 负责管理元数据获取队列和任务调度
pub struct MetadataScheduler {
// 输入通道
hash_rx: mpsc::Receiver<HashDiscovered>,
// 配置
max_queue_size: usize,
max_concurrent: usize,
// 元数据获取器
fetcher: Arc<RbitFetcher>,
// 回调
callback: Arc<RwLock<Option<TorrentCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
// 统计(使用 Atomic 支持多线程访问)
total_received: Arc<AtomicU64>,
total_dropped: Arc<AtomicU64>,
total_dispatched: Arc<AtomicU64>,
// 共享的队列长度计数器(用于向 Server 反馈背压)
queue_len: Arc<AtomicUsize>,
}
@@ -44,7 +31,7 @@ impl MetadataScheduler {
max_concurrent: usize,
callback: Arc<RwLock<Option<TorrentCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
queue_len: Arc<AtomicUsize>, // 新增参数
queue_len: Arc<AtomicUsize>,
) -> 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::<HashDiscovered>(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,
}
}
}
}
/// 处理单个 hashWorker 调用)
async fn process_hash(
hash: HashDiscovered,
fetcher: &Arc<RbitFetcher>,
@@ -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<HashDiscovered>) {
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}%)",
+75 -144
View File
@@ -24,11 +24,9 @@ const BOOTSTRAP_NODES: &[&str] = &[
"dht.aelitis.com:6881",
];
// 类型定义
pub type BoxedBoolFuture = Pin<Box<dyn Future<Output = bool> + Send>>;
pub type MetadataFetchCallback = Arc<dyn Fn(String) -> 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<dyn Fn(TorrentInfo) + Send + Sync>;
type FilterCallback = Arc<dyn Fn(&str) -> bool + Send + Sync>;
@@ -49,18 +45,11 @@ pub struct DHTServer {
socket: Arc<UdpSocket>,
socket_v6: Option<Arc<UdpSocket>>,
token_secret: Vec<u8>,
callback: Arc<RwLock<Option<TorrentCallback>>>,
filter: Arc<RwLock<Option<FilterCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
// 使用分片锁,大幅减少竞争
node_queue: Arc<ShardedNodeQueue>,
// 发送 hash 发现事件
hash_tx: mpsc::Sender<HashDiscovered>,
// Metadata 队列长度(用于自适应爬取速度)
metadata_queue_len: Arc<AtomicUsize>,
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默认是仅IPv6Linux/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<u8> = (0..10).map(|_| rng.r#gen::<u8>()).collect();
// 使用分片队列
// 队列容量:从配置获取
let node_queue = ShardedNodeQueue::new(options.node_queue_capacity);
// -----------------------------------------------------------
// 内部初始化 MetadataScheduler
// -----------------------------------------------------------
let (hash_tx, hash_rx) = mpsc::channel::<HashDiscovered>(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<UdpSocket> {
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<F, Fut>(&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<F>(&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<F>(&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<String, serde_bencode::value::Value> = 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<u8> {
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<UdpSocket>,
socket_v6: Option<&Arc<UdpSocket>>,
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<String, serde_bencode::value::Value> = 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<u8> {
let mut rng = rand::thread_rng();
(0..20).map(|_| rng.r#gen::<u8>()).collect()
+4 -66
View File
@@ -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<Mutex<Bloom<[u8; 20]>>>,
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<u8>,
pub addr: SocketAddr,
}
/// 单个队列分片
struct NodeQueueShard {
queue: VecDeque<NodeTuple>,
index: HashSet<SocketAddr>,
@@ -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<Mutex<NodeQueueShard>>, // IPv4 节点分片
shards_v6: Vec<Mutex<NodeQueueShard>>, // IPv6 节点分片
shards_v4: Vec<Mutex<NodeQueueShard>>,
shards_v6: Vec<Mutex<NodeQueueShard>>,
}
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<bool>) -> Vec<NodeTuple> {
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<bool>) -> Vec<NodeTuple> {
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<NodeQueueShard>], count: usize) -> Vec<NodeTuple> {
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<usize> = (0..shard_len).collect();
// 只 shuffle 前 to_take 个(部分 shuffleFisher-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>) -> 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();
+1 -20
View File
@@ -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,
}
}