1use crate::runtime_stats::DhtRuntimeStats;
2use crate::types::FileInfo;
3use ahash::AHashMap;
4use bytes::Bytes;
5#[cfg(feature = "metrics")]
6use metrics::{counter, gauge, histogram};
7use rbit::peer::ExtensionMessage;
8use rbit::{
9 ExtensionHandshake, Message, MetadataMessage, MetadataMessageType, PeerConnection, PeerId,
10 metadata_piece_count,
11};
12use sha1::{Digest, Sha1};
13use std::collections::{BTreeMap, VecDeque};
14use std::net::SocketAddr;
15use std::sync::{Arc, Mutex};
16use std::time::{Duration, Instant};
17use tokio::time::timeout;
18
19pub(crate) type FetchedMetadata = (String, u64, Vec<FileInfo>, u64);
20
21pub(crate) enum MetadataFetchOutcome {
22 Fetched(FetchedMetadata),
23 Failed,
24 SkippedCached,
25}
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28enum MetadataFetchFailure {
29 Connect,
30 NoExtension,
31 Send,
32 SizeLimit,
33 Sha1,
34 Parse,
35 Other,
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq)]
39enum PeerFailureReason {
40 Timeout,
41 ConnectFailed,
42}
43
44#[cfg(feature = "metrics")]
45impl PeerFailureReason {
46 fn as_str(self) -> &'static str {
47 match self {
48 Self::Timeout => "timeout",
49 Self::ConnectFailed => "connect_failed",
50 }
51 }
52}
53
54#[derive(Debug, Clone, Copy)]
55struct PeerFailureEntry {
56 expires_at: Instant,
57 reason: PeerFailureReason,
58}
59
60#[derive(Default)]
61struct PeerFailureCacheInner {
62 entries: AHashMap<SocketAddr, PeerFailureEntry>,
63 expiry: VecDeque<(Instant, SocketAddr)>,
64}
65
66struct PeerFailureCache {
67 inner: Mutex<PeerFailureCacheInner>,
68 capacity: usize,
69 ttl: Duration,
70}
71
72impl PeerFailureCache {
73 fn new(capacity: usize, ttl: Duration) -> Self {
74 Self {
75 inner: Mutex::new(PeerFailureCacheInner {
76 entries: AHashMap::with_capacity(capacity.min(16_384)),
77 expiry: VecDeque::with_capacity(capacity.min(16_384)),
78 }),
79 capacity,
80 ttl,
81 }
82 }
83
84 fn get(&self, addr: SocketAddr, now: Instant) -> (Option<PeerFailureReason>, usize) {
85 if self.capacity == 0 || self.ttl.is_zero() {
86 return (None, 0);
87 }
88 let mut inner = self
89 .inner
90 .lock()
91 .unwrap_or_else(std::sync::PoisonError::into_inner);
92 Self::expire(&mut inner, now);
93 (
94 inner.entries.get(&addr).map(|entry| entry.reason),
95 inner.entries.len(),
96 )
97 }
98
99 fn insert(&self, addr: SocketAddr, reason: PeerFailureReason, now: Instant) -> usize {
100 if self.capacity == 0 || self.ttl.is_zero() {
101 return 0;
102 }
103 let mut inner = self
104 .inner
105 .lock()
106 .unwrap_or_else(std::sync::PoisonError::into_inner);
107 Self::expire(&mut inner, now);
108 while inner.entries.len() >= self.capacity && !inner.entries.contains_key(&addr) {
109 let Some((expires_at, oldest_addr)) = inner.expiry.pop_front() else {
110 break;
111 };
112 if inner
113 .entries
114 .get(&oldest_addr)
115 .is_some_and(|entry| entry.expires_at == expires_at)
116 {
117 inner.entries.remove(&oldest_addr);
118 }
119 }
120
121 let expires_at = now + self.ttl;
122 inner
123 .entries
124 .insert(addr, PeerFailureEntry { expires_at, reason });
125 inner.expiry.push_back((expires_at, addr));
126 inner.entries.len()
127 }
128
129 fn remove(&self, addr: &SocketAddr, now: Instant) -> usize {
130 if self.capacity == 0 || self.ttl.is_zero() {
131 return 0;
132 }
133 let mut inner = self
134 .inner
135 .lock()
136 .unwrap_or_else(std::sync::PoisonError::into_inner);
137 Self::expire(&mut inner, now);
138 inner.entries.remove(addr);
139 inner.entries.len()
140 }
141
142 fn expire(inner: &mut PeerFailureCacheInner, now: Instant) {
143 while let Some((expires_at, addr)) = inner.expiry.front().copied() {
144 if expires_at > now {
145 break;
146 }
147 inner.expiry.pop_front();
148 if inner
149 .entries
150 .get(&addr)
151 .is_some_and(|entry| entry.expires_at == expires_at)
152 {
153 inner.entries.remove(&addr);
154 }
155 }
156 }
157}
158
159#[derive(Clone)]
160pub struct RbitFetcher {
162 total_timeout: Duration,
163 runtime_stats: DhtRuntimeStats,
164 peer_failure_cache: Arc<PeerFailureCache>,
165}
166
167impl RbitFetcher {
168 pub fn new(timeout_secs: u64) -> Self {
173 Self::new_with_runtime_stats(timeout_secs, 200_000, 60, DhtRuntimeStats::default())
174 }
175
176 pub(crate) fn new_with_runtime_stats(
177 timeout_secs: u64,
178 peer_failure_cache_capacity: usize,
179 peer_failure_ttl_secs: u64,
180 runtime_stats: DhtRuntimeStats,
181 ) -> Self {
182 Self {
183 total_timeout: Duration::from_secs(if timeout_secs == 0 { 15 } else { timeout_secs }),
184 runtime_stats,
185 peer_failure_cache: Arc::new(PeerFailureCache::new(
186 peer_failure_cache_capacity,
187 Duration::from_secs(peer_failure_ttl_secs),
188 )),
189 }
190 }
191
192 #[cfg(test)]
198 pub(crate) async fn fetch(
199 &self,
200 info_hash: &[u8; 20],
201 peer_addr: SocketAddr,
202 ) -> MetadataFetchOutcome {
203 self.fetch_with_attempt_observer(info_hash, peer_addr, || {})
204 .await
205 }
206
207 pub(crate) async fn fetch_with_attempt_observer<F>(
208 &self,
209 info_hash: &[u8; 20],
210 peer_addr: SocketAddr,
211 on_attempt: F,
212 ) -> MetadataFetchOutcome
213 where
214 F: FnOnce() + Send,
215 {
216 let (cached_reason, cache_entries) = self.peer_failure_cache.get(peer_addr, Instant::now());
217 self.set_peer_failure_cache_entries(cache_entries);
218 if let Some(reason) = cached_reason {
219 self.runtime_stats.metadata_peer_failure_cache_hit();
220 match reason {
221 PeerFailureReason::Timeout => self.runtime_stats.peer_cache_hit_timeout(),
222 PeerFailureReason::ConnectFailed => self.runtime_stats.peer_cache_hit_connect(),
223 }
224 #[cfg(feature = "metrics")]
225 counter!("dht_metadata_peer_failure_cache_hits_total", "reason" => reason.as_str())
226 .increment(1);
227 #[cfg(not(feature = "metrics"))]
228 let _ = reason;
229 return MetadataFetchOutcome::SkippedCached;
230 }
231
232 on_attempt();
233 self.runtime_stats.metadata_peer_attempt();
234 #[cfg(feature = "metrics")]
235 {
236 counter!("dht_metadata_fetch_attempts_total").increment(1);
237 counter!("dht_metadata_peer_attempts_total").increment(1);
238 }
239
240 let started = Instant::now();
241 let result = timeout(
242 self.total_timeout,
243 self.fetch_with_peer(info_hash, peer_addr),
244 )
245 .await;
246
247 #[cfg(feature = "metrics")]
248 histogram!("dht_metadata_fetch_duration_seconds").record(started.elapsed().as_secs_f64());
249 self.runtime_stats.observe_metadata_fetch_duration(
250 started.elapsed().as_millis().min(u128::from(u64::MAX)) as u64,
251 );
252
253 match result {
254 Ok(Ok(metadata)) => {
255 let cache_entries = self.peer_failure_cache.remove(&peer_addr, Instant::now());
256 self.set_peer_failure_cache_entries(cache_entries);
257 self.runtime_stats.metadata_peer_succeeded();
258 #[cfg(feature = "metrics")]
259 {
260 counter!("dht_metadata_fetch_success_total").increment(1);
261 counter!("dht_metadata_fetch_result_total", "result" => "success").increment(1);
262 }
263 MetadataFetchOutcome::Fetched(metadata)
264 }
265 Ok(Err(reason)) => {
266 self.runtime_stats.metadata_peer_failed();
267 match reason {
268 MetadataFetchFailure::Connect => self.runtime_stats.metadata_failure_connect(),
269 MetadataFetchFailure::NoExtension => {
270 self.runtime_stats.metadata_failure_no_extension()
271 }
272 MetadataFetchFailure::Send => self.runtime_stats.metadata_failure_send(),
273 MetadataFetchFailure::SizeLimit => {
274 self.runtime_stats.metadata_failure_size_limit()
275 }
276 MetadataFetchFailure::Sha1 => self.runtime_stats.metadata_failure_sha1(),
277 MetadataFetchFailure::Parse => self.runtime_stats.metadata_failure_parse(),
278 MetadataFetchFailure::Other => self.runtime_stats.metadata_failure_other(),
279 }
280 #[cfg(feature = "metrics")]
281 counter!("dht_metadata_fetch_result_total", "result" => "failed").increment(1);
282 MetadataFetchOutcome::Failed
283 }
284 Err(_) => {
285 self.record_peer_failure(peer_addr, PeerFailureReason::Timeout);
286 self.runtime_stats.metadata_peer_failed();
287 self.runtime_stats.metadata_peer_timeout();
288 self.runtime_stats.metadata_failure_timeout();
289 #[cfg(feature = "metrics")]
290 {
291 counter!("dht_metadata_fetch_fail_total", "reason" => "timeout").increment(1);
292 counter!("dht_metadata_fetch_result_total", "result" => "timeout").increment(1);
293 }
294 MetadataFetchOutcome::Failed
295 }
296 }
297 }
298
299 fn record_peer_failure(&self, peer_addr: SocketAddr, reason: PeerFailureReason) {
300 let cache_entries = self
301 .peer_failure_cache
302 .insert(peer_addr, reason, Instant::now());
303 self.set_peer_failure_cache_entries(cache_entries);
304 #[cfg(feature = "metrics")]
305 counter!("dht_metadata_peer_failure_cache_inserts_total", "reason" => reason.as_str())
306 .increment(1);
307 }
308
309 fn set_peer_failure_cache_entries(&self, count: usize) {
310 self.runtime_stats
311 .set_metadata_peer_failure_cache_entries(count);
312 #[cfg(feature = "metrics")]
313 gauge!("dht_metadata_peer_failure_cache_entries").set(count as f64);
314 }
315
316 async fn fetch_with_peer(
317 &self,
318 info_hash: &[u8; 20],
319 peer_addr: SocketAddr,
320 ) -> Result<FetchedMetadata, MetadataFetchFailure> {
321 let peer_id = PeerId::generate();
322 let mut conn = match PeerConnection::connect(peer_addr, *info_hash, *peer_id.as_bytes())
323 .await
324 {
325 Ok(conn) => {
326 #[cfg(feature = "metrics")]
327 counter!("dht_metadata_connection_result_total", "result" => "success")
328 .increment(1);
329 conn
330 }
331 Err(_) => {
332 self.record_peer_failure(peer_addr, PeerFailureReason::ConnectFailed);
333 self.runtime_stats.metadata_connect_failed();
334 #[cfg(feature = "metrics")]
335 counter!("dht_metadata_connection_result_total", "result" => "failed").increment(1);
336 return Err(MetadataFetchFailure::Connect);
337 }
338 };
339
340 if !conn.supports_extension {
341 self.runtime_stats.metadata_no_extension();
342 #[cfg(feature = "metrics")]
343 counter!("dht_metadata_handshake_result_total", "result" => "no_extension_support")
344 .increment(1);
345 return Err(MetadataFetchFailure::NoExtension);
346 }
347
348 let my_ut_metadata_id = 1;
349 let handshake = ExtensionHandshake::with_extensions(&[("ut_metadata", my_ut_metadata_id)]);
350 let handshake_bytes = handshake.encode().map_err(|_| MetadataFetchFailure::Send)?;
351 if conn
352 .send(Message::Extended {
353 id: 0,
354 payload: handshake_bytes,
355 })
356 .await
357 .is_err()
358 {
359 #[cfg(feature = "metrics")]
360 counter!("dht_metadata_fetch_fail_total", "reason" => "send_error").increment(1);
361 return Err(MetadataFetchFailure::Send);
362 }
363
364 let mut metadata_size = 0;
365 let mut remote_ut_metadata_id = 0;
366 let mut pieces: BTreeMap<u32, Bytes> = BTreeMap::new();
367 let mut total_received = 0usize;
368 let mut request_sent = false;
369
370 let info_bytes = loop {
371 let msg = conn
372 .receive()
373 .await
374 .map_err(|_| MetadataFetchFailure::Other)?;
375 let Message::Extended { id, payload } = msg else {
376 continue;
377 };
378
379 if id == 0 {
380 if let Ok(ExtensionMessage::Handshake(remote_hs)) =
381 ExtensionMessage::decode(id, &payload)
382 {
383 if let Some(size) = remote_hs.metadata_size {
384 metadata_size = size as u32;
385 }
386 if let Some(ext_id) = remote_hs.get_extension_id("ut_metadata") {
387 remote_ut_metadata_id = ext_id;
388 }
389 }
390
391 if metadata_size > 0 && remote_ut_metadata_id > 0 && !request_sent {
392 if metadata_size > 10 * 1024 * 1024 {
393 #[cfg(feature = "metrics")]
394 counter!("dht_metadata_fetch_fail_total", "reason" => "size_limit")
395 .increment(1);
396 return Err(MetadataFetchFailure::SizeLimit);
397 }
398
399 let count = metadata_piece_count(metadata_size as usize);
400 for piece in 0..count {
401 let encoded = MetadataMessage::request(piece as u32)
402 .encode()
403 .map_err(|_| MetadataFetchFailure::Send)?;
404 if conn
405 .send(Message::Extended {
406 id: remote_ut_metadata_id,
407 payload: encoded,
408 })
409 .await
410 .is_err()
411 {
412 #[cfg(feature = "metrics")]
413 counter!("dht_metadata_fetch_fail_total", "reason" => "send_error")
414 .increment(1);
415 return Err(MetadataFetchFailure::Send);
416 }
417 }
418 request_sent = true;
419 }
420 continue;
421 }
422
423 if id != my_ut_metadata_id {
424 continue;
425 }
426 let Ok(meta_msg) = MetadataMessage::decode(&payload) else {
427 continue;
428 };
429 if meta_msg.msg_type != MetadataMessageType::Data {
430 continue;
431 }
432 let Some(data) = meta_msg.data else {
433 continue;
434 };
435
436 #[cfg(feature = "metrics")]
437 counter!("dht_metadata_bytes_downloaded_total").increment(data.len() as u64);
438 self.runtime_stats.metadata_bytes_downloaded(data.len());
439
440 let data_len = data.len();
441 if let Some(previous) = pieces.insert(meta_msg.piece, data) {
442 total_received = total_received.saturating_sub(previous.len());
443 }
444 total_received = total_received.saturating_add(data_len);
445
446 if metadata_size == 0 || total_received < metadata_size as usize {
447 continue;
448 }
449
450 let count = metadata_piece_count(metadata_size as usize);
451 let mut full_data = Vec::with_capacity(metadata_size as usize);
452 for piece in 0..count {
453 let data = pieces
454 .get(&(piece as u32))
455 .ok_or(MetadataFetchFailure::Other)?;
456 full_data.extend_from_slice(data);
457 }
458
459 let info_hash_copy = *info_hash;
460 let validated = tokio::task::spawn_blocking(move || {
461 let mut hasher = Sha1::new();
462 hasher.update(&full_data);
463 let digest: [u8; 20] = hasher.finalize().into();
464 (digest == info_hash_copy).then_some(full_data)
465 })
466 .await
467 .ok()
468 .flatten();
469
470 match validated {
471 Some(data) => {
472 #[cfg(feature = "metrics")]
473 counter!("dht_metadata_handshake_result_total", "result" => "success")
474 .increment(1);
475 break data;
476 }
477 None => {
478 #[cfg(feature = "metrics")]
479 counter!("dht_metadata_fetch_fail_total", "reason" => "sha1_mismatch")
480 .increment(1);
481 return Err(MetadataFetchFailure::Sha1);
482 }
483 }
484 };
485
486 self.runtime_stats.observe_metadata_size(info_bytes.len());
487 match parse_metadata(&info_bytes) {
488 Some(metadata) => {
489 #[cfg(feature = "metrics")]
490 histogram!("dht_metadata_size_bytes").record(info_bytes.len() as f64);
491 Ok(metadata)
492 }
493 None => {
494 #[cfg(feature = "metrics")]
495 counter!("dht_metadata_fetch_fail_total", "reason" => "parse_error").increment(1);
496 Err(MetadataFetchFailure::Parse)
497 }
498 }
499 }
500}
501
502fn parse_metadata(info_bytes: &[u8]) -> Option<FetchedMetadata> {
503 let value = rbit::decode(info_bytes).ok()?;
504 let dict = value.as_dict()?;
505 let name = dict
506 .get(&b"name"[..])
507 .and_then(|value| value.as_str())
508 .unwrap_or("Unknown")
509 .to_string();
510 let piece_length = dict
511 .get(&b"piece length"[..])
512 .and_then(|value| value.as_integer())
513 .unwrap_or(0) as u64;
514
515 let mut total_size = 0;
516 let mut file_list = Vec::new();
517 if let Some(files) = dict.get(&b"files"[..]).and_then(|value| value.as_list()) {
518 for file in files {
519 let Some(file_dict) = file.as_dict() else {
520 continue;
521 };
522 let Some(length) = file_dict
523 .get(&b"length"[..])
524 .and_then(|value| value.as_integer())
525 else {
526 continue;
527 };
528 let length = length as u64;
529 total_size += length;
530 let path = file_dict
531 .get(&b"path"[..])
532 .and_then(|value| value.as_list())
533 .map(|parts| {
534 parts
535 .iter()
536 .filter_map(|part| part.as_str())
537 .collect::<Vec<_>>()
538 .join("/")
539 })
540 .unwrap_or_default();
541 file_list.push(FileInfo { path, size: length });
542 }
543 } else if let Some(length) = dict
544 .get(&b"length"[..])
545 .and_then(|value| value.as_integer())
546 {
547 total_size = length as u64;
548 file_list.push(FileInfo {
549 path: name.clone(),
550 size: total_size,
551 });
552 }
553
554 (total_size > 0).then_some((name, total_size, file_list, piece_length))
555}
556
557#[cfg(test)]
558mod tests {
559 use super::*;
560 use tokio::net::TcpListener;
561
562 #[test]
563 fn peer_failure_cache_is_socket_specific_and_expires() {
564 let start = Instant::now();
565 let cache = PeerFailureCache::new(10, Duration::from_secs(60));
566 let first: SocketAddr = "127.0.0.1:1000".parse().unwrap();
567 let same_ip_other_port: SocketAddr = "127.0.0.1:1001".parse().unwrap();
568
569 assert_eq!(cache.insert(first, PeerFailureReason::Timeout, start), 1);
570 assert_eq!(cache.get(first, start).0, Some(PeerFailureReason::Timeout));
571 assert_eq!(cache.get(same_ip_other_port, start).0, None);
572 assert_eq!(cache.get(first, start + Duration::from_secs(61)), (None, 0));
573 }
574
575 #[test]
576 fn peer_failure_cache_evicts_oldest_at_capacity() {
577 let start = Instant::now();
578 let cache = PeerFailureCache::new(1, Duration::from_secs(60));
579 let first: SocketAddr = "127.0.0.1:1000".parse().unwrap();
580 let second: SocketAddr = "127.0.0.1:1001".parse().unwrap();
581
582 cache.insert(first, PeerFailureReason::Timeout, start);
583 cache.insert(second, PeerFailureReason::ConnectFailed, start);
584
585 assert_eq!(cache.get(first, start).0, None);
586 assert_eq!(
587 cache.get(second, start).0,
588 Some(PeerFailureReason::ConnectFailed)
589 );
590 }
591
592 #[tokio::test]
593 async fn total_timeout_covers_peer_handshake() {
594 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
595 let addr = listener.local_addr().unwrap();
596 let accept_task = tokio::spawn(async move {
597 let (_stream, _) = listener.accept().await.unwrap();
598 std::future::pending::<()>().await;
599 });
600
601 let stats = DhtRuntimeStats::default();
602 let fetcher = RbitFetcher::new_with_runtime_stats(1, 10, 60, stats.clone());
603 let started = Instant::now();
604 assert!(matches!(
605 fetcher.fetch(&[7; 20], addr).await,
606 MetadataFetchOutcome::Failed
607 ));
608 assert!(started.elapsed() < Duration::from_secs(2));
609
610 let cached_started = Instant::now();
611 assert!(matches!(
612 fetcher.fetch(&[8; 20], addr).await,
613 MetadataFetchOutcome::SkippedCached
614 ));
615 assert!(cached_started.elapsed() < Duration::from_millis(100));
616
617 let snapshot = stats.snapshot();
618 assert_eq!(snapshot.metadata_peer_attempts, 1);
619 assert_eq!(snapshot.metadata_peer_failed, 1);
620 assert_eq!(snapshot.metadata_peer_timeouts, 1);
621 assert_eq!(snapshot.metadata_peer_failure_cache_hits, 1);
622 assert_eq!(snapshot.metadata_peer_failure_cache_entries, 1);
623
624 accept_task.abort();
625 }
626}