Explorar o código

migrate other protocols (server and seed) to Arc<Self>

narodnik %!s(int64=5) %!d(string=hai) anos
pai
achega
6df824ffb6
Modificáronse 3 ficheiros con 81 adicións e 114 borrados
  1. 8 19
      src/bin/dfi.rs
  2. 46 60
      src/net/protocol/seed_protocol.rs
  3. 27 35
      src/net/protocol/server_protocol.rs

+ 8 - 19
src/bin/dfi.rs

@@ -6,7 +6,7 @@ use async_std::sync::Mutex;
 use easy_parallel::Parallel;
 use log::*;
 use std::collections::HashMap;
-use std::net::{IpAddr, Ipv4Addr, SocketAddr};
+use std::net::SocketAddr;
 use std::sync::Arc;
 
 use sapvi::{ClientProtocol, Result, SeedProtocol, ServerProtocol};
@@ -23,11 +23,9 @@ async fn start(executor: Arc<Executor<'_>>, options: ProgramOptions) -> Result<(
     if let Some(accept_addr) = options.accept_addr {
         let accept_addr = accept_addr.clone();
 
-        let mut protocol = ServerProtocol::new(connections.clone());
+        let protocol = ServerProtocol::new(connections.clone(), accept_addr, stored_addrs2);
         server_task = Some(executor.spawn(async move {
-            protocol
-                .start(accept_addr, stored_addrs2, executor2)
-                .await?;
+            protocol.start(executor2).await?;
             Ok::<(), sapvi::Error>(())
         }));
     }
@@ -35,18 +33,11 @@ async fn start(executor: Arc<Executor<'_>>, options: ProgramOptions) -> Result<(
     let mut seed_protocols = Vec::with_capacity(options.seed_addrs.len());
 
     // Normally we query this from a server
-    let local_addr = options.accept_addr.clone();
+    let accept_addr = options.accept_addr.clone();
 
     for seed_addr in options.seed_addrs.iter() {
-        let mut protocol = SeedProtocol::new();
-        protocol
-            .start(
-                seed_addr.clone(),
-                local_addr,
-                stored_addrs.clone(),
-                executor.clone(),
-            )
-            .await;
+        let protocol = SeedProtocol::new(seed_addr.clone(), accept_addr, stored_addrs.clone());
+        protocol.clone().start(executor.clone()).await;
         seed_protocols.push(protocol);
     }
 
@@ -58,13 +49,11 @@ async fn start(executor: Arc<Executor<'_>>, options: ProgramOptions) -> Result<(
 
     debug!("Seed nodes queried.");
 
-    let accept_addr = options.accept_addr.clone();
-
     let mut client_slots = vec![];
     for i in 0..options.connection_slots {
         debug!("Starting connection slot {}", i);
 
-        let mut client = ClientProtocol::new(
+        let client = ClientProtocol::new(
             connections.clone(),
             accept_addr.clone(),
             stored_addrs.clone(),
@@ -76,7 +65,7 @@ async fn start(executor: Arc<Executor<'_>>, options: ProgramOptions) -> Result<(
     for remote_addr in options.manual_connects {
         debug!("Starting connection (manual) to {}", remote_addr);
 
-        let mut client = ClientProtocol::new(
+        let client = ClientProtocol::new(
             connections.clone(),
             accept_addr.clone(),
             stored_addrs.clone(),

+ 46 - 60
src/net/protocol/seed_protocol.rs

@@ -1,3 +1,4 @@
+use async_std::sync::Mutex;
 use log::*;
 use smol::{Async, Executor};
 use std::net::{SocketAddr, TcpStream};
@@ -14,7 +15,11 @@ type Clock = Arc<AtomicU64>;
 pub struct SeedProtocol {
     send_sx: async_channel::Sender<net::Message>,
     send_rx: async_channel::Receiver<net::Message>,
-    main_process: Option<smol::Task<()>>,
+    main_process: Mutex<Option<smol::Task<()>>>,
+
+    seed_addr: SocketAddr,
+    accept_addr: Option<SocketAddr>,
+    stored_addrs: AddrsStorage,
 }
 
 #[derive(PartialEq)]
@@ -25,111 +30,97 @@ enum ProtocolSignal {
 }
 
 impl SeedProtocol {
-    pub fn new() -> Self {
+    pub fn new(
+        seed_addr: SocketAddr,
+        accept_addr: Option<SocketAddr>,
+        stored_addrs: AddrsStorage,
+    ) -> Arc<Self> {
         let (send_sx, send_rx) = async_channel::unbounded::<net::Message>();
-        Self {
+        Arc::new(Self {
             send_sx,
             send_rx,
-            main_process: None,
-        }
+            main_process: Mutex::new(None),
+            seed_addr,
+            accept_addr,
+            stored_addrs,
+        })
     }
 
-    pub async fn start(
-        &mut self,
-        seed_addr: SocketAddr,
-        local_addr: Option<SocketAddr>,
-        stored_addrs: AddrsStorage,
-        executor: Arc<Executor<'_>>,
-    ) {
-        let (send_sx, send_rx) = (self.send_sx.clone(), self.send_rx.clone());
-        let ex = executor.clone();
-        self.main_process = Some(ex.spawn(async move {
-            match Async::<TcpStream>::connect(seed_addr.clone()).await {
+    pub async fn start(self: Arc<Self>, executor: Arc<Executor<'_>>) {
+        let executor2 = executor.clone();
+        let self2 = self.clone();
+
+        *self2.main_process.lock().await = Some(executor.spawn(async move {
+            match Async::<TcpStream>::connect(self.seed_addr).await {
                 Ok(stream) => {
-                    let _ = Self::handle_connect(
-                        stream,
-                        &stored_addrs,
-                        seed_addr.clone(),
-                        local_addr,
-                        (send_sx.clone(), send_rx.clone()),
-                        executor.clone(),
-                    )
-                    .await;
+                    let _ = self.handle_connect(stream, executor2).await;
                 }
                 Err(err) => {
-                    warn!("Unable to connect to seed {}: {}", seed_addr, err)
+                    warn!("Unable to connect to seed {}: {}", self.seed_addr, err)
                 }
             }
         }));
     }
 
-    pub async fn await_finish(self) {
-        if let Some(process) = self.main_process {
+    pub async fn await_finish(self: Arc<Self>) {
+        let mut process = self.main_process.lock().await;
+        if let Some(process) = &mut *process {
             process.await;
         }
     }
 
     async fn handle_connect(
+        &self,
         stream: Async<TcpStream>,
-        stored_addrs: &AddrsStorage,
-        seed_addr: SocketAddr,
-        local_addr: Option<SocketAddr>,
-        (send_sx, send_rx): (
-            async_channel::Sender<net::Message>,
-            async_channel::Receiver<net::Message>,
-        ),
         executor: Arc<Executor<'_>>,
     ) -> Result<()> {
-        if let Some(local_addr) = local_addr {
-            send_sx
+        if let Some(accept_addr) = self.accept_addr {
+            self.send_sx
                 .send(net::Message::Addrs(net::AddrsMessage {
-                    addrs: vec![local_addr],
+                    addrs: vec![accept_addr],
                 }))
                 .await?;
         }
 
-        send_sx
+        self.send_sx
             .send(net::Message::GetAddrs(net::GetAddrsMessage {}))
             .await?;
 
         let stream = async_dup::Arc::new(stream);
 
         // Run event loop
-        match Self::event_loop_process(stream, stored_addrs.clone(), (send_sx, send_rx), executor)
-            .await
-        {
+        match self.event_loop_process(stream, executor).await {
             Ok(ProtocolSignal::Finished) => {
-                info!("Seed node queried successfully: {}", seed_addr);
+                info!("Seed node queried successfully: {}", self.seed_addr);
             }
             Ok(ProtocolSignal::Timeout) => {
-                warn!("Seed node timeout: {}", seed_addr);
+                warn!("Seed node timeout: {}", self.seed_addr);
             }
             Ok(_) => {
                 unreachable!();
             }
             Err(err) => {
-                warn!("Seed disconnected: {} {}", seed_addr, err);
+                warn!("Seed disconnected: {} {}", self.seed_addr, err);
             }
         }
         Ok(())
     }
 
     async fn event_loop_process(
+        &self,
         mut stream: net::AsyncTcpStream,
-        stored_addrs: AddrsStorage,
-        (send_sx, send_rx): (
-            async_channel::Sender<net::Message>,
-            async_channel::Receiver<net::Message>,
-        ),
         executor: Arc<Executor<'_>>,
     ) -> Result<ProtocolSignal> {
         let inactivity_timer = net::InactivityTimer::new(executor.clone());
 
         let clock = Arc::new(AtomicU64::new(0));
-        let _ping_task = executor.spawn(protocol_base::repeat_ping(send_sx.clone(), clock.clone()));
+        let _ping_task = executor.spawn(protocol_base::repeat_ping(
+            self.send_sx.clone(),
+            clock.clone(),
+        ));
 
         loop {
-            let event = net::select_event(&mut stream, &send_rx, &inactivity_timer).await?;
+            let event = net::select_event(&mut stream, &self.send_rx, &inactivity_timer).await?;
 
             match event {
                 net::Event::Send(message) => {
@@ -137,7 +128,7 @@ impl SeedProtocol {
                 }
                 net::Event::Receive(message) => {
                     inactivity_timer.reset().await?;
-                    let signal = Self::protocol(message, &stored_addrs, &send_sx, &clock).await?;
+                    let signal = self.protocol(message, &clock).await?;
 
                     if signal == ProtocolSignal::Finished {
                         return Ok(ProtocolSignal::Finished);
@@ -152,12 +143,7 @@ impl SeedProtocol {
         //inactivity_timer.stop().await;
     }
 
-    async fn protocol(
-        message: net::Message,
-        stored_addrs: &AddrsStorage,
-        _send_sx: &async_channel::Sender<net::Message>,
-        clock: &Clock,
-    ) -> Result<ProtocolSignal> {
+    async fn protocol(&self, message: net::Message, clock: &Clock) -> Result<ProtocolSignal> {
         match message {
             net::Message::Pong => {
                 let current_time = get_current_time();
@@ -166,7 +152,7 @@ impl SeedProtocol {
             }
             net::Message::Addrs(message) => {
                 info!("received AddrMessage");
-                let mut stored_addrs = stored_addrs.lock().await;
+                let mut stored_addrs = self.stored_addrs.lock().await;
                 for addr in message.addrs {
                     if !stored_addrs.contains(&addr) {
                         stored_addrs.push(addr);

+ 27 - 35
src/net/protocol/server_protocol.rs

@@ -13,29 +13,34 @@ pub struct ServerProtocol {
     send_sx: async_channel::Sender<net::Message>,
     send_rx: async_channel::Receiver<net::Message>,
     connections: ConnectionsMap,
+
+    accept_addr: SocketAddr,
+    stored_addrs: AddrsStorage,
 }
 
 impl ServerProtocol {
-    pub fn new(connections: ConnectionsMap) -> Self {
+    pub fn new(
+        connections: ConnectionsMap,
+        accept_addr: SocketAddr,
+        stored_addrs: AddrsStorage,
+    ) -> Arc<Self> {
         let (send_sx, send_rx) = async_channel::unbounded::<net::Message>();
-        Self {
+        Arc::new(Self {
             send_sx,
             send_rx,
             connections,
-        }
+
+            accept_addr,
+            stored_addrs,
+        })
     }
 
     pub fn get_send_pipe(&self) -> async_channel::Sender<net::Message> {
         self.send_sx.clone()
     }
 
-    pub async fn start(
-        &mut self,
-        address: SocketAddr,
-        stored_addrs: AddrsStorage,
-        executor: std::sync::Arc<Executor<'_>>,
-    ) -> Result<()> {
-        let listener = Async::<TcpListener>::bind(address)?;
+    pub async fn start(self: Arc<Self>, executor: Arc<Executor<'_>>) -> Result<()> {
+        let listener = Async::<TcpListener>::bind(self.accept_addr)?;
         info!("Listening on {}", listener.get_ref().local_addr()?);
 
         loop {
@@ -43,25 +48,17 @@ impl ServerProtocol {
             info!("Accepted client: {}", peer_addr);
             let stream = async_dup::Arc::new(stream);
 
-            let (send_sx, send_rx) = (self.send_sx.clone(), self.send_rx.clone());
-
-            let connections = self.connections.clone();
-            connections.lock().await.insert(peer_addr, send_sx.clone());
+            self.connections
+                .lock()
+                .await
+                .insert(peer_addr, self.send_sx.clone());
 
-            let stored_addrs = stored_addrs.clone();
             let executor2 = executor.clone();
+            let self2 = self.clone();
 
             executor
                 .spawn(async move {
-                    match Self::event_loop_process(
-                        stream,
-                        stored_addrs,
-                        (send_sx, send_rx),
-                        connections.clone(),
-                        executor2,
-                    )
-                    .await
-                    {
+                    match self2.clone().event_loop_process(stream, executor2).await {
                         Ok(()) => {
                             warn!("Peer {} timeout", peer_addr);
                         }
@@ -69,26 +66,21 @@ impl ServerProtocol {
                             warn!("Peer {} disconnected: {}", peer_addr, err);
                         }
                     }
-                    connections.lock().await.remove(&peer_addr);
+                    self2.connections.lock().await.remove(&peer_addr);
                 })
                 .detach();
         }
     }
 
     pub async fn event_loop_process(
+        self: Arc<Self>,
         mut stream: net::AsyncTcpStream,
-        stored_addrs: AddrsStorage,
-        (send_sx, send_rx): (
-            async_channel::Sender<net::Message>,
-            async_channel::Receiver<net::Message>,
-        ),
-        connections: ConnectionsMap,
         executor: Arc<Executor<'_>>,
     ) -> Result<()> {
         let inactivity_timer = net::InactivityTimer::new(executor.clone());
 
         loop {
-            let event = net::select_event(&mut stream, &send_rx, &inactivity_timer).await?;
+            let event = net::select_event(&mut stream, &self.send_rx, &inactivity_timer).await?;
 
             match event {
                 net::Event::Send(message) => {
@@ -98,10 +90,10 @@ impl ServerProtocol {
                     inactivity_timer.reset().await?;
                     protocol_base::protocol(
                         message,
-                        &stored_addrs,
-                        &send_sx,
+                        &self.stored_addrs,
+                        &self.send_sx,
                         None,
-                        connections.clone(),
+                        self.connections.clone(),
                     )
                     .await?;
                 }