This commit is contained in:
桥下红药
2025-12-22 21:32:14 +08:00
commit 7cc8896918
10 changed files with 1672 additions and 0 deletions
+59
View File
@@ -0,0 +1,59 @@
[package]
name = "dht_crawler"
version = "3.0.0"
edition = "2021"
[dependencies]
tokio = { version = "1.35", features = ["full"] }
serde = { version = "1.0", features = ["derive"] }
serde_bencode = "0.2" # 更成熟的 bencode 库,支持 UTF-8
sha1 = "0.10"
hex = "0.4"
rand = "0.8"
log = "0.4"
tracing = "0.1.43"
tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt"] }
tracing-appender = "0.2"
tracing-log = "0.2"
thiserror = "1.0"
socket2 = { version = "0.5", features = ["all"] }
rbit = "0.1"
bytes = "1.0"
mimalloc = "0.1"
bloomfilter = "1.0"
ahash = "0.8" # 快速哈希算法(比 SHA1 快 10倍+)
# 可选:用于 Web API
actix-web = { version = "4.4", optional = true }
encoding_rs = { version = "0.8.35", optional = true }
serde_bytes = "0.11.19"
[features]
default = []
web = ["actix-web"]
encoding = ["dep:encoding_rs"]
# ==================== 性能优化配置 ====================
[profile.release]
opt-level = 3 # 最高优化级别
lto = "fat" # 链接时优化(LTO)- 显著提升性能
codegen-units = 1 # 单编译单元 - 更好的优化但编译慢
panic = "abort" # panic时直接终止 - 减少二进制大小
strip = true # 移除调试符号 - 减小二进制
# 开发时快速编译
[profile.dev]
opt-level = 0
debug = true
# 性能测试配置
[profile.bench]
inherits = "release"
debug = true # 保留符号以便性能分析
+24
View File
@@ -0,0 +1,24 @@
use thiserror::Error;
#[derive(Error, Debug)]
pub enum DHTError {
#[error("网络错误: {0}")]
Network(#[from] std::io::Error),
#[error("Bencode 解析错误: {0}")]
Bencode(String),
#[error("协议错误: {0}")]
Protocol(String),
#[error("超时")]
Timeout,
#[error("未找到 Peer")]
NoPeersAvailable,
#[error("元数据验证失败")]
InvalidMetadata,
}
pub type Result<T> = std::result::Result<T, DHTError>;
+21
View File
@@ -0,0 +1,21 @@
mod error;
mod server;
pub mod protocol;
pub mod types;
pub mod metadata; // 公开 metadata 模块
mod sharded; // 分片锁模块
pub mod scheduler; // 元数据调度器
pub use error::{DHTError, Result};
pub use server::{DHTServer, HashDiscovered};
pub use types::{DHTOptions, FileInfo, TorrentInfo};
pub use sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple};
pub use scheduler::MetadataScheduler;
// 重新导出常用类型
pub mod prelude {
pub use crate::error::{DHTError, Result};
pub use crate::server::DHTServer;
pub use crate::types::{DHTOptions, FileInfo, TorrentInfo};
pub use crate::scheduler::MetadataScheduler;
}
+105
View File
@@ -0,0 +1,105 @@
#[global_allocator]
static GLOBAL: MiMalloc = MiMalloc;
use dht_crawler::prelude::*;
use std::sync::Arc;
use mimalloc::MiMalloc;
use tracing_subscriber::EnvFilter;
use std::sync::atomic::{AtomicUsize, Ordering};
#[tokio::main]
async fn main() -> Result<()> {
if std::env::var("RUST_LOG").is_err() {
std::env::set_var("RUST_LOG", "info");
}
// 直接输出到 stdout,避免 _guard 被 drop 导致日志丢失
tracing_subscriber::fmt()
.with_env_filter(EnvFilter::from_default_env())
.with_ansi(true)
.init();
let options = DHTOptions {
port: 45452,
auto_metadata: true,
metadata_timeout: 3, // ✅ 快速超时,快速失败
max_metadata_queue_size: 100000, // ✅ 大缓冲区(防止饱和)
max_metadata_worker_count: 1000, // ✅ 激进并发(最大化吞吐)
};
// 统计计数器
let torrent_count = Arc::new(AtomicUsize::new(0));
let torrent_count_clone = torrent_count.clone();
// 🚀 初始化 DHT Server
log::info!("🔧 正在初始化 DHT Server...");
let server = DHTServer::new(options.clone()).await?;
log::info!("🚀 DHT Server 启动,监听端口: {}", options.port);
// 设置 torrent 回调
server.on_torrent(move |_torrent| {
let _count = torrent_count_clone.fetch_add(1, Ordering::Relaxed) + 1;
// 🔇 取消打印 torrent 信息,减少日志输出
// let total_size: u64 = torrent.files.iter().map(|f| f.size).sum();
// let files_display = if torrent.files.len() <= 3 {
// torrent.files.iter()
// .map(|f| format!("{} ({})", f.path, format_size(f.size)))
// .collect::<Vec<_>>()
// .join(", ")
// } else {
// format!("{}个文件", torrent.files.len())
// };
//
// log::info!(
// "🎉 [{}] {} ({}, {})",
// count,
// torrent.name,
// format_size(total_size),
// files_display
// );
});
// 设置元数据获取前的检查回调
server.on_metadata_fetch(|_hash| async move {
true
});
server.set_filter(|_hash| {
true
});
server.on_duplicate(|_hash| {
});
// 启动监控任务
let dht_monitor = server.clone();
let count_monitor = torrent_count.clone();
tokio::spawn(async move {
let mut interval = tokio::time::interval(std::time::Duration::from_secs(5));
let start_time = std::time::Instant::now();
loop {
interval.tick().await;
let success_fetch = count_monitor.load(Ordering::Relaxed);
let uptime = start_time.elapsed().as_secs();
// ✅ 监控:布隆过滤器的位使用情况反映了爬虫的活跃度
log::info!(
"📊 [监控] 时长: {}s | 成功抓取: ✨ {} | 活跃指纹: {}",
uptime, success_fetch, dht_monitor.get_seen_count()
);
if uptime > 0 && success_fetch > 0 {
let speed = (success_fetch as f64) / (uptime as f64 / 60.0);
log::info!("📈 平均抓取速度: {:.2} 种子/分钟", speed);
}
}
});
server.start().await?;
Ok(())
}
+171
View File
@@ -0,0 +1,171 @@
use std::collections::BTreeMap;
use std::net::SocketAddr;
use std::time::Duration;
use bytes::Bytes;
use sha1::{Digest, Sha1};
use tokio::time::timeout;
use rbit::{
metadata_piece_count, ExtensionHandshake, Message, MetadataMessage,
MetadataMessageType, PeerConnection, PeerId,
};
use rbit::peer::ExtensionMessage;
use crate::types::FileInfo;
#[derive(Clone)]
pub struct RbitFetcher {
timeout: Duration,
}
impl RbitFetcher {
pub fn new(timeout_secs: u64) -> Self {
Self {
timeout: Duration::from_secs(if timeout_secs == 0 { 15 } else { timeout_secs }),
}
}
pub async fn fetch(
&self,
info_hash: &[u8; 20],
peer_addr: SocketAddr,
) -> Option<(String, u64, Vec<FileInfo>)> {
let info_hash_hex = hex::encode(info_hash);
log::debug!("[Metadata] 开始获取: {} @ {}", info_hash_hex, peer_addr);
let peer_id = PeerId::generate();
// 🔥 修改点:缩短连接超时到 3 秒
// DHT 网络很不稳定,如果 3 秒连不上,基本就是连不上了,不要浪费时间
let mut conn = match timeout(
Duration::from_secs(3),
PeerConnection::connect(peer_addr, *info_hash, *peer_id.as_bytes()),
).await {
Ok(Ok(c)) => c,
Ok(Err(_)) => return None,
Err(_) => return None,
};
if !conn.supports_extension {
return None;
}
let my_ut_metadata_id = 1;
let handshake = ExtensionHandshake::with_extensions(&[("ut_metadata", my_ut_metadata_id)]);
if let Ok(handshake_bytes) = handshake.encode() {
let _ = conn.send(Message::Extended { id: 0, payload: handshake_bytes }).await;
} else {
return None;
}
let mut metadata_size = 0;
let mut remote_ut_metadata_id = 0;
let mut pieces: BTreeMap<u32, Bytes> = BTreeMap::new();
let mut request_sent = false;
let result = timeout(self.timeout, async {
loop {
let msg = conn.receive().await.ok()?;
match msg {
Message::Extended { id, payload } => {
if id == 0 {
if let Ok(ExtensionMessage::Handshake(remote_hs)) = ExtensionMessage::decode(id, &payload) {
if let Some(size) = remote_hs.metadata_size {
metadata_size = size as u32;
}
if let Some(ext_id) = remote_hs.get_extension_id("ut_metadata") {
remote_ut_metadata_id = ext_id;
}
}
if metadata_size > 0 && remote_ut_metadata_id > 0 && !request_sent {
if metadata_size > 10 * 1024 * 1024 { return None; }
let count = metadata_piece_count(metadata_size as usize);
for i in 0..count {
let req = MetadataMessage::request(i as u32);
if let Ok(encoded) = req.encode() {
let _ = conn.send(Message::Extended { id: remote_ut_metadata_id, payload: encoded }).await;
}
}
request_sent = true;
}
} else if id == my_ut_metadata_id {
if let Ok(meta_msg) = MetadataMessage::decode(&payload) {
if meta_msg.msg_type == MetadataMessageType::Data {
if let Some(data) = meta_msg.data {
pieces.insert(meta_msg.piece, data);
}
}
}
if metadata_size > 0 {
let total_received: usize = pieces.values().map(|p| p.len()).sum();
if total_received >= metadata_size as usize {
let mut full_data = Vec::with_capacity(metadata_size as usize);
let count = metadata_piece_count(metadata_size as usize);
let mut success = true;
for i in 0..count {
if let Some(p) = pieces.get(&(i as u32)) {
full_data.extend_from_slice(p);
} else {
success = false; break;
}
}
if success {
let info_hash_copy = *info_hash;
let full_data_clone = full_data.clone();
let is_valid = tokio::task::spawn_blocking(move || {
let mut hasher = Sha1::new();
hasher.update(&full_data_clone);
let digest: [u8; 20] = hasher.finalize().into();
digest == info_hash_copy
}).await.unwrap_or(false);
if is_valid {
return Some(full_data);
}
return None;
}
}
}
}
}
_ => {}
}
}
}).await;
match result {
Ok(Some(info_bytes)) => {
if let Ok(value) = rbit::decode(&info_bytes) {
if let Some(dict) = value.as_dict() {
let name = dict.get(&b"name"[..]).and_then(|v| v.as_str()).unwrap_or("Unknown").to_string();
let mut total_size = 0;
let mut file_list = Vec::new();
if let Some(files) = dict.get(&b"files"[..]).and_then(|v| v.as_list()) {
for file in files {
if let Some(f_dict) = file.as_dict() {
if let Some(len) = f_dict.get(&b"length"[..]).and_then(|v| v.as_integer()) {
let len = len as u64;
total_size += len;
let mut path_parts = Vec::new();
if let Some(path_list) = f_dict.get(&b"path"[..]).and_then(|v| v.as_list()) {
for p in path_list {
if let Some(p_str) = p.as_str() { path_parts.push(p_str); }
}
}
file_list.push(FileInfo { path: path_parts.join("/"), size: len });
}
}
}
} else if let Some(len) = dict.get(&b"length"[..]).and_then(|v| v.as_integer()) {
total_size = len as u64;
file_list.push(FileInfo { path: name.clone(), size: total_size });
}
if total_size > 0 { return Some((name, total_size, file_list)); }
}
}
None
}
_ => None,
}
}
}
+33
View File
@@ -0,0 +1,33 @@
use serde::Deserialize;
#[derive(Deserialize, Debug)]
#[allow(dead_code)]
pub struct DhtMessage {
pub t: serde_bytes::ByteBuf,
#[allow(dead_code)] // 用于快速预检查,不在反序列化后使用
pub y: String,
#[allow(dead_code)] // 用于快速预检查,不在反序列化后使用
pub q: Option<String>,
pub a: Option<DhtArgs>,
pub r: Option<DhtResponse>,
}
#[derive(Deserialize, Debug)]
pub struct DhtArgs {
pub id: Option<serde_bytes::ByteBuf>,
pub target: Option<serde_bytes::ByteBuf>,
pub info_hash: Option<serde_bytes::ByteBuf>,
pub token: Option<serde_bytes::ByteBuf>,
pub port: Option<u16>,
pub implied_port: Option<u8>,
}
#[derive(Deserialize, Debug)]
pub struct DhtResponse {
#[serde(default)]
#[allow(dead_code)]
pub id: Option<serde_bytes::ByteBuf>,
#[serde(default)]
pub nodes: Option<serde_bytes::ByteBuf>,
}
+264
View File
@@ -0,0 +1,264 @@
use crate::server::HashDiscovered;
use crate::types::TorrentInfo;
use crate::metadata::RbitFetcher;
use std::sync::{Arc, RwLock};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use tokio::sync::{mpsc, Mutex};
use std::time::Duration;
type TorrentCallback = Arc<dyn Fn(TorrentInfo) + Send + Sync>;
type MetadataFetchCallback = Arc<dyn Fn(String) -> std::pin::Pin<Box<dyn std::future::Future<Output = bool> + Send>> + Send + Sync>;
/// 元数据调度器(优雅版:Worker 池 + Channel
/// 负责管理元数据获取队列和任务调度
pub struct MetadataScheduler {
// 输入通道
hash_rx: mpsc::Receiver<HashDiscovered>,
// 配置
max_queue_size: usize,
max_concurrent: usize,
// 元数据获取器
fetcher: Arc<RbitFetcher>,
// 回调
callback: Arc<RwLock<Option<TorrentCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
// 统计(使用 Atomic 支持多线程访问)
total_received: Arc<AtomicU64>,
total_dropped: Arc<AtomicU64>,
total_dispatched: Arc<AtomicU64>,
// 共享的队列长度计数器(用于向 Server 反馈背压)
queue_len: Arc<AtomicUsize>,
}
impl MetadataScheduler {
pub fn new(
hash_rx: mpsc::Receiver<HashDiscovered>,
fetcher: Arc<RbitFetcher>,
max_queue_size: usize,
max_concurrent: usize,
callback: Arc<RwLock<Option<TorrentCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
queue_len: Arc<AtomicUsize>, // 新增参数
) -> Self {
Self {
hash_rx,
max_queue_size,
max_concurrent,
fetcher,
callback,
on_metadata_fetch,
total_received: Arc::new(AtomicU64::new(0)),
total_dropped: Arc::new(AtomicU64::new(0)),
total_dispatched: Arc::new(AtomicU64::new(0)),
queue_len,
}
}
/// 设置 torrent 回调
pub fn set_callback(&mut self, callback: TorrentCallback) {
if let Ok(mut guard) = self.callback.try_write() {
*guard = Some(callback);
}
}
/// 设置元数据获取前的检查回调
pub fn set_metadata_fetch_callback(&mut self, callback: MetadataFetchCallback) {
if let Ok(mut guard) = self.on_metadata_fetch.try_write() {
*guard = Some(callback);
}
}
/// 运行调度器(完全事件驱动)
pub async fn run(mut self) {
// 创建任务队列(channel 自带背压)
let (task_tx, task_rx) = mpsc::channel::<HashDiscovered>(self.max_queue_size);
let task_rx = Arc::new(Mutex::new(task_rx));
// 启动 Worker 池
for worker_id in 0..self.max_concurrent {
let task_rx = task_rx.clone();
let fetcher = self.fetcher.clone();
let callback = self.callback.clone();
let on_metadata_fetch = self.on_metadata_fetch.clone();
let total_dispatched = self.total_dispatched.clone();
let queue_len = self.queue_len.clone(); // 传递计数器
tokio::spawn(async move {
log::trace!("Worker {} 启动", worker_id);
loop {
// Worker 从队列取任务(阻塞等待,零延迟)
let hash = {
let mut rx = task_rx.lock().await;
let h = rx.recv().await;
// 取出任务后,减少计数器
if h.is_some() {
queue_len.fetch_sub(1, Ordering::Relaxed);
}
h
};
let hash = match hash {
Some(h) => h,
None => break, // Channel 关闭,退出
};
total_dispatched.fetch_add(1, Ordering::Relaxed);
// 执行任务
Self::process_hash(
hash,
&fetcher,
&callback,
&on_metadata_fetch,
).await;
}
log::trace!("Worker {} 退出", worker_id);
});
}
// 主循环:只负责接收 hash 并转发到 worker 队列
let mut stats_interval = tokio::time::interval(Duration::from_secs(60));
stats_interval.tick().await;
loop {
tokio::select! {
Some(hash) = self.hash_rx.recv() => {
self.total_received.fetch_add(1, Ordering::Relaxed);
// 尝试发送到 worker 队列
match task_tx.try_send(hash) {
Ok(_) => {
// 成功入队,增加计数器
self.queue_len.fetch_add(1, Ordering::Relaxed);
}
Err(mpsc::error::TrySendError::Full(_)) => {
// 队列满,丢弃
self.total_dropped.fetch_add(1, Ordering::Relaxed);
}
Err(_) => break, // Channel 关闭
}
}
_ = stats_interval.tick() => {
self.print_stats(&task_tx);
}
else => break,
}
}
// log::info!("🛑 Metadata 调度器停止");
}
/// 处理单个 hashWorker 调用)
async fn process_hash(
hash: HashDiscovered,
fetcher: &Arc<RbitFetcher>,
callback: &Arc<RwLock<Option<TorrentCallback>>>,
on_metadata_fetch: &Arc<RwLock<Option<MetadataFetchCallback>>>,
) {
let info_hash = hash.info_hash.clone();
let peer_addr = hash.peer_addr;
// 检查是否需要获取(获取回调快照并释放锁)
let maybe_check_fn = {
match on_metadata_fetch.read() {
Ok(guard) => guard.clone(),
Err(_) => return, // 锁中毒
}
};
if let Some(f) = maybe_check_fn {
if !f(info_hash.clone()).await {
return;
}
}
// 解码 info_hash
let info_hash_bytes: [u8; 20] = match hex::decode(&info_hash) {
Ok(bytes) if bytes.len() == 20 => {
let mut arr = [0u8; 20];
arr.copy_from_slice(&bytes);
arr
}
_ => return,
};
// 获取元数据
if let Some((name, total_size, files)) = fetcher.fetch(&info_hash_bytes, peer_addr).await {
let metadata = TorrentInfo {
info_hash,
name,
total_size,
files,
magnet_link: format!("magnet:?xt=urn:btih:{}", hash.info_hash),
peers: vec![peer_addr.to_string()],
piece_length: 0,
timestamp: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs(),
};
// 获取回调快照并释放锁
let maybe_torrent_cb = {
match callback.read() {
Ok(guard) => guard.clone(),
Err(_) => return, // 锁中毒
}
};
if let Some(cb) = maybe_torrent_cb {
cb(metadata);
}
}
}
/// 输出统计信息
fn print_stats(&self, task_tx: &mpsc::Sender<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}%)",
queue_size,
self.max_queue_size,
queue_pressure,
received,
dispatched,
dropped,
drop_rate
);
} else {
log::info!(
"📊 Metadata 调度器统计:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)",
queue_size,
self.max_queue_size,
queue_pressure,
received,
dispatched,
dropped,
drop_rate
);
}
}
}
+626
View File
@@ -0,0 +1,626 @@
use crate::error::Result;
use crate::metadata::RbitFetcher;
use crate::protocol::{DhtMessage, DhtArgs, DhtResponse};
use crate::scheduler::MetadataScheduler;
use crate::types::{DHTOptions, TorrentInfo};
use crate::sharded::{ShardedBloom, ShardedNodeQueue, NodeTuple};
use rand::Rng;
use ahash::AHasher;
use std::hash::{Hash, Hasher};
use std::net::{IpAddr, SocketAddr};
use std::sync::{Arc, RwLock};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::net::UdpSocket;
use tokio::sync::{mpsc, Semaphore};
use socket2::{Socket, Domain, Type, Protocol};
use std::pin::Pin;
use std::future::Future;
const BOOTSTRAP_NODES: &[&str] = &[
"router.bittorrent.com:6881",
"dht.transmissionbt.com:6881",
"router.utorrent.com:6881",
"dht.aelitis.com:6881",
];
// 类型定义
pub type BoxedBoolFuture = Pin<Box<dyn Future<Output = bool> + Send>>;
pub type MetadataFetchCallback = Arc<dyn Fn(String) -> BoxedBoolFuture + Send + Sync>;
// Hash 发现事件
/// DHT Server 发现 hash 后发送此事件,由独立的 MetadataScheduler 处理
#[derive(Debug, Clone)]
pub struct HashDiscovered {
pub info_hash: String,
pub peer_addr: SocketAddr,
pub discovered_at: std::time::Instant,
}
// ---------------------------------------------------------------
type TorrentCallback = Arc<dyn Fn(TorrentInfo) + Send + Sync>;
type FilterCallback = Arc<dyn Fn(&str) -> bool + Send + Sync>;
type DuplicateCallback = Arc<dyn Fn(&str) + Send + Sync>;
#[derive(Clone)]
pub struct DHTServer {
#[allow(dead_code)]
options: DHTOptions,
node_id: Vec<u8>,
socket: Arc<UdpSocket>,
token_secret: Vec<u8>,
// 这些回调现在与 MetadataScheduler 共享
callback: Arc<RwLock<Option<TorrentCallback>>>,
filter: Arc<RwLock<Option<FilterCallback>>>,
on_duplicate: Arc<RwLock<Option<DuplicateCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
// 使用分片锁,大幅减少竞争
node_queue: Arc<ShardedNodeQueue>,
seen_hashes: Arc<ShardedBloom>,
// 发送 hash 发现事件
hash_tx: mpsc::Sender<HashDiscovered>,
// Metadata 队列长度(用于自适应爬取速度)
metadata_queue_len: Arc<AtomicUsize>,
max_metadata_queue_size: usize,
}
impl DHTServer {
pub async fn new(options: DHTOptions) -> Result<Self> {
let socket = Socket::new(Domain::IPV4, Type::DGRAM, Some(Protocol::UDP))?;
#[cfg(not(windows))]
{ let _ = socket.set_reuse_port(true); }
let _ = socket.set_reuse_address(true);
socket.set_nonblocking(true)?;
// 增加网络缓冲区以应对高QPS
let _ = socket.set_recv_buffer_size(32 * 1024 * 1024); // 32MB(原16MB
let _ = socket.set_send_buffer_size(8 * 1024 * 1024); // 8MB(原4MB
let addr: SocketAddr = format!("0.0.0.0:{}", options.port).parse().unwrap();
socket.bind(&addr.into())?;
let socket = UdpSocket::from_std(socket.into())?;
let node_id = generate_random_id();
let mut rng = rand::thread_rng();
let token_secret: Vec<u8> = (0..10).map(|_| rng.gen()).collect();
// 使用分片队列和分片布隆过滤器
// 队列容量:100000 个节点(扩容以适应 DHT 网络裂变速度)
let node_queue = ShardedNodeQueue::new(100000);
// 布隆过滤器:预期500万元素,0.1%误判率
// 内存使用:约 90MB32分片 × 2.8MB
let bloom = ShardedBloom::new_for_fp_rate(5_000_000, 0.001);
// -----------------------------------------------------------
// 内部初始化 MetadataScheduler
// -----------------------------------------------------------
let (hash_tx, hash_rx) = mpsc::channel::<HashDiscovered>(10000);
let fetcher = Arc::new(RbitFetcher::new(options.metadata_timeout));
// 创建共享的回调状态
let callback = Arc::new(RwLock::new(None));
let on_metadata_fetch = Arc::new(RwLock::new(None));
// 创建共享的队列长度计数器
let metadata_queue_len = Arc::new(AtomicUsize::new(0));
let scheduler = MetadataScheduler::new(
hash_rx,
fetcher,
options.max_metadata_queue_size,
options.max_metadata_worker_count,
callback.clone(),
on_metadata_fetch.clone(),
metadata_queue_len.clone(),
);
// 启动 Scheduler
tokio::spawn(async move {
scheduler.run().await;
});
let server = Self {
options: options.clone(),
node_id: node_id.clone(),
socket: Arc::new(socket),
token_secret,
callback,
on_metadata_fetch,
node_queue: Arc::new(node_queue),
seen_hashes: Arc::new(bloom),
filter: Arc::new(RwLock::new(None)),
on_duplicate: Arc::new(RwLock::new(None)),
hash_tx,
metadata_queue_len,
max_metadata_queue_size: options.max_metadata_queue_size,
};
Ok(server)
}
pub fn local_addr(&self) -> Result<SocketAddr> {
Ok(self.socket.local_addr()?)
}
/// 设置元数据获取前的检查回调
///
/// 此回调在发现新的 info_hash 后,但在实际连接对等端获取元数据之前执行。
/// 你可以在这里进行去重检查(如查询数据库),返回 `true` 表示继续获取,`false` 表示跳过。
///
/// # 注意事项
/// - 回调是在 `MetadataScheduler` 的 Worker 线程中异步执行的(通过 `.await`)。
/// - 支持耗时操作(如数据库查询),但请注意 Worker 数量限制(默认 500)。
/// - 如果回调执行过慢,可能会导致任务队列堆积。
///
/// # 示例
/// ```rust
/// server.on_metadata_fetch(|hash| async move {
/// // 检查数据库是否存在
/// // let exists = db.has(hash).await;
/// // !exists
/// true
/// });
/// ```
pub fn on_metadata_fetch<F, Fut>(&self, callback: F)
where
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))
}));
}
/// 设置成功获取到种子信息的回调
///
/// 当成功从对等端下载并解析出种子元数据(Metadata)后调用。
///
/// # 注意事项
/// - 此回调是在 Worker 线程中同步执行的。
/// - 如果包含耗时操作(如写入大量数据或复杂计算),**必须**在回调内部手动使用 `tokio::spawn`。
/// - 否则会阻塞当前的元数据获取 Worker,降低系统吞吐量。
///
/// # 示例
/// ```rust
/// server.on_torrent(|info| {
/// // 简单操作可以直接做
/// println!("Got torrent: {}", info.name);
///
/// // 耗时操作建议 spawn
/// tokio::spawn(async move {
/// save_to_db(info).await;
/// });
/// });
/// ```
pub fn on_torrent<F>(&self, callback: F) where F: Fn(TorrentInfo) + Send + Sync + 'static {
*self.callback.write().unwrap() = Some(Arc::new(callback));
}
/// 设置 Hash 过滤器
///
/// 在处理 `announce_peer` 消息时,用于快速判断是否应该处理该 Hash。
/// 这通常用于布隆过滤器之前的黑名单或白名单机制。
///
/// # 注意事项
/// - 此回调是在 UDP 处理线程中**同步执行**的。
/// - **绝对禁止**执行任何耗时操作(如 IO、数据库查询、锁等待)。
/// - 任何延迟都会直接阻塞网络包的接收,导致丢包。
/// - 应仅进行纯内存的快速判断。
pub fn set_filter<F>(&self, filter: F) where F: Fn(&str) -> bool + Send + Sync + 'static {
*self.filter.write().unwrap() = Some(Arc::new(filter));
}
/// 设置重复 Hash 发现的回调
///
/// 当接收到的 Hash 已经被布隆过滤器标记为“已存在”时调用。
///
/// # 注意事项
/// - 库内部已经自动为每次调用包裹了 `tokio::spawn`。
/// - 因此你可以放心地在回调中执行耗时操作(如数据库记录),而不用担心阻塞 UDP 线程。
/// - 虽然内部有 spawn,但频繁触发仍会产生大量任务,请注意资源控制。
pub fn on_duplicate<F>(&self, callback: F) where F: Fn(&str) + Send + Sync + 'static {
*self.on_duplicate.write().unwrap() = Some(Arc::new(callback));
}
pub fn get_seen_count(&self) -> usize {
// 分片布隆过滤器的位数统计
self.seen_hashes.number_of_bits() as usize
}
pub fn get_node_pool_size(&self) -> usize {
self.node_queue.len()
}
pub async fn start(&self) -> Result<()> {
self.start_receiver();
self.bootstrap().await;
let server = self.clone();
tokio::spawn(async move {
let semaphore = Arc::new(Semaphore::new(2000));
let mut loop_tick = 0;
loop {
// 自适应爬取速度:根据 Metadata 队列负载调整爬取策略
let queue_len = server.metadata_queue_len.load(Ordering::Relaxed);
let queue_pressure = queue_len as f64 / server.max_metadata_queue_size as f64;
// 动态计算批次大小和休眠时间
let (batch_size, sleep_duration) = if queue_pressure < 0.5 {
// 🟢 绿区:队列空闲,全速爬取
(200, Duration::from_millis(10))
} else if queue_pressure < 0.8 {
// 🟡 黄区:队列有压力,适度减速
(200, Duration::from_millis(20))
} else if queue_pressure < 0.95 {
// 🟠 橙区:队列高压,大幅减速
(20, Duration::from_millis(500))
} else {
// 🔴 红区:队列爆满,暂停主动爬取
(0, Duration::from_millis(1000))
};
let nodes_batch = {
if server.node_queue.is_empty() || batch_size == 0 {
None
} else {
Some(server.node_queue.pop_batch(batch_size))
}
};
loop_tick += 1;
if nodes_batch.is_none() || loop_tick % 50 == 0 {
server.bootstrap().await;
if nodes_batch.is_none() {
tokio::time::sleep(sleep_duration).await;
continue;
}
}
if let Some(nodes) = nodes_batch {
for node in nodes {
let permit = semaphore.clone().acquire_owned().await.unwrap();
let server_clone = server.clone();
tokio::spawn(async move {
let neighbor_id = generate_neighbor_target(&node.id, &server_clone.node_id);
let random_target = generate_random_id();
let _ = server_clone.send_find_node(node.addr, &random_target, &neighbor_id).await;
drop(permit);
});
}
}
tokio::time::sleep(sleep_duration).await;
}
});
std::future::pending::<()>().await;
Ok(())
}
fn start_receiver(&self) {
let socket = self.socket.clone();
let server = self.clone();
let num_workers = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(8);
let queue_size = 5000;
let mut senders = Vec::with_capacity(num_workers);
for _ in 0..num_workers {
let (tx, mut rx) = mpsc::channel::<(Vec<u8>, SocketAddr)>(queue_size);
senders.push(tx);
let server_clone = server.clone();
tokio::spawn(async move {
while let Some((data, addr)) = rx.recv().await {
let _ = server_clone.handle_message(&data, addr).await;
}
});
}
tokio::spawn(async move {
let mut buf = [0u8; 65536];
let mut next_worker_idx = 0;
loop {
match socket.recv_from(&mut buf).await {
Ok((size, addr)) => {
// 🛡️ 安全检查1:拒绝异常大的包(DHT 消息通常 < 2KB
if size > 8192 {
#[cfg(debug_assertions)]
log::trace!("⚠️ 拒绝异常大的 UDP 包: {} 字节 from {}", size, addr);
continue;
}
// 🛡️ 安全检查2:快速检查是否是有效的 Bencode 字典
// DHT KRPC 消息(BEP-5)必须是字典,首字符必须是 'd'
if size == 0 || buf[0] != b'd' {
continue;
}
let data = buf[..size].to_vec();
let tx = &senders[next_worker_idx];
next_worker_idx = (next_worker_idx + 1) % num_workers;
match tx.try_send((data, addr)) {
Ok(_) => {},
Err(mpsc::error::TrySendError::Full(_)) => {
#[cfg(debug_assertions)]
log::trace!("UDP worker queue full, dropping packet");
},
Err(_) => { break; }
}
}
Err(_e) => {
tokio::time::sleep(Duration::from_millis(1)).await;
}
}
}
});
}
async fn handle_message(&self, data: &[u8], addr: SocketAddr) -> Result<()> {
let msg: DhtMessage = match serde_bencode::from_bytes(data) {
Ok(m) => m,
Err(_) => return Ok(()),
};
match msg.y.as_str() {
"q" => {
if let Some(q_type) = &msg.q {
self.handle_query(&msg, q_type.as_bytes(), addr).await?;
}
}
"r" => {
if let Some(response) = &msg.r {
self.handle_response(response).await?;
}
}
_ => {}
}
Ok(())
}
async fn handle_query(&self, msg: &DhtMessage, query_type: &[u8], addr: SocketAddr) -> Result<()> {
let args = match &msg.a {
Some(a) => a,
None => return Ok(()),
};
let transaction_id = &msg.t;
let sender_id: Option<&[u8]> = args.id.as_deref().map(|v| v.as_slice());
let target_id_fallback: Option<&[u8]> = args.target.as_deref()
.or(args.info_hash.as_deref())
.map(|v| v.as_slice());
let q_str = std::str::from_utf8(query_type).unwrap_or("");
if q_str == "announce_peer" {
self.handle_announce_peer(args, addr).await?;
}
self.send_response(transaction_id, addr, q_str, sender_id, target_id_fallback).await?;
Ok(())
}
async fn handle_announce_peer(&self, args: &DhtArgs, addr: SocketAddr) -> Result<()> {
if let Some(token) = &args.token {
if !self.validate_token(token, addr) { return Ok(()); }
} else {
return Ok(());
}
if let Some(info_hash) = &args.info_hash {
let info_hash_arr: [u8; 20] = match info_hash.as_ref().try_into() {
Ok(arr) => arr, Err(_) => return Ok(()),
};
let hash_hex = hex::encode(info_hash_arr);
// 使用分片布隆过滤器进行高效去重
let is_duplicate = self.seen_hashes.check_and_set(&info_hash_arr);
if is_duplicate {
let dup_cb = self.on_duplicate.read().unwrap().clone();
if let Some(cb) = dup_cb {
let hash_hex_clone = hash_hex.clone();
tokio::spawn(async move {
cb(&hash_hex_clone);
});
}
return Ok(());
}
let filter_cb = self.filter.read().unwrap().clone();
if let Some(f) = filter_cb {
if !f(&hash_hex) { return Ok(()); }
}
#[cfg(debug_assertions)]
log::debug!("🔥 新 Hash: {} 来自 {}", hash_hex, addr);
// 解耦:发送 hash 发现事件
let port = if let Some(implied) = args.implied_port {
if implied != 0 { addr.port() } else { args.port.unwrap_or(0) }
} else {
args.port.unwrap_or(addr.port())
};
if port > 0 {
let event = HashDiscovered {
info_hash: hash_hex,
peer_addr: SocketAddr::new(addr.ip(), port),
discovered_at: std::time::Instant::now(),
};
// 使用 try_send,队列满时直接丢弃(背压)
if let Err(_) = self.hash_tx.try_send(event) {
#[cfg(debug_assertions)]
log::trace!("⚠️ Hash 队列满,丢弃 hash");
}
}
}
Ok(())
}
async fn handle_response(&self, response: &DhtResponse) -> Result<()> {
if let Some(nodes_bytes) = &response.nodes {
self.process_compact_nodes(nodes_bytes);
}
Ok(())
}
fn process_compact_nodes(&self, nodes_bytes: &[u8]) {
if nodes_bytes.len() % 26 != 0 { return; }
// 使用分片队列,直接并发插入(无锁竞争)
for chunk in nodes_bytes.chunks(26) {
let id = chunk[0..20].to_vec();
let port = u16::from_be_bytes([chunk[24], chunk[25]]);
let ip = std::net::Ipv4Addr::new(chunk[20], chunk[21], chunk[22], chunk[23]);
let addr = SocketAddr::new(std::net::IpAddr::V4(ip), port);
self.node_queue.push(NodeTuple { id, addr });
}
}
async fn send_response(
&self,
tid: &[u8],
addr: SocketAddr,
query_type: &str,
sender_id: Option<&[u8]>,
target_id_fallback: Option<&[u8]>,
) -> Result<()> {
let mut r_dict = std::collections::HashMap::new();
let reference_id = sender_id.or(target_id_fallback);
let my_id = if let Some(target) = reference_id {
generate_neighbor_target(target, &self.node_id)
} else {
self.node_id.clone()
};
r_dict.insert(b"id".to_vec(), serde_bencode::value::Value::Bytes(my_id));
let token = self.generate_token(addr);
r_dict.insert(b"token".to_vec(), serde_bencode::value::Value::Bytes(token));
if query_type == "get_peers" || query_type == "find_node" {
// 使用分片队列获取随机节点(无锁竞争)
let nodes = self.node_queue.get_random_nodes(8);
let mut nodes_data = Vec::new();
for node in nodes {
nodes_data.extend_from_slice(&node.id);
match node.addr.ip() {
IpAddr::V4(ip) => nodes_data.extend_from_slice(&ip.octets()),
_ => continue,
}
nodes_data.extend_from_slice(&node.addr.port().to_be_bytes());
}
r_dict.insert(b"nodes".to_vec(), serde_bencode::value::Value::Bytes(nodes_data));
}
let mut response: std::collections::HashMap<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) {
let _ = self.socket.send_to(&encoded, addr).await;
}
Ok(())
}
async fn bootstrap(&self) {
let target = generate_random_id();
for node in BOOTSTRAP_NODES {
match tokio::net::lookup_host(node).await {
Ok(addrs) => {
for addr in addrs {
if addr.is_ipv6() { continue; }
let _ = self.send_find_node(addr, &target, &self.node_id).await;
}
}
Err(_) => {}
}
}
}
async fn send_find_node(&self, addr: SocketAddr, target: &[u8], sender_id: &[u8]) -> Result<()> {
let mut args = std::collections::HashMap::new();
args.insert(b"id".to_vec(), serde_bencode::value::Value::Bytes(sender_id.to_vec()));
args.insert(b"target".to_vec(), serde_bencode::value::Value::Bytes(target.to_vec()));
let mut msg: std::collections::HashMap<String, serde_bencode::value::Value> = std::collections::HashMap::new();
msg.insert("t".to_string(), serde_bencode::value::Value::Bytes(vec![0, 1]));
msg.insert("y".to_string(), serde_bencode::value::Value::Bytes(b"q".to_vec()));
msg.insert("q".to_string(), serde_bencode::value::Value::Bytes(b"find_node".to_vec()));
msg.insert("a".to_string(), serde_bencode::value::Value::Dict(args));
if let Ok(encoded) = serde_bencode::to_bytes(&msg) {
let _ = self.socket.send_to(&encoded, addr).await;
}
Ok(())
}
fn generate_token(&self, addr: SocketAddr) -> Vec<u8> {
let mut hasher = AHasher::default();
// Hash IP地址
match addr.ip() {
IpAddr::V4(ip) => ip.octets().hash(&mut hasher),
IpAddr::V6(ip) => ip.octets().hash(&mut hasher),
}
// Hash 密钥
self.token_secret.hash(&mut hasher);
// 返回 8 字节 token
let hash = hasher.finish();
hash.to_le_bytes().to_vec()
}
fn validate_token(&self, token: &[u8], addr: SocketAddr) -> bool {
if token.len() != 8 {
return false;
}
let expected = self.generate_token(addr);
token == expected.as_slice()
}
}
fn generate_random_id() -> Vec<u8> {
let mut rng = rand::thread_rng();
(0..20).map(|_| rng.gen()).collect()
}
fn generate_neighbor_target(remote_id: &[u8], local_id: &[u8]) -> Vec<u8> {
let mut id = Vec::with_capacity(20);
let prefix_len = std::cmp::min(remote_id.len(), 6);
id.extend_from_slice(&remote_id[..prefix_len]);
if local_id.len() > prefix_len {
id.extend_from_slice(&local_id[prefix_len..]);
} else {
while id.len() < 20 {
id.push(rand::random());
}
}
id
}
+289
View File
@@ -0,0 +1,289 @@
// 分片锁实现 - 大幅减少锁竞争,提升并发性能
//
// 核心思想:1个大锁 → N个小锁
// 性能提升:预期 3-4 倍
use bloomfilter::Bloom;
use std::collections::{HashSet, VecDeque};
use std::net::SocketAddr;
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
// 配置:分片数量
const BLOOM_SHARD_COUNT: usize = 32; // 32个布隆过滤器分片
const QUEUE_SHARD_COUNT: usize = 16; // 16个队列分片
// ==================== 分片布隆过滤器 ====================
/// 分片布隆过滤器 - 减少锁竞争
///
/// 将单个布隆过滤器拆分为32个分片,每个分片独立锁
/// 不同的hash会落到不同的分片上,大幅减少竞争
pub struct ShardedBloom {
shards: Vec<Mutex<Bloom<[u8; 20]>>>,
count: AtomicUsize,
}
impl ShardedBloom {
/// 创建新的分片布隆过滤器
pub fn new_for_fp_rate(expected_items: usize, fp_rate: f64) -> Self {
let items_per_shard = (expected_items + BLOOM_SHARD_COUNT - 1) / BLOOM_SHARD_COUNT;
let shards = (0..BLOOM_SHARD_COUNT)
.map(|_| Mutex::new(Bloom::new_for_fp_rate(items_per_shard, fp_rate)))
.collect();
Self {
shards,
count: AtomicUsize::new(0),
}
}
/// 检查并设置元素(原子操作)
pub fn check_and_set(&self, hash: &[u8; 20]) -> bool {
let shard_idx = self.hash_to_shard(hash);
let mut shard = self.shards[shard_idx].lock().unwrap();
let present = shard.check_and_set(hash);
// 如果之前不存在,增加计数
if !present {
self.count.fetch_add(1, Ordering::Relaxed);
}
present
}
/// 获取实际发现的唯一 InfoHash 数量
pub fn number_of_bits(&self) -> u64 {
self.count.load(Ordering::Relaxed) as u64
}
/// 根据hash计算分片索引
#[inline]
fn hash_to_shard(&self, hash: &[u8; 20]) -> usize {
// 使用hash的前两个字节计算分片
let idx = (hash[0] as usize) | ((hash[1] as usize) << 8);
idx % BLOOM_SHARD_COUNT
}
}
// ==================== 分片节点队列 ====================
/// 节点信息
#[derive(Debug, Clone)]
pub struct NodeTuple {
pub id: Vec<u8>,
pub addr: SocketAddr,
}
/// 单个队列分片
struct NodeQueueShard {
queue: VecDeque<NodeTuple>,
index: HashSet<SocketAddr>,
capacity: usize,
}
impl NodeQueueShard {
fn new(capacity: usize) -> Self {
Self {
queue: VecDeque::with_capacity(capacity),
index: HashSet::with_capacity(capacity),
capacity,
}
}
fn push(&mut self, node: NodeTuple) {
if self.index.contains(&node.addr) {
return;
}
// 如果满了,移除最早的一个(保持流动性,优胜劣汰)
if self.queue.len() >= self.capacity {
if let Some(removed) = self.queue.pop_front() {
self.index.remove(&removed.addr);
}
}
self.index.insert(node.addr);
self.queue.push_back(node);
}
fn pop_batch(&mut self, count: usize) -> Vec<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);
nodes.push(node);
}
}
nodes
}
fn len(&self) -> usize {
self.queue.len()
}
fn is_empty(&self) -> bool {
self.queue.is_empty()
}
}
/// 分片节点队列 - 支持高并发
pub struct ShardedNodeQueue {
shards: Vec<Mutex<NodeQueueShard>>,
}
impl ShardedNodeQueue {
/// 创建新的分片队列
pub fn new(total_capacity: usize) -> Self {
let capacity_per_shard = (total_capacity + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT;
let shards = (0..QUEUE_SHARD_COUNT)
.map(|_| Mutex::new(NodeQueueShard::new(capacity_per_shard)))
.collect();
Self { shards }
}
/// 添加节点
pub fn push(&self, node: NodeTuple) {
let shard_idx = self.addr_to_shard(&node.addr);
let mut shard = self.shards[shard_idx].lock().unwrap();
shard.push(node);
}
/// 批量弹出节点
pub fn pop_batch(&self, count: usize) -> Vec<NodeTuple> {
let mut result = Vec::with_capacity(count);
let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT;
// 从所有分片获取
for shard in &self.shards {
if result.len() >= count {
break;
}
let mut s = shard.lock().unwrap();
let nodes = s.pop_batch(per_shard);
result.extend(nodes);
}
result
}
/// 获取随机节点(用于DHT响应)
/// 🚀 优化:使用储层采样算法,O(n)时间,无需clone全部节点
pub fn get_random_nodes(&self, count: usize) -> Vec<NodeTuple> {
use rand::Rng;
let mut rng = rand::thread_rng();
// 🚀 策略1:小规模请求用快速路径(最常见:8个节点)
if count <= 16 {
return self.get_random_nodes_fast(count);
}
// 🚀 策略2:大规模请求用储层采样
let mut result = Vec::with_capacity(count);
let mut seen = 0usize;
// 储层采样算法
for shard in &self.shards {
let s = shard.lock().unwrap();
for node in s.queue.iter() {
seen += 1;
if result.len() < count {
// 前 count 个直接加入
result.push(node.clone());
} else {
// 后续以 count/seen 的概率替换
let j = rng.gen_range(0..seen);
if j < count {
result[j] = node.clone();
}
}
}
}
result
}
/// 快速路径:小规模随机选择(针对常见的8节点请求)
fn get_random_nodes_fast(&self, count: usize) -> Vec<NodeTuple> {
use rand::Rng;
let mut rng = rand::thread_rng();
let mut result = Vec::with_capacity(count);
// 从每个分片随机选择几个节点
let per_shard = (count + QUEUE_SHARD_COUNT - 1) / QUEUE_SHARD_COUNT;
for shard in &self.shards {
if result.len() >= count {
break;
}
let s = shard.lock().unwrap();
let shard_len = s.queue.len();
if shard_len == 0 {
continue;
}
// 从当前分片随机选择最多 per_shard 个节点
let to_take = per_shard.min(shard_len).min(count - result.len());
// 生成随机索引(不重复)
let mut indices: Vec<usize> = (0..shard_len).collect();
// 只 shuffle 前 to_take 个(部分 shuffleFisher-Yates 优化)
for i in 0..to_take {
let j = rng.gen_range(i..shard_len);
indices.swap(i, j);
}
// 取前 to_take 个索引对应的节点
for i in 0..to_take {
if let Some(node) = s.queue.get(indices[i]) {
result.push(node.clone());
}
}
}
result
}
/// 获取总长度
pub fn len(&self) -> usize {
self.shards
.iter()
.map(|shard| shard.lock().unwrap().len())
.sum()
}
/// 检查是否为空
pub fn is_empty(&self) -> bool {
self.shards
.iter()
.all(|shard| shard.lock().unwrap().is_empty())
}
/// 根据地址计算分片索引
#[inline]
fn addr_to_shard(&self, addr: &SocketAddr) -> usize {
// 使用端口和IP最后一个字节
let hash = match addr.ip() {
std::net::IpAddr::V4(ip) => {
let octets = ip.octets();
(octets[3] as usize) ^ (addr.port() as usize)
}
std::net::IpAddr::V6(ip) => {
let octets = ip.octets();
(octets[15] as usize) ^ (addr.port() as usize)
}
};
hash % QUEUE_SHARD_COUNT
}
}
+80
View File
@@ -0,0 +1,80 @@
use serde::{Deserialize, Serialize};
/// 完整的种子信息(包含元数据)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TorrentInfo {
pub info_hash: String,
pub magnet_link: String,
pub name: String,
pub total_size: u64,
pub files: Vec<FileInfo>,
pub piece_length: u64,
pub peers: Vec<String>,
pub timestamp: u64,
}
/// 文件信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FileInfo {
pub path: String,
pub size: u64,
}
impl TorrentInfo {
pub fn format_size(&self) -> String {
format_bytes(self.total_size)
}
}
impl FileInfo {
pub fn format_size(&self) -> String {
format_bytes(self.size)
}
}
fn format_bytes(bytes: u64) -> String {
const UNITS: &[&str] = &["B", "KB", "MB", "GB", "TB"];
let mut size = bytes as f64;
let mut unit_index = 0;
while size >= 1024.0 && unit_index < UNITS.len() - 1 {
size /= 1024.0;
unit_index += 1;
}
format!("{:.2} {}", size, UNITS[unit_index])
}
/// DHT 服务器配置
#[derive(Debug, Clone)]
pub struct DHTOptions {
/// DHT 端口
pub port: u16,
/// 是否自动获取元数据
pub auto_metadata: bool,
/// 元数据获取超时(秒)
pub metadata_timeout: u64,
/// 元数据获取队列大小(背压限制)
pub max_metadata_queue_size: usize,
/// 并发元数据获取工作线程数
pub max_metadata_worker_count: usize,
}
impl Default for DHTOptions {
fn default() -> Self {
Self {
port: 0,
auto_metadata: true,
// 缩短超时,快速失败,不等待慢节点
metadata_timeout: 10,
// 加大队列,防止流量高峰丢包
max_metadata_queue_size: 10000,
// 提高并发,模拟 Node.js 的高并发 IO
max_metadata_worker_count: 1000,
}
}
}