Bläddra i källkod

Fixed 3 bugs:

1. Skip writing or reading anything from the socket when payload size is 0
2. Ping/Pong write their payload in pack()
3. Race conditions caused by having protocol subscriptions inside spawned tasks. Instead these have now been moved to class constructors. This enables messages to be queued for that protocol subsystem while it is setting up itself. Once it is ready the messages can be read and processed. See the changes for files protocols/protocol_address.rs and protocols/protocol_version.rs
narodnik 5 år sedan
förälder
incheckning
2c0d8c9a8d

+ 1 - 3
src/bin/dfi.rs

@@ -344,9 +344,7 @@ fn main() -> Result<()> {
 
     let options = ProgramOptions::load()?;
 
-    let logger_config = ConfigBuilder::new()
-        .set_time_format_str("%T%.6f")
-        .build();
+    let logger_config = ConfigBuilder::new().set_time_format_str("%T%.6f").build();
 
     CombinedLogger::init(vec![
         TermLogger::new(LevelFilter::Debug, logger_config, TerminalMode::Mixed).unwrap(),

+ 19 - 10
src/net/messages.rs

@@ -25,10 +25,10 @@ pub type CiphertextHash = [u8; 32];
 #[derive(IntoPrimitive, TryFromPrimitive, Copy, Clone, PartialEq, Eq, Hash, Debug)]
 #[repr(u8)]
 pub enum PacketType {
-    Ping = 0,
-    Pong = 1,
-    GetAddrs = 2,
-    Addrs = 3,
+    Ping = 1,
+    Pong = 2,
+    GetAddrs = 3,
+    Addrs = 4,
     Inv = 5,
     GetSlabs = 6,
     Slab = 7,
@@ -237,7 +237,7 @@ impl Message {
                 message.encode(&mut payload)?;
                 Ok(Packet {
                     command: PacketType::Ping,
-                    payload: Vec::new(),
+                    payload,
                 })
             }
             Message::Pong(message) => {
@@ -245,7 +245,7 @@ impl Message {
                 message.encode(&mut payload)?;
                 Ok(Packet {
                     command: PacketType::Pong,
-                    payload: Vec::new(),
+                    payload,
                 })
             }
             Message::GetAddrs(message) => {
@@ -352,36 +352,45 @@ pub async fn read_packet<R: AsyncRead + Unpin>(stream: &mut R) -> Result<Packet>
 
     // The type of the message
     let command = AsyncReadExt::read_u8(stream).await?;
-    //debug!(target: "net", "read command: {}", command);
+    debug!(target: "net", "read command: {}", command);
     let command = PacketType::try_from(command).map_err(|_| Error::MalformedPacket)?;
 
     let payload_len = VarInt::decode_async(stream).await?.0 as usize;
 
     // The message-dependent data (see message types)
     let mut payload = vec![0u8; payload_len];
-    stream.read_exact(&mut payload).await?;
-    //debug!(target: "net", "read payload");
+    if payload_len > 0 {
+        stream.read_exact(&mut payload).await?;
+    }
+    debug!(target: "net", "read payload {} bytes", payload_len);
 
     Ok(Packet { command, payload })
 }
 
 pub async fn send_packet<W: AsyncWrite + Unpin>(stream: &mut W, packet: Packet) -> Result<()> {
+    debug!(target: "net", "sending magic...");
     stream.write_all(&MAGIC_BYTES).await?;
+    debug!(target: "net", "sent magic...");
 
     AsyncWriteExt::write_u8(stream, packet.command as u8).await?;
+    debug!(target: "net", "sent command: {}", packet.command as u8);
 
     assert_eq!(std::mem::size_of::<usize>(), std::mem::size_of::<u64>());
     VarInt(packet.payload.len() as u64)
         .encode_async(stream)
         .await?;
 
-    stream.write_all(&packet.payload).await?;
+    if packet.payload.len() > 0 {
+        stream.write_all(&packet.payload).await?;
+    }
+    debug!(target: "net", "sent payload {} bytes", packet.payload.len() as u64);
 
     Ok(())
 }
 
 pub async fn receive_message<R: AsyncRead + Unpin>(stream: &mut R) -> Result<Message> {
     let packet = read_packet(stream).await?;
+    debug!(target: "net", "unpacking packet: {:?}", packet.command);
     let message = Message::unpack(packet)?;
     debug!(target: "net", "received Message::{}", message.name());
     Ok(message)

+ 20 - 15
src/net/protocols/protocol_address.rs

@@ -3,12 +3,17 @@ use smol::Executor;
 use std::sync::Arc;
 
 use crate::net::error::NetResult;
+use crate::net::message_subscriber::MessageSubscription;
 use crate::net::messages;
 use crate::net::protocols::{ProtocolJobsManager, ProtocolJobsManagerPtr};
 use crate::net::{ChannelPtr, HostsPtr, SettingsPtr};
 
 pub struct ProtocolAddress {
     channel: ChannelPtr,
+
+    addrs_sub: MessageSubscription,
+    get_addrs_sub: MessageSubscription,
+
     hosts: HostsPtr,
     settings: SettingsPtr,
 
@@ -16,9 +21,21 @@ pub struct ProtocolAddress {
 }
 
 impl ProtocolAddress {
-    pub fn new(channel: ChannelPtr, hosts: HostsPtr, settings: SettingsPtr) -> Arc<Self> {
+    pub async fn new(channel: ChannelPtr, hosts: HostsPtr, settings: SettingsPtr) -> Arc<Self> {
+        let addrs_sub = channel
+            .clone()
+            .subscribe_msg(messages::PacketType::Addrs)
+            .await;
+
+        let get_addrs_sub = channel
+            .clone()
+            .subscribe_msg(messages::PacketType::GetAddrs)
+            .await;
+
         Arc::new(Self {
             channel: channel.clone(),
+            addrs_sub,
+            get_addrs_sub,
             hosts,
             settings,
             jobsman: ProtocolJobsManager::new("ProtocolAddress", channel),
@@ -45,14 +62,8 @@ impl ProtocolAddress {
 
     async fn handle_receive_addrs(self: Arc<Self>) -> NetResult<()> {
         debug!(target: "net", "ProtocolAddress::handle_receive_addrs() [START]");
-        let addrs_sub = self
-            .channel
-            .clone()
-            .subscribe_msg(messages::PacketType::Addrs)
-            .await;
-
         loop {
-            let addrs_msg = receive_message!(addrs_sub, messages::Message::Addrs);
+            let addrs_msg = receive_message!(self.addrs_sub, messages::Message::Addrs);
 
             debug!(target: "net", "ProtocolAddress::handle_receive_addrs() storing address in hosts");
             self.hosts.store(addrs_msg.addrs.clone()).await;
@@ -61,14 +72,8 @@ impl ProtocolAddress {
 
     async fn handle_receive_get_addrs(self: Arc<Self>) -> NetResult<()> {
         debug!(target: "net", "ProtocolAddress::handle_receive_get_addrs() [START]");
-        let get_addrs_sub = self
-            .channel
-            .clone()
-            .subscribe_msg(messages::PacketType::GetAddrs)
-            .await;
-
         loop {
-            let _get_addrs = receive_message!(get_addrs_sub, messages::Message::GetAddrs);
+            let _get_addrs = receive_message!(self.get_addrs_sub, messages::Message::GetAddrs);
 
             debug!(target: "net", "ProtocolAddress::handle_receive_get_addrs() received GetAddrs message");
 

+ 23 - 16
src/net/protocols/protocol_version.rs

@@ -4,18 +4,36 @@ use smol::Executor;
 use std::sync::Arc;
 
 use crate::net::error::{NetError, NetResult};
+use crate::net::message_subscriber::MessageSubscription;
 use crate::net::messages;
 use crate::net::utility::sleep;
 use crate::net::{ChannelPtr, SettingsPtr};
 
 pub struct ProtocolVersion {
     channel: ChannelPtr,
+    version_sub: MessageSubscription,
+    verack_sub: MessageSubscription,
     settings: SettingsPtr,
 }
 
 impl ProtocolVersion {
-    pub fn new(channel: ChannelPtr, settings: SettingsPtr) -> Arc<Self> {
-        Arc::new(Self { channel, settings })
+    pub async fn new(channel: ChannelPtr, settings: SettingsPtr) -> Arc<Self> {
+        let version_sub = channel
+            .clone()
+            .subscribe_msg(messages::PacketType::Version)
+            .await;
+
+        let verack_sub = channel
+            .clone()
+            .subscribe_msg(messages::PacketType::Verack)
+            .await;
+
+        Arc::new(Self {
+            channel,
+            version_sub,
+            verack_sub,
+            settings,
+        })
     }
 
     pub async fn run(self: Arc<Self>, executor: Arc<Executor<'_>>) -> NetResult<()> {
@@ -34,6 +52,7 @@ impl ProtocolVersion {
 
     async fn exchange_versions(self: Arc<Self>, executor: Arc<Executor<'_>>) -> NetResult<()> {
         debug!(target: "net", "ProtocolVersion::exchange_versions() [START]");
+
         let send = executor.spawn(self.clone().send_version());
         let recv = executor.spawn(self.recv_version());
 
@@ -44,17 +63,11 @@ impl ProtocolVersion {
 
     async fn send_version(self: Arc<Self>) -> NetResult<()> {
         debug!(target: "net", "ProtocolVersion::send_version() [START]");
-        let verack_sub = self
-            .channel
-            .clone()
-            .subscribe_msg(messages::PacketType::Verack)
-            .await;
-
         let version = messages::Message::Version(messages::VersionMessage {});
         self.channel.clone().send(version).await?;
 
         // Wait for version acknowledgement
-        let _verack_msg = verack_sub.receive().await?;
+        let _verack_msg = self.verack_sub.receive().await?;
 
         debug!(target: "net", "ProtocolVersion::send_version() [END]");
         Ok(())
@@ -62,13 +75,7 @@ impl ProtocolVersion {
 
     async fn recv_version(self: Arc<Self>) -> NetResult<()> {
         debug!(target: "net", "ProtocolVersion::recv_version() [START]");
-        let version_sub = self
-            .channel
-            .clone()
-            .subscribe_msg(messages::PacketType::Version)
-            .await;
-
-        let _version_msg = version_sub.receive().await?;
+        let _version_msg = self.version_sub.receive().await?;
 
         // Check the message is OK
 

+ 2 - 2
src/net/sessions/inbound_session.rs

@@ -108,9 +108,9 @@ impl InboundSession {
         let hosts = self.p2p().hosts().clone();
 
         let protocol_ping = ProtocolPing::new(channel.clone(), settings.clone());
-        protocol_ping.start(executor.clone()).await;
+        let protocol_addr = ProtocolAddress::new(channel, hosts, settings).await;
 
-        let protocol_addr = ProtocolAddress::new(channel, hosts, settings);
+        protocol_ping.start(executor.clone()).await;
         protocol_addr.start(executor).await;
 
         Ok(())

+ 2 - 2
src/net/sessions/outbound_session.rs

@@ -113,9 +113,9 @@ impl OutboundSession {
         let hosts = self.p2p().hosts().clone();
 
         let protocol_ping = ProtocolPing::new(channel.clone(), settings.clone());
-        protocol_ping.start(executor.clone()).await;
+        let protocol_addr = ProtocolAddress::new(channel, hosts, settings).await;
 
-        let protocol_addr = ProtocolAddress::new(channel, hosts, settings);
+        protocol_ping.start(executor.clone()).await;
         protocol_addr.start(executor).await;
 
         Ok(())

+ 9 - 6
src/net/sessions/session.rs

@@ -31,7 +31,10 @@ pub trait Session: Sync {
         executor: Arc<Executor<'_>>,
     ) -> NetResult<()> {
         debug!(target: "net", "Session::register_channel() [START]");
-        let handshake_task = self.perform_handshake_protocols(channel.clone(), executor.clone());
+
+        let protocol_version = ProtocolVersion::new(channel.clone(), self.p2p().settings()).await;
+        let handshake_task =
+            self.perform_handshake_protocols(protocol_version, channel.clone(), executor.clone());
 
         // start channel
         channel.start(executor);
@@ -44,22 +47,22 @@ pub trait Session: Sync {
 
     async fn perform_handshake_protocols(
         &self,
+        protocol_version: Arc<ProtocolVersion>,
         channel: ChannelPtr,
         executor: Arc<Executor<'_>>,
     ) -> NetResult<()> {
-        let p2p = self.p2p();
-
         // Perform handshake
-        let protocol_version = ProtocolVersion::new(channel.clone(), p2p.settings());
         protocol_version.run(executor.clone()).await?;
 
         // Channel is now initialized
 
         // Add channel to p2p
-        p2p.clone().store(channel.clone()).await;
+        self.p2p().clone().store(channel.clone()).await;
 
         // Subscribe to stop, so can remove from p2p
-        executor.spawn(remove_sub_on_stop(p2p, channel)).detach();
+        executor
+            .spawn(remove_sub_on_stop(self.p2p(), channel))
+            .detach();
 
         // Channel is ready for use
         Ok(())