Skip to main content

dht_crawler/
metadata.rs

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)]
160/// BEP-9 Metadata fetcher with an end-to-end timeout and shared Peer failure cache.
161pub struct RbitFetcher {
162    total_timeout: Duration,
163    runtime_stats: DhtRuntimeStats,
164    peer_failure_cache: Arc<PeerFailureCache>,
165}
166
167impl RbitFetcher {
168    /// Creates a standalone fetcher with the default failure-cache capacity and TTL.
169    ///
170    /// [`DHTServer`](crate::DHTServer) normally constructs this component from
171    /// [`MetadataOptions`](crate::MetadataOptions).
172    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    /// Fetch metadata from one peer under a single end-to-end deadline.
193    ///
194    /// The deadline covers TCP connect, both BitTorrent handshakes, all metadata
195    /// piece I/O, hash validation and bencode parsing. Inner library timeouts can
196    /// therefore never stack on top of the configured metadata timeout.
197    #[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}