性能优化
This commit is contained in:
+1
-1
@@ -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"
|
||||
|
||||
@@ -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, // ✅ 激进并发(最大化吞吐)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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(())
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user