优化部分语法警告

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