Przeglądaj źródła

net/session/direct_session: spawn a task to create a channel, fix race condition

epiphany 9 miesięcy temu
rodzic
commit
27dd21cc50
1 zmienionych plików z 184 dodań i 140 usunięć
  1. 184 140
      src/net/session/direct_session.rs

+ 184 - 140
src/net/session/direct_session.rs

@@ -36,8 +36,8 @@ use std::{
 };
 };
 
 
 use async_trait::async_trait;
 use async_trait::async_trait;
-use smol::lock::Mutex as AsyncMutex;
-use tracing::{debug, error, warn};
+use smol::lock::{Mutex as AsyncMutex, OnceCell};
+use tracing::{error, warn};
 use url::Url;
 use url::Url;
 
 
 use super::{
 use super::{
@@ -52,7 +52,9 @@ use super::{
 };
 };
 use crate::{
 use crate::{
     net::ChannelPtr,
     net::ChannelPtr,
-    system::{sleep, timeout::timeout, CondVar, PublisherPtr, StoppableTask, StoppableTaskPtr},
+    system::{
+        msleep, sleep, timeout::timeout, CondVar, PublisherPtr, StoppableTask, StoppableTaskPtr,
+    },
     util::logger::verbose,
     util::logger::verbose,
     Error, Result,
     Error, Result,
 };
 };
@@ -63,8 +65,8 @@ pub type DirectSessionPtr = Arc<DirectSession>;
 pub struct DirectSession {
 pub struct DirectSession {
     /// Weak pointer to parent p2p object
     /// Weak pointer to parent p2p object
     pub(in crate::net) p2p: Weak<P2p>,
     pub(in crate::net) p2p: Weak<P2p>,
-    /// Service to create direct channels
-    channel_builder: Arc<AsyncMutex<ChannelBuilder>>,
+    /// Connector to create direct connections
+    connector: OnceCell<Connector>,
     /// Tasks that are trying to create a direct channel (they retry until they succeed).
     /// Tasks that are trying to create a direct channel (they retry until they succeed).
     /// A task is removed once the channel is successfully created.
     /// A task is removed once the channel is successfully created.
     retries_tasks: Arc<AsyncMutex<HashMap<Url, Arc<StoppableTask>>>>,
     retries_tasks: Arc<AsyncMutex<HashMap<Url, Arc<StoppableTask>>>>,
@@ -72,6 +74,8 @@ pub struct DirectSession {
     peer_discovery: Arc<PeerDiscovery>,
     peer_discovery: Arc<PeerDiscovery>,
     /// Channel ID -> usage count
     /// Channel ID -> usage count
     channels_usage: Arc<AsyncMutex<HashMap<u32, u32>>>,
     channels_usage: Arc<AsyncMutex<HashMap<u32, u32>>>,
+    /// Pending channel creation tasks
+    tasks: Arc<AsyncMutex<HashMap<Url, Weak<ChannelTask>>>>,
 }
 }
 
 
 impl DirectSession {
 impl DirectSession {
@@ -79,10 +83,11 @@ impl DirectSession {
     pub fn new(p2p: Weak<P2p>) -> DirectSessionPtr {
     pub fn new(p2p: Weak<P2p>) -> DirectSessionPtr {
         Arc::new_cyclic(|session| Self {
         Arc::new_cyclic(|session| Self {
             p2p,
             p2p,
-            channel_builder: Arc::new(AsyncMutex::new(ChannelBuilder::new(session.clone()))),
+            connector: OnceCell::new(),
             retries_tasks: Arc::new(AsyncMutex::new(HashMap::new())),
             retries_tasks: Arc::new(AsyncMutex::new(HashMap::new())),
             peer_discovery: PeerDiscovery::new(session.clone()),
             peer_discovery: PeerDiscovery::new(session.clone()),
             channels_usage: Arc::new(AsyncMutex::new(HashMap::new())),
             channels_usage: Arc::new(AsyncMutex::new(HashMap::new())),
+            tasks: Arc::new(AsyncMutex::new(HashMap::new())),
         })
         })
     }
     }
 
 
@@ -112,65 +117,102 @@ impl DirectSession {
     /// If there is an existing channel to the same address, this method will
     /// If there is an existing channel to the same address, this method will
     /// return it (even if the channel was not created by the direct session).
     /// return it (even if the channel was not created by the direct session).
     /// Otherwise it will create a new channel to `addr` in the direct session.
     /// Otherwise it will create a new channel to `addr` in the direct session.
-    pub async fn get_channel(&self, addr: &Url) -> Result<ChannelPtr> {
+    pub async fn get_channel(self: Arc<Self>, addr: &Url) -> Result<ChannelPtr> {
+        // Check existing channels
         let channels = self.p2p().hosts().channels();
         let channels = self.p2p().hosts().channels();
-        let channel = channels.iter().find(|&chan| chan.address() == addr).cloned();
-        if let Some(channel) = channel {
+        if let Some(channel) =
+            channels.iter().find(|&chan| chan.info.connect_addr == *addr).cloned()
+        {
             let mut channels_usage = self.channels_usage.lock().await;
             let mut channels_usage = self.channels_usage.lock().await;
-            let usage_count = channels_usage.get_mut(&channel.info.id);
-            if let Some(count) = usage_count {
-                *count += 1;
+            if channel.is_stopped() {
+                channel.clone().start(self.p2p().executor());
+            }
+            if channel.session_type_id() & SESSION_DIRECT != 0 {
+                channels_usage.entry(channel.info.id).and_modify(|count| *count += 1).or_insert(1);
             }
             }
-            return Ok(channel)
+            return Ok(channel);
         }
         }
-        self.channel_builder.lock().await.new_channel(addr).await
+
+        let mut tasks = self.tasks.lock().await;
+
+        // Check if task is already running for this addr
+        if let Some(task) = tasks.get(addr) {
+            let task = task.upgrade().unwrap();
+            drop(tasks);
+
+            // Wait for the existing task to complete
+            while task.output.lock().await.is_none() {
+                msleep(100).await; // Wait for completion
+            }
+            return task.output.lock().await.clone().unwrap();
+        }
+
+        // If no task running, create one
+        let task = Arc::new(ChannelTask {
+            session: Arc::downgrade(&self.clone()),
+            addr: addr.clone(),
+            output: Arc::new(AsyncMutex::new(None)),
+        });
+        tasks.insert(addr.clone(), Arc::downgrade(&task));
+        drop(tasks);
+
+        // Spawn a new task to create the channel
+        let ex = self.p2p().executor();
+        let addr_ = addr.clone();
+        let self_ = self.clone();
+        let task_ = task.clone();
+        ex.spawn(async move {
+            let res = self_.clone().new_channel(addr_.clone()).await;
+
+            let mut output = task_.output.lock().await;
+            *output = Some(res);
+        })
+        .detach();
+
+        // Wait for completion
+        while task.output.lock().await.is_none() {
+            msleep(100).await;
+        }
+        let res = task.output.lock().await.as_ref().unwrap().clone();
+        if let Ok(ref channel) = res {
+            self.inc_channel_usage(channel, Arc::strong_count(&task).try_into().unwrap()).await;
+        }
+        res
+    }
+
+    /// Increment channel usage
+    async fn inc_channel_usage(&self, channel: &ChannelPtr, n: u32) {
+        let mut channels_usage = self.channels_usage.lock().await;
+        channels_usage.entry(channel.info.id).and_modify(|count| *count += n).or_insert(n);
     }
     }
 
 
     /// Try to create a new channel until it succeeds, then notify `channel_pub`.
     /// Try to create a new channel until it succeeds, then notify `channel_pub`.
     /// If it fails to create a channel, a task will sleep
     /// If it fails to create a channel, a task will sleep
     /// `outbound_connect_timeout` seconds and try again.
     /// `outbound_connect_timeout` seconds and try again.
-    pub async fn get_channel_with_retries(&self, addr: Url, channel_pub: PublisherPtr<ChannelPtr>) {
-        let channel_builder = self.channel_builder.clone();
+    pub async fn get_channel_with_retries(
+        self: Arc<Self>,
+        addr: Url,
+        channel_pub: PublisherPtr<ChannelPtr>,
+    ) {
         let task = StoppableTask::new();
         let task = StoppableTask::new();
-        let retries_tasks_lock = self.retries_tasks.clone();
+        let self_ = self.clone();
         let mut retries_tasks = self.retries_tasks.lock().await;
         let mut retries_tasks = self.retries_tasks.lock().await;
-        let channels_usage = self.channels_usage.clone();
-        let p2p = self.p2p().clone();
         retries_tasks.insert(addr.clone(), task.clone());
         retries_tasks.insert(addr.clone(), task.clone());
         drop(retries_tasks);
         drop(retries_tasks);
 
 
         task.clone().start(
         task.clone().start(
             async move {
             async move {
                 loop {
                 loop {
-                    // Check if there is already a channel to this addr
-                    let channels = p2p.hosts().channels();
-                    let channel = channels.iter().find(|&chan| chan.address() == &addr).cloned();
-                    if let Some(channel) = channel {
-                        let mut chan_usage = channels_usage.lock().await;
-                        let usage_count = chan_usage.get_mut(&channel.info.id);
-                        if let Some(count) = usage_count {
-                            *count += 1;
-                        }
-                        channel_pub.notify(channel).await;
-                        let mut retries_tasks = retries_tasks_lock.lock().await;
-                        retries_tasks.remove(&addr);
-                        break
-                    }
-
-                    // Try to create a new channel
-                    let mut builder = channel_builder.lock().await;
-                    let res = builder.new_channel(&addr).await;
+                    let res = self_.clone().get_channel(&addr).await;
                     match res {
                     match res {
                         Ok(channel) => {
                         Ok(channel) => {
                             channel_pub.notify(channel).await;
                             channel_pub.notify(channel).await;
-                            let mut retries_tasks = retries_tasks_lock.lock().await;
+                            let mut retries_tasks = self_.retries_tasks.lock().await;
                             retries_tasks.remove(&addr);
                             retries_tasks.remove(&addr);
                             break
                             break
                         }
                         }
-                        Err(Error::HostDoesNotExist) => break,
                         Err(_) => {
                         Err(_) => {
-                            drop(builder);
-                            let settings = p2p.settings().read_arc().await;
+                            let settings = self_.p2p().settings().read_arc().await;
                             sleep(settings.outbound_connect_timeout).await;
                             sleep(settings.outbound_connect_timeout).await;
                         }
                         }
                     }
                     }
@@ -182,7 +224,7 @@ impl DirectSession {
                 match res {
                 match res {
                     Ok(()) | Err(Error::DetachedTaskStopped) => { /* Do nothing */ }
                     Ok(()) | Err(Error::DetachedTaskStopped) => { /* Do nothing */ }
                     Err(e) => {
                     Err(e) => {
-                        error!(target: "net::direct_session::create_channel_with_retries()", "{e}")
+                        error!(target: "net::direct_session::get_channel_with_retries()", "{e}")
                     }
                     }
                 }
                 }
             },
             },
@@ -191,97 +233,27 @@ impl DirectSession {
         );
         );
     }
     }
 
 
-    /// Close a direct channel if it's not used by anything.
-    /// `AsyncDrop` would be great here (<https://doc.rust-lang.org/std/future/trait.AsyncDrop.html>)
-    /// but it's still in nightly. For now you must call this method manually
-    /// once you are done with a direct channel.
-    /// Returns `true` if the channel is stopped.
-    pub async fn cleanup_channel(&self, channel: ChannelPtr) -> bool {
-        if channel.session_type_id() & SESSION_DIRECT == 0 {
-            // Do nothing if this is not a channel created by the direct session
-            return false
-        }
-
-        let mut channels_usage = self.channels_usage.lock().await;
-        let usage_count = channels_usage.get_mut(&channel.info.id);
-        if usage_count.is_none() {
-            channel.stop().await;
-            return true
-        }
-        let usage_count = usage_count.unwrap();
-        if *usage_count > 0 {
-            *usage_count -= 1;
-        }
-
-        if *usage_count == 0 {
-            channels_usage.remove(&channel.info.id);
-            channel.stop().await;
-            return true
-        }
-
-        false
-    }
-}
-
-#[async_trait]
-impl Session for DirectSession {
-    fn p2p(&self) -> P2pPtr {
-        self.p2p.upgrade().unwrap()
-    }
-
-    fn type_id(&self) -> SessionBitFlag {
-        SESSION_DIRECT
-    }
-}
-
-pub struct ChannelBuilder {
-    /// Weak pointer to parent object
-    session: Weak<DirectSession>,
-    connector: Option<Arc<Connector>>,
-}
-
-impl ChannelBuilder {
-    pub fn new(session: Weak<DirectSession>) -> Self {
-        Self { session: session.clone(), connector: None }
-    }
-
-    fn session(&self) -> DirectSessionPtr {
-        self.session.upgrade().unwrap()
-    }
-
-    fn p2p(&self) -> P2pPtr {
-        self.session().p2p()
-    }
-
-    fn connector(&mut self) -> Arc<Connector> {
-        match &self.connector {
-            Some(c) => c.clone(),
-            None => {
-                self.connector = Some(Arc::new(Connector::new(
-                    self.session().p2p().settings(),
-                    self.session.clone(),
-                )));
-                self.connector.clone().unwrap()
-            }
+    async fn new_channel(self: Arc<Self>, addr: Url) -> Result<ChannelPtr> {
+        if !self.connector.is_initialized() {
+            let _ = self
+                .connector
+                .set(Connector::new(self.p2p().settings(), Arc::downgrade(&self.clone()).clone()))
+                .await;
         }
         }
-    }
 
 
-    /// Create a new channel to `addr` in the direct session.
-    /// The transport is verified before the connection is started.
-    pub async fn new_channel(&mut self, addr: &Url) -> Result<ChannelPtr> {
         verbose!(
         verbose!(
             target: "net::direct_session",
             target: "net::direct_session",
             "[P2P] Connecting to direct outbound [{addr}]",
             "[P2P] Connecting to direct outbound [{addr}]",
         );
         );
 
 
-        let settings = self.session().p2p().settings().read_arc().await;
+        let settings = self.p2p().settings().read_arc().await;
         let seeds = settings.seeds.clone();
         let seeds = settings.seeds.clone();
         let allowed_transports = settings.allowed_transports.clone();
         let allowed_transports = settings.allowed_transports.clone();
         drop(settings);
         drop(settings);
 
 
         // Do not establish a connection to a host that is also configured as a seed.
         // Do not establish a connection to a host that is also configured as a seed.
         // This indicates a user misconfiguration.
         // This indicates a user misconfiguration.
-        if seeds.contains(addr) {
+        if seeds.contains(&addr) {
             error!(
             error!(
                 target: "net::direct_session",
                 target: "net::direct_session",
                 "[P2P] Suspending direct connection to seed [{}]", addr.clone(),
                 "[P2P] Suspending direct connection to seed [{}]", addr.clone(),
@@ -290,9 +262,9 @@ impl ChannelBuilder {
         }
         }
 
 
         // Abort if we are trying to connect to our own external address.
         // Abort if we are trying to connect to our own external address.
-        let hosts = self.session().p2p().hosts();
+        let hosts = self.p2p().hosts();
         let external_addrs = hosts.external_addrs().await;
         let external_addrs = hosts.external_addrs().await;
-        if external_addrs.contains(addr) {
+        if external_addrs.contains(&addr) {
             warn!(
             warn!(
                 target: "net::hosts::check_addrs",
                 target: "net::hosts::check_addrs",
                 "[P2P] Suspending direct connection to external addr [{}]", addr.clone(),
                 "[P2P] Suspending direct connection to external addr [{}]", addr.clone(),
@@ -308,21 +280,35 @@ impl ChannelBuilder {
         }
         }
 
 
         // Abort if this peer is IPv6 and we do not support it.
         // Abort if this peer is IPv6 and we do not support it.
-        if !hosts.ipv6_available.load(Ordering::SeqCst) && hosts.is_ipv6(addr) {
+        if !hosts.ipv6_available.load(Ordering::SeqCst) && hosts.is_ipv6(&addr) {
             return Err(Error::ConnectFailed(format!("[{addr}]: IPv6 is unavailable")))
             return Err(Error::ConnectFailed(format!("[{addr}]: IPv6 is unavailable")))
         }
         }
 
 
-        if let Err(e) = hosts.try_register(addr.clone(), HostState::Connect) {
-            debug!(target: "net::direct_session",
-                "[P2P] Cannot connect to direct={addr}, err={e}");
-            return Err(e)
+        // Set the addr to HostState::Connect
+        loop {
+            if let Err(e) = hosts.try_register(addr.clone(), HostState::Connect) {
+                // If `try_register` failed because the addr is being refined, try again in a bit.
+                if let Error::HostStateBlocked(from, _) = &e {
+                    if from == "Refine" {
+                        // TODO: Add a setting or have a way to wait for the refinery to complete
+                        sleep(5).await;
+                        continue
+                    }
+                }
+
+                error!(target: "net::direct_session",
+                    "[P2P] Cannot connect to direct={addr}, err={e}");
+                return Err(e)
+            }
+            break
         }
         }
 
 
         dnetev!(self, DirectConnecting, {
         dnetev!(self, DirectConnecting, {
             connect_addr: addr.clone(),
             connect_addr: addr.clone(),
         });
         });
 
 
-        match self.connector().connect(addr).await {
+        // Attempt channel creation
+        match self.connector.get().unwrap().connect(&addr).await {
             Ok((_, channel)) => {
             Ok((_, channel)) => {
                 verbose!(
                 verbose!(
                     target: "net::direct_session",
                     target: "net::direct_session",
@@ -337,15 +323,8 @@ impl ChannelBuilder {
                 });
                 });
 
 
                 // Register the new channel
                 // Register the new channel
-                match self
-                    .session()
-                    .register_channel(channel.clone(), self.session().p2p().executor())
-                    .await
-                {
-                    Ok(()) => {
-                        self.session().channels_usage.lock().await.insert(channel.info.id, 1);
-                        Ok(channel)
-                    }
+                match self.register_channel(channel.clone(), self.p2p().executor()).await {
+                    Ok(()) => Ok(channel),
                     Err(e) => {
                     Err(e) => {
                         warn!(
                         warn!(
                             target: "net::direct_session",
                             target: "net::direct_session",
@@ -359,7 +338,7 @@ impl ChannelBuilder {
                         });
                         });
 
 
                         // Free up this addr for future operations.
                         // Free up this addr for future operations.
-                        if let Err(e) = self.session().p2p().hosts().unregister(channel.address()) {
+                        if let Err(e) = self.p2p().hosts().unregister(channel.address()) {
                             warn!(target: "net::direct_session", "[P2P] Error while unregistering addr={}, err={e}", channel.display_address());
                             warn!(target: "net::direct_session", "[P2P] Error while unregistering addr={}, err={e}", channel.display_address());
                         }
                         }
 
 
@@ -379,7 +358,7 @@ impl ChannelBuilder {
                 });
                 });
 
 
                 // Free up this addr for future operations.
                 // Free up this addr for future operations.
-                if let Err(e) = self.session().p2p().hosts().unregister(addr) {
+                if let Err(e) = self.p2p().hosts().unregister(&addr) {
                     warn!(target: "net::direct_session", "[P2P] Error while unregistering addr={addr}, err={e}");
                     warn!(target: "net::direct_session", "[P2P] Error while unregistering addr={addr}, err={e}");
                 }
                 }
 
 
@@ -387,6 +366,71 @@ impl ChannelBuilder {
             }
             }
         }
         }
     }
     }
+
+    /// Close a direct channel if it's not used by anything.
+    /// `AsyncDrop` would be great here (<https://doc.rust-lang.org/std/future/trait.AsyncDrop.html>)
+    /// but it's still in nightly. For now you must call this method manually
+    /// once you are done with a direct channel.
+    /// Returns `true` if the channel is stopped.
+    pub async fn cleanup_channel(self: Arc<Self>, channel: ChannelPtr) -> bool {
+        if channel.session_type_id() & SESSION_DIRECT == 0 {
+            // Do nothing if this is not a channel created by the direct session
+            return false
+        }
+
+        let mut channels_usage = self.channels_usage.lock().await;
+        let usage_count = channels_usage.get_mut(&channel.info.id);
+        if usage_count.is_none() {
+            let _ = self.p2p().hosts().unregister(channel.address());
+            channel.stop().await;
+            return true
+        }
+        let usage_count = usage_count.unwrap();
+        if *usage_count > 0 {
+            *usage_count -= 1;
+        }
+
+        if *usage_count == 0 {
+            channels_usage.remove(&channel.info.id);
+            let _ = self.p2p().hosts().unregister(channel.address());
+            channel.stop().await;
+            return true
+        }
+
+        false
+    }
+}
+
+#[async_trait]
+impl Session for DirectSession {
+    fn p2p(&self) -> P2pPtr {
+        self.p2p.upgrade().unwrap()
+    }
+
+    fn type_id(&self) -> SessionBitFlag {
+        SESSION_DIRECT
+    }
+}
+
+struct ChannelTask {
+    session: Weak<DirectSession>,
+    addr: Url,
+    output: Arc<AsyncMutex<Option<Result<ChannelPtr>>>>,
+}
+
+impl Drop for ChannelTask {
+    fn drop(&mut self) {
+        let session = self.session.upgrade().unwrap();
+        let addr = self.addr.clone();
+        session
+            .p2p()
+            .executor()
+            .spawn(async move {
+                let mut tasks = session.tasks.lock().await;
+                tasks.remove(&addr);
+            })
+            .detach();
+    }
 }
 }
 
 
 /// PeerDiscovery process for that sends `GetAddrs` messages to a random
 /// PeerDiscovery process for that sends `GetAddrs` messages to a random
@@ -580,7 +624,7 @@ impl PeerDiscovery {
 
 
             // Stop the channel we created for peer discovery
             // Stop the channel we created for peer discovery
             if let Some(ch) = channel {
             if let Some(ch) = channel {
-                ch.stop().await;
+                self.p2p().session_direct().cleanup_channel(ch).await;
             }
             }
 
 
             // Give some time for new connections to be established
             // Give some time for new connections to be established