性能优化

This commit is contained in:
桥下红药
2026-03-12 21:43:18 +08:00
parent 76c1998ed9
commit 2ab9103e5b
9 changed files with 145 additions and 223 deletions
+1 -1
View File
@@ -28,10 +28,10 @@ thiserror = "1.0"
socket2 = { version = "0.5", features = ["all"] }
rbit = "0.2"
bytes = "1.0"
bloomfilter = "1.0"
ahash = "0.8"
serde_bytes = "0.11.19"
metrics = { version = "0.24", optional = true }
async-channel = "2.5.0"
[dev-dependencies]
tracing = "0.1"
-1
View File
@@ -32,7 +32,6 @@ 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, // ✅ 激进并发(最大化吞吐)
-15
View File
@@ -5,21 +5,6 @@ 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),
}
+1 -1
View File
@@ -9,7 +9,7 @@ pub mod types;
pub use error::{DHTError, Result};
pub use scheduler::MetadataScheduler;
pub use server::{DHTServer, HashDiscovered};
pub use sharded::{NodeTuple, ShardedBloom, ShardedNodeQueue};
pub use sharded::{NodeTuple, ShardedNodeQueue};
pub use types::{DHTOptions, FileInfo, NetMode, TorrentInfo};
pub mod prelude {
+16 -9
View File
@@ -29,7 +29,7 @@ impl RbitFetcher {
&self,
info_hash: &[u8; 20],
peer_addr: SocketAddr,
) -> Option<(String, u64, Vec<FileInfo>)> {
) -> Option<(String, u64, Vec<FileInfo>, u64)> {
#[cfg(feature = "metrics")]
counter!("dht_metadata_fetch_attempts_total").increment(1);
@@ -138,18 +138,21 @@ impl RbitFetcher {
}
if success {
let info_hash_copy = *info_hash;
let full_data_clone = full_data.clone();
let is_valid = tokio::task::spawn_blocking(move || {
let validated = tokio::task::spawn_blocking(move || {
let mut hasher = Sha1::new();
hasher.update(&full_data_clone);
hasher.update(&full_data);
let digest: [u8; 20] = hasher.finalize().into();
digest == info_hash_copy
}).await.unwrap_or(false);
if digest == info_hash_copy {
Some(full_data)
} else {
None
}
}).await.unwrap_or(None);
if is_valid {
if validated.is_some() {
#[cfg(feature = "metrics")]
counter!("dht_metadata_handshake_result_total", "result" => "success").increment(1);
return Some(full_data);
return validated;
}
#[cfg(feature = "metrics")]
counter!("dht_metadata_fetch_fail_total", "reason" => "sha1_mismatch").increment(1);
@@ -172,6 +175,10 @@ impl RbitFetcher {
.and_then(|v| v.as_str())
.unwrap_or("Unknown")
.to_string();
let piece_length = dict
.get(&b"piece length"[..])
.and_then(|v| v.as_integer())
.unwrap_or(0) as u64;
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()) {
@@ -212,7 +219,7 @@ impl RbitFetcher {
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, piece_length));
}
}
#[cfg(feature = "metrics")]
+81 -109
View File
@@ -3,9 +3,7 @@ use crate::server::HashDiscovered;
use crate::types::TorrentInfo;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, RwLock};
#[cfg(debug_assertions)]
use std::time::Duration;
use tokio::sync::{Mutex, mpsc};
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
type TorrentCallback = Arc<dyn Fn(TorrentInfo) + Send + Sync>;
@@ -69,8 +67,7 @@ impl MetadataScheduler {
}
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 (task_tx, task_rx) = async_channel::bounded::<HashDiscovered>(self.max_queue_size);
let shutdown = self.shutdown.clone();
#[cfg_attr(not(debug_assertions), allow(unused_variables))]
@@ -94,20 +91,13 @@ impl MetadataScheduler {
log::trace!("Worker {} 收到关闭信号,退出", worker_id);
break;
}
result = async {
let mut rx = task_rx.lock().await;
rx.recv().await
} => {
let hash = {
if result.is_some() {
result = task_rx.recv() => {
let hash = match result {
Ok(h) => {
queue_len.fetch_sub(1, Ordering::Relaxed);
h
}
result
};
let hash = match hash {
Some(h) => h,
None => break,
Err(_) => break,
};
total_dispatched.fetch_add(1, Ordering::Relaxed);
@@ -127,74 +117,52 @@ impl MetadataScheduler {
});
}
#[cfg(debug_assertions)]
let mut stats_interval = tokio::time::interval(Duration::from_secs(60));
#[cfg(debug_assertions)]
stats_interval.tick().await;
let mut stats_interval = if cfg!(debug_assertions) {
Some(tokio::time::interval(std::time::Duration::from_secs(60)))
} else {
None
};
if let Some(ref mut interval) = stats_interval {
interval.tick().await;
}
let shutdown = self.shutdown.clone();
loop {
#[cfg(debug_assertions)]
{
tokio::select! {
_ = shutdown.cancelled() => {
#[cfg(debug_assertions)]
log::trace!("MetadataScheduler 主循环收到关闭信号,退出");
break;
}
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);
}
Err(mpsc::error::TrySendError::Full(_)) => {
self.total_dropped.fetch_add(1, Ordering::Relaxed);
}
Err(_) => break,
}
}
_ = stats_interval.tick() => {
self.print_stats(&task_tx);
}
else => break,
tokio::select! {
_ = shutdown.cancelled() => {
#[cfg(debug_assertions)]
log::trace!("MetadataScheduler 主循环收到关闭信号,退出");
break;
}
}
result = self.hash_rx.recv() => {
match result {
Some(hash) => {
self.total_received.fetch_add(1, Ordering::Relaxed);
#[cfg(not(debug_assertions))]
{
tokio::select! {
_ = shutdown.cancelled() => {
#[cfg(debug_assertions)]
log::trace!("MetadataScheduler 主循环收到关闭信号,退出");
break;
}
result = self.hash_rx.recv() => {
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);
}
Err(mpsc::error::TrySendError::Full(_)) => {
self.total_dropped.fetch_add(1, Ordering::Relaxed);
}
Err(_) => break,
match task_tx.try_send(hash) {
Ok(_) => {
self.queue_len.fetch_add(1, Ordering::Relaxed);
}
Err(async_channel::TrySendError::Full(_)) => {
self.total_dropped.fetch_add(1, Ordering::Relaxed);
}
Err(_) => break,
}
None => break,
}
None => break,
}
}
_ = async {
match stats_interval.as_mut() {
Some(interval) => interval.tick().await,
None => std::future::pending().await,
}
} => {
self.print_stats_inline();
}
}
}
// 显式关闭 task_tx,让所有 worker 任务能够退出
drop(task_tx);
#[cfg(debug_assertions)]
log::trace!("MetadataScheduler 主循环退出,等待 worker 任务完成");
@@ -231,7 +199,9 @@ impl MetadataScheduler {
_ => return,
};
if let Some((name, total_size, files)) = fetcher.fetch(&info_hash_bytes, peer_addr).await {
if let Some((name, total_size, files, piece_length)) =
fetcher.fetch(&info_hash_bytes, peer_addr).await
{
let metadata = TorrentInfo {
info_hash,
name,
@@ -239,10 +209,10 @@ impl MetadataScheduler {
files,
magnet_link: format!("magnet:?xt=urn:btih:{}", hash.info_hash),
peers: vec![peer_addr.to_string()],
piece_length: 0,
piece_length,
timestamp: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.unwrap_or_default()
.as_secs(),
};
@@ -259,43 +229,45 @@ impl MetadataScheduler {
}
}
#[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);
fn print_stats_inline(&self) {
#[cfg(debug_assertions)]
{
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 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;
let queue_len = self.queue_len.load(Ordering::Relaxed);
let queue_pressure = (queue_len 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
);
if queue_pressure > 80.0 {
log::warn!(
"Metadata 队列高压:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)",
queue_len,
self.max_queue_size,
queue_pressure,
received,
dispatched,
dropped,
drop_rate
);
} else {
log::info!(
"Metadata 调度器统计:队列={}/{}({:.1}%), 接收={}, 调度={}, 丢弃={}({:.2}%)",
queue_len,
self.max_queue_size,
queue_pressure,
received,
dispatched,
dropped,
drop_rate
);
}
}
}
}
+46 -39
View File
@@ -45,9 +45,9 @@ type FilterCallback = Arc<dyn Fn(&str) -> bool + Send + Sync>;
pub struct DHTServer {
#[allow(dead_code)]
options: DHTOptions,
node_id: Vec<u8>,
node_id: [u8; 20],
socket_providers: Arc<HashMap<SocketAddr, Arc<UdpSocket>>>,
token_secret: Vec<u8>,
token_secret: [u8; 10],
callback: Arc<RwLock<Option<TorrentCallback>>>,
filter: Arc<RwLock<Option<FilterCallback>>>,
on_metadata_fetch: Arc<RwLock<Option<MetadataFetchCallback>>>,
@@ -111,8 +111,8 @@ impl DHTServer {
};
let node_id = generate_random_id();
let mut rng = rand::thread_rng();
let token_secret: Vec<u8> = (0..10).map(|_| rng.r#gen::<u8>()).collect();
let mut token_secret = [0u8; 10];
rand::thread_rng().fill(&mut token_secret);
let node_queue = ShardedNodeQueue::new(options.node_queue_capacity);
@@ -146,7 +146,7 @@ impl DHTServer {
let max_metadata_queue_size = options.max_metadata_queue_size;
let server = Self {
options,
node_id: node_id.clone(),
node_id,
socket_providers: Arc::new(socket_providers),
token_secret,
callback,
@@ -264,13 +264,15 @@ impl DHTServer {
}
if let Some(nodes) = nodes_batch {
let node_id = server.node_id.clone();
let node_id = server.node_id;
for node in nodes {
let permit = semaphore.clone().acquire_owned().await.unwrap();
let node_id_clone = node_id.clone();
// Pick a random avaliable socket
let socket = match server.socket_providers.values().next().cloned() {
let permit = match semaphore.clone().acquire_owned().await {
Ok(p) => p,
Err(_) => break,
};
let node_id_clone = node_id;
let socket = match server.socket_for_addr(&node.addr) {
Some(sock) => sock,
None => {
log::warn!("未绑定任何地址");
@@ -599,12 +601,15 @@ impl DHTServer {
let my_id = if let Some(target) = reference_id {
generate_neighbor_target(target, &self.node_id)
} else {
self.node_id.clone()
self.node_id.to_vec()
};
r_dict.insert(b"id".to_vec(), serde_bencode::value::Value::Bytes(my_id));
let token = self.generate_token(remote_addr);
r_dict.insert(b"token".to_vec(), serde_bencode::value::Value::Bytes(token));
r_dict.insert(
b"token".to_vec(),
serde_bencode::value::Value::Bytes(token.to_vec()),
);
if query_type == "get_peers" || query_type == "find_node" {
let requestor_is_ipv6 = remote_addr.is_ipv6();
@@ -698,13 +703,20 @@ impl DHTServer {
}
}
fn socket_for_addr(&self, addr: &SocketAddr) -> Option<Arc<UdpSocket>> {
self.socket_providers
.iter()
.find(|(bind_addr, _)| bind_addr.is_ipv4() == addr.is_ipv4())
.map(|(_, sock)| sock.clone())
}
async fn send_find_node(&self, target_addr: &SocketAddr, target: &[u8], sender_id: &[u8]) {
if let Some(sock) = self.socket_providers.values().next().cloned() {
if let Some(sock) = self.socket_for_addr(target_addr) {
send_find_node_impl(target_addr, target, sender_id, sock).await
}
}
fn generate_token(&self, addr: SocketAddr) -> Vec<u8> {
fn generate_token(&self, addr: SocketAddr) -> [u8; 8] {
let mut hasher = ahash::AHasher::default();
match addr.ip() {
@@ -714,8 +726,7 @@ impl DHTServer {
self.token_secret.hash(&mut hasher);
let hash = hasher.finish();
hash.to_le_bytes().to_vec()
hasher.finish().to_le_bytes()
}
fn validate_token(&self, token: &[u8], addr: SocketAddr) -> bool {
@@ -723,7 +734,7 @@ impl DHTServer {
return false;
}
let expected = self.generate_token(addr);
token == expected.as_slice()
token == expected
}
}
@@ -811,9 +822,10 @@ async fn send_find_node_impl(
}
}
fn generate_random_id() -> Vec<u8> {
let mut rng = rand::thread_rng();
(0..20).map(|_| rng.r#gen::<u8>()).collect()
fn generate_random_id() -> [u8; 20] {
let mut id = [0u8; 20];
rand::thread_rng().fill(&mut id);
id
}
/// 生成邻居目标节点 ID
@@ -945,7 +957,8 @@ fn process_udp_packet(
return Err(ProcessUdpPacketError::InvalidPacket);
}
let mut data = Some(buffer[..size].to_owned().into_boxed_slice());
let mut choked_count = 0;
let mut attempts = 0;
let max_attempts = workers.len();
while let Some(packet) = data.take() {
let worker = &workers[*worker_index];
@@ -956,34 +969,28 @@ fn process_udp_packet(
break;
}
Err(mpsc::error::TrySendError::Full((packet, _, _))) => {
choked_count += 1;
if *worker_index == 0 {
if choked_count >= workers.len() {
// all workers choked
#[cfg(feature = "metrics")]
counter!("dht_udp_packets_received_total", "status" => "queue_full")
.increment(1);
attempts += 1;
if attempts >= max_attempts {
#[cfg(feature = "metrics")]
counter!("dht_udp_packets_received_total", "status" => "queue_full")
.increment(1);
#[cfg(debug_assertions)]
log::trace!("UDP worker queue full, dropping packet");
return Err(ProcessUdpPacketError::ChokedWorkers);
}
choked_count = 0
#[cfg(debug_assertions)]
log::trace!("UDP worker queue full, dropping packet");
return Err(ProcessUdpPacketError::ChokedWorkers);
}
let _ = data.insert(packet); // choose the next worker
let _ = data.insert(packet);
}
Err(mpsc::error::TrySendError::Closed((packet, _, _))) => {
log::warn!("UDP worker dropped.");
workers.swap_remove(*worker_index); // remove the dead worker. faster but does not retain ordering.
let _ = data.insert(packet); // choose the next worker
workers.swap_remove(*worker_index);
let _ = data.insert(packet);
}
}
if workers.is_empty() {
// no live workers
return Err(ProcessUdpPacketError::NoLiveWorkers);
}
// dispatch messages to workers in round-robin style
*worker_index = (*worker_index + 1) % workers.len(); // note: a dead worker may be removed
*worker_index = (*worker_index + 1) % workers.len();
}
Ok(())
-45
View File
@@ -1,54 +1,9 @@
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;
const QUEUE_SHARD_COUNT: usize = 16;
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 {
#[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 {
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);
idx % BLOOM_SHARD_COUNT
}
}
#[derive(Debug, Clone)]
pub struct NodeTuple {
pub id: Vec<u8>,
-3
View File
@@ -55,8 +55,6 @@ fn format_bytes(bytes: u64) -> String {
pub struct DHTOptions {
pub port: u16,
pub auto_metadata: bool,
pub metadata_timeout: u64,
pub max_metadata_queue_size: usize,
@@ -74,7 +72,6 @@ impl Default for DHTOptions {
fn default() -> Self {
Self {
port: 6881,
auto_metadata: true,
metadata_timeout: 3,
max_metadata_queue_size: 100000,
max_metadata_worker_count: 1000,