Просмотр исходного кода

net: Use transports when polling GetAddrsMessage to seed/peer

This allows the seed/peer to only send back addresses with our desired transports
aggstam 2 лет назад
Родитель
Сommit
af9124c93a

+ 20 - 3
src/net/hosts.rs

@@ -267,7 +267,7 @@ impl Hosts {
         self.addrs.read().await.iter().cloned().collect()
     }
 
-    /// Get up to n random hosts from the hosts set.
+    /// Get up to n random peers from the hosts set.
     pub async fn fetch_n_random(&self, n: u32) -> Vec<Url> {
         let n = n as usize;
         if n == 0 {
@@ -275,8 +275,25 @@ impl Hosts {
         }
         let addrs = self.addrs.read().await;
         let urls = addrs.iter().choose_multiple(&mut OsRng, n.min(addrs.len()));
-        let urls = urls.iter().map(|&url| url.clone()).collect();
-        urls
+        urls.iter().map(|&url| url.clone()).collect()
+    }
+
+    /// Get up to n random peers that match the given transport schemes from the hosts set.
+    pub async fn fetch_n_random_with_schemes(&self, schemes: &[String], n: u32) -> Vec<Url> {
+        let n = n as usize;
+        if n == 0 {
+            return vec![]
+        }
+
+        // Retrieve all peers corresponding to that transport schemes
+        let hosts = self.fetch_with_schemes(schemes, None).await;
+        if hosts.is_empty() {
+            return hosts
+        }
+
+        // Grab random ones
+        let urls = hosts.iter().choose_multiple(&mut OsRng, n.min(hosts.len()));
+        urls.iter().map(|&url| url.clone()).collect()
     }
 
     /// Get up to limit peers that match the given transport schemes from the hosts set.

+ 3 - 1
src/net/message.rs

@@ -57,10 +57,12 @@ pub struct PongMessage {
 impl_p2p_message!(PongMessage, "pong");
 
 /// Requests address of outbound connecction.
-#[derive(Debug, Copy, Clone, SerialEncodable, SerialDecodable)]
+#[derive(Debug, Clone, SerialEncodable, SerialDecodable)]
 pub struct GetAddrsMessage {
     /// Maximum number of addresses to receive
     pub max: u32,
+    /// Preferred addresses transports
+    pub transports: Vec<String>,
 }
 impl_p2p_message!(GetAddrsMessage, "getaddr");
 

+ 23 - 2
src/net/protocol/protocol_address.rs

@@ -93,6 +93,11 @@ impl ProtocolAddress {
 
             // TODO: We might want to close the channel here if we're getting
             // corrupted addresses.
+            // Validate addreses length
+            if addrs_msg.addrs.len() > self.settings.outbound_connections {
+                continue
+            }
+
             self.hosts.store(&addrs_msg.addrs).await;
         }
     }
@@ -113,7 +118,20 @@ impl ProtocolAddress {
                 "Received GetAddrs({}) message from {}", get_addrs_msg.max, self.channel.address(),
             );
 
-            let addrs = self.hosts.fetch_n_random(get_addrs_msg.max).await;
+            // Validate transports length
+            // TODO: Verify this limit. It should be the max number of all our allowed transports,
+            //       plus their mixing.
+            if get_addrs_msg.transports.len() > 20 {
+                // TODO: Should this error out, effectively ending the connection?
+                let addrs_msg = AddrsMessage { addrs: vec![] };
+                self.channel.send(&addrs_msg).await?;
+                continue
+            }
+
+            let addrs = self
+                .hosts
+                .fetch_n_random_with_schemes(&get_addrs_msg.transports, get_addrs_msg.max)
+                .await;
             debug!(
                 target: "net::protocol_address::handle_receive_get_addrs()",
                 "Sending {} addresses to {}", addrs.len(), self.channel.address(),
@@ -160,7 +178,10 @@ impl ProtocolBase for ProtocolAddress {
         self.jobsman.spawn(self.clone().handle_receive_get_addrs(), ex).await;
 
         // Send get_address message.
-        let get_addrs = GetAddrsMessage { max: self.settings.outbound_connections as u32 };
+        let get_addrs = GetAddrsMessage {
+            max: self.settings.outbound_connections as u32,
+            transports: self.settings.allowed_transports.clone(),
+        };
         self.channel.send(&get_addrs).await?;
 
         debug!(target: "net::protocol_address::start()", "END => address={}", self.channel.address());

+ 4 - 1
src/net/protocol/protocol_seed.rs

@@ -93,7 +93,10 @@ impl ProtocolBase for ProtocolSeed {
         self.send_self_address().await?;
 
         // Send get address message
-        let get_addr = GetAddrsMessage { max: self.settings.outbound_connections as u32 };
+        let get_addr = GetAddrsMessage {
+            max: self.settings.outbound_connections as u32,
+            transports: self.settings.allowed_transports.clone(),
+        };
         self.channel.send(&get_addr).await?;
 
         // Receive addresses

+ 4 - 1
src/net/session/outbound_session.rs

@@ -529,7 +529,10 @@ impl PeerDiscovery {
                     state: "getaddr",
                 });
 
-                let get_addrs = GetAddrsMessage { max: p2p.settings().outbound_connections as u32 };
+                let get_addrs = GetAddrsMessage {
+                    max: p2p.settings().outbound_connections as u32,
+                    transports: p2p.settings().allowed_transports.clone(),
+                };
                 p2p.broadcast(&get_addrs).await;
 
                 // Wait for a hosts store update event