优化部分语法警告
This commit is contained in:
@@ -15,3 +15,6 @@ torrents/
|
||||
|
||||
# 日志
|
||||
*.log
|
||||
|
||||
# 个人脚本(不提交到仓库)
|
||||
scripts/
|
||||
|
||||
+13
-17
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -32,4 +32,3 @@ pub struct DhtResponse {
|
||||
#[serde(default)]
|
||||
pub nodes6: Option<serde_bytes::ByteBuf>,
|
||||
}
|
||||
|
||||
|
||||
+42
-38
@@ -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
@@ -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 socket(IPv6 模式下 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
@@ -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
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user