diff --git a/src-tauri/src/network/mod.rs b/src-tauri/src/network/mod.rs index dbca543..652a4d7 100644 --- a/src-tauri/src/network/mod.rs +++ b/src-tauri/src/network/mod.rs @@ -1,4 +1,5 @@ use std::collections::HashMap; +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -15,9 +16,9 @@ use wyd_common::{ use crate::db::AppDatabase; use crate::keypair::AppKeypair; -use crate::remotes::{self, Remote}; +use crate::remotes::{self, Remote, RemotesChanged}; -type Statuses = Arc>>; +type Statuses = Arc>>; #[derive(Debug, Clone, Serialize, Deserialize, Type)] #[serde(rename_all = "camelCase")] @@ -42,18 +43,28 @@ pub struct NetworkStatusChanged { pub statuses: Vec, } +struct Connection { + remote: Remote, + sender: mpsc::Sender, + task: tauri::async_runtime::JoinHandle<()>, + generation: u64, +} + pub struct Network { - senders: Mutex>>, + connections: Mutex>, statuses: Statuses, profile: watch::Sender, + keypair: AppKeypair, + next_generation: AtomicU64, } impl Network { #[allow(dead_code)] // Ready for the first domain message sender. pub fn send(&self, remote_id: &str, payload: String) -> Result<(), String> { - let senders = self.senders.lock().map_err(|error| error.to_string())?; - let sender = senders + let connections = self.connections.lock().map_err(|error| error.to_string())?; + let sender = connections .get(remote_id) + .map(|connection| &connection.sender) .ok_or_else(|| "remote is not configured".to_string())?; sender .try_send(payload) @@ -63,35 +74,86 @@ impl Network { pub fn update_profile(&self, profile: crate::user::User) { self.profile.send_replace(profile); } + + fn sync_remotes(&self, handle: &AppHandle, remotes: Vec) -> Result<(), String> { + let desired: HashMap<_, _> = remotes + .into_iter() + .map(|remote| (remote.id.clone(), remote)) + .collect(); + let mut connections = self.connections.lock().map_err(|error| error.to_string())?; + + let stale: Vec<_> = connections + .iter() + .filter(|(id, connection)| desired.get(*id) != Some(&connection.remote)) + .map(|(id, _)| id.clone()) + .collect(); + + for id in stale { + if let Some(connection) = connections.remove(&id) { + connection.task.abort(); + remove_status(&self.statuses, &id, connection.generation); + } + } + + for remote in desired.into_values() { + if connections.contains_key(&remote.id) { + continue; + } + + let generation = self.next_generation.fetch_add(1, Ordering::Relaxed); + let (sender, receiver) = mpsc::channel(32); + set_initial( + &self.statuses, + &remote, + generation, + ConnectionState::Connecting, + ); + let task = tauri::async_runtime::spawn(run( + handle.clone(), + self.statuses.clone(), + remote.clone(), + generation, + self.profile.subscribe(), + self.keypair.clone(), + receiver, + )); + connections.insert( + remote.id.clone(), + Connection { + remote, + sender, + task, + generation, + }, + ); + } + + drop(connections); + emit_statuses(handle, &self.statuses) + } } pub async fn init(handle: &AppHandle) -> Result<(), Box> { let database = handle.state::(); let keypair = handle.state::().inner().clone(); let profile = crate::profile::get(&database, keypair.public_key()).await?; - let (profile_sender, profile_receiver) = watch::channel(profile); - let remotes = remotes::all(&database).await?; - let mut senders = HashMap::new(); - let statuses = Statuses::default(); + let (profile, _) = watch::channel(profile); + let network = Network { + connections: Mutex::new(HashMap::new()), + statuses: Statuses::default(), + profile, + keypair, + next_generation: AtomicU64::new(1), + }; + network.sync_remotes(handle, remotes::all(&database).await?)?; + handle.manage(network); - for remote in remotes { - let (sender, receiver) = mpsc::channel(32); - senders.insert(remote.id.clone(), sender); - set(&statuses, &remote, ConnectionState::Connecting); - tauri::async_runtime::spawn(run( - handle.clone(), - statuses.clone(), - remote, - profile_receiver.clone(), - keypair.clone(), - receiver, - )); - } - - handle.manage(Network { - senders: Mutex::new(senders), - statuses, - profile: profile_sender, + let listener_handle = handle.clone(); + RemotesChanged::listen(handle, move |event| { + let network = listener_handle.state::(); + if let Err(error) = network.sync_remotes(&listener_handle, event.payload.remotes) { + eprintln!("failed to synchronize remote connections: {error}"); + } }); Ok(()) } @@ -100,16 +162,24 @@ async fn run( handle: AppHandle, statuses: Statuses, remote: Remote, + generation: u64, mut profiles: watch::Receiver, keypair: AppKeypair, mut outgoing: mpsc::Receiver, ) { loop { - changed(&handle, &statuses, &remote, ConnectionState::Connecting); + changed( + &handle, + &statuses, + &remote, + generation, + ConnectionState::Connecting, + ); if let Err(error) = connect( &handle, &statuses, &remote, + generation, &mut profiles, &keypair, &mut outgoing, @@ -118,7 +188,13 @@ async fn run( { eprintln!("remote {} disconnected: {error}", remote.id); } - changed(&handle, &statuses, &remote, ConnectionState::Disconnected); + changed( + &handle, + &statuses, + &remote, + generation, + ConnectionState::Disconnected, + ); tokio::time::sleep(Duration::from_secs(5)).await; } } @@ -127,6 +203,7 @@ async fn connect( handle: &AppHandle, statuses: &Statuses, remote: &Remote, + generation: u64, profiles: &mut watch::Receiver, keypair: &AppKeypair, outgoing: &mut mpsc::Receiver, @@ -157,7 +234,13 @@ async fn connect( if !matches!(recv(&mut reader).await?, ServerMessage::Registered) { return Err("server rejected registration".into()); } - changed(handle, statuses, remote, ConnectionState::Connected); + changed( + handle, + statuses, + remote, + generation, + ConnectionState::Connected, + ); loop { tokio::select! { @@ -204,33 +287,71 @@ pub fn list_statuses( Ok(statuses) } -fn changed(handle: &AppHandle, statuses: &Statuses, remote: &Remote, state: ConnectionState) { - set(statuses, remote, state); - if let Ok(statuses) = snapshot(statuses) { - let _ = NetworkStatusChanged { statuses }.emit(handle); +fn changed( + handle: &AppHandle, + statuses: &Statuses, + remote: &Remote, + generation: u64, + state: ConnectionState, +) { + if set(statuses, remote, generation, state) { + let _ = emit_statuses(handle, statuses); } } -fn set(statuses: &Statuses, remote: &Remote, state: ConnectionState) { +fn set_initial(statuses: &Statuses, remote: &Remote, generation: u64, state: ConnectionState) { if let Ok(mut statuses) = statuses.lock() { - statuses.insert( - remote.id.clone(), - ConnectionStatus { - remote_id: remote.id.clone(), - address: remote.address.clone(), - name: remote.name.clone(), - state, - }, - ); + statuses.insert(remote.id.clone(), (generation, status(remote, state))); } } +fn set(statuses: &Statuses, remote: &Remote, generation: u64, state: ConnectionState) -> bool { + let Ok(mut statuses) = statuses.lock() else { + return false; + }; + let Some((current_generation, current)) = statuses.get_mut(&remote.id) else { + return false; + }; + if *current_generation != generation { + return false; + } + *current = status(remote, state); + true +} + +fn remove_status(statuses: &Statuses, remote_id: &str, generation: u64) { + if let Ok(mut statuses) = statuses.lock() + && statuses + .get(remote_id) + .is_some_and(|(current, _)| *current == generation) + { + statuses.remove(remote_id); + } +} + +fn status(remote: &Remote, state: ConnectionState) -> ConnectionStatus { + ConnectionStatus { + remote_id: remote.id.clone(), + address: remote.address.clone(), + name: remote.name.clone(), + state, + } +} + +fn emit_statuses(handle: &AppHandle, statuses: &Statuses) -> Result<(), String> { + NetworkStatusChanged { + statuses: snapshot(statuses)?, + } + .emit(handle) + .map_err(|error| error.to_string()) +} + fn snapshot(statuses: &Statuses) -> Result, String> { let mut statuses: Vec<_> = statuses .lock() .map_err(|error| error.to_string())? .values() - .cloned() + .map(|(_, status)| status.clone()) .collect(); statuses.sort_by(|a, b| a.remote_id.cmp(&b.remote_id)); Ok(statuses)