diff --git a/src-common/src/lib.rs b/src-common/src/lib.rs index 4fd4cca..d20d1c2 100644 --- a/src-common/src/lib.rs +++ b/src-common/src/lib.rs @@ -12,9 +12,24 @@ pub struct Profile { #[derive(Debug, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "camelCase")] pub enum ClientMessage { - Register { profile: Profile, signature: String }, - ProfileUpdated { profile: Profile, signature: String }, - Signed { payload: String, signature: String }, + Register { + profile: Profile, + friends: Vec, + signature: String, + }, + ProfileUpdated { + profile: Profile, + signature: String, + }, + FriendsUpdated { + friends: Vec, + signature: String, + }, + SyncFriendProfiles, + Signed { + payload: String, + signature: String, + }, } #[derive(Debug, Serialize, Deserialize)] @@ -22,12 +37,16 @@ pub enum ClientMessage { pub enum ServerMessage { Challenge { version: u8, challenge: String }, Registered, + FriendProfileUpdated { profile: Profile }, + FriendProfiles { profiles: Vec }, } -pub fn register_bytes(challenge: &str, profile: &Profile) -> Vec { +pub fn register_bytes(challenge: &str, profile: &Profile, friends: &[String]) -> Vec { format!( - "wyd-register-v{VERSION}\n{challenge}\n{}\n{}", - profile.id, profile.display_name + "wyd-register-v{VERSION}\n{challenge}\n{}\n{}\n{}", + profile.id, + profile.display_name, + friends.join("\n") ) .into_bytes() } @@ -43,3 +62,7 @@ pub fn profile_bytes(profile: &Profile) -> Vec { ) .into_bytes() } + +pub fn friends_bytes(friends: &[String]) -> Vec { + format!("wyd-friends-v{VERSION}\n{}", friends.join("\n")).into_bytes() +} diff --git a/src-server/src/network/mod.rs b/src-server/src/network/mod.rs index 2e04288..f895fa3 100644 --- a/src-server/src/network/mod.rs +++ b/src-server/src/network/mod.rs @@ -14,7 +14,8 @@ use futures_util::{SinkExt, StreamExt}; use tokio::sync::{Mutex, mpsc}; use uuid::Uuid; use wyd_common::{ - ClientMessage, Profile, ServerMessage, message_bytes, profile_bytes, register_bytes, + ClientMessage, Profile, ServerMessage, friends_bytes, message_bytes, profile_bytes, + register_bytes, }; type Clients = Arc>>; @@ -24,6 +25,8 @@ struct Client { key: VerifyingKey, #[allow(dead_code)] // Used when presence and profile lookup are exposed. profile: Profile, + #[allow(dead_code)] // Used when friend-authorized message routing is added. + friends: Vec, #[allow(dead_code)] // Used when server-side message routing is added. sender: mpsc::Sender, } @@ -56,11 +59,20 @@ async fn connected(mut socket: WebSocket, clients: Clients) { let Some(Ok(Message::Text(text))) = socket.recv().await else { return; }; - let Ok(ClientMessage::Register { profile, signature }) = serde_json::from_str(&text) else { + let Ok(ClientMessage::Register { + profile, + friends, + signature, + }) = serde_json::from_str(&text) + else { return; }; let Ok(key) = key(&profile.id) else { return }; - if !verify(&key, ®ister_bytes(&challenge, &profile), &signature) { + if !verify( + &key, + ®ister_bytes(&challenge, &profile, &friends), + &signature, + ) { return; } @@ -73,6 +85,7 @@ async fn connected(mut socket: WebSocket, clients: Clients) { connection_id, key, profile, + friends, sender, }, ); @@ -107,6 +120,20 @@ async fn connected(mut socket: WebSocket, clients: Clients) { break; } } + Ok(ClientMessage::FriendsUpdated { friends, signature }) => { + if !update_friends(&clients, &public_key, connection_id, friends, &signature).await { + break; + } + } + Ok(ClientMessage::SyncFriendProfiles) => { + let Some(profiles) = friend_profiles(&clients, &public_key, connection_id).await else { + break; + }; + let message = ServerMessage::FriendProfiles { profiles }; + if writer.send(Message::Text(serde_json::to_string(&message).unwrap().into())).await.is_err() { + break; + } + } _ => break, } Some(Ok(Message::Ping(data))) => { @@ -127,6 +154,45 @@ async fn update_profile( connection_id: Uuid, profile: Profile, signature: &str, +) -> bool { + let mut clients = clients.lock().await; + { + let Some(client) = clients + .get_mut(public_key) + .filter(|client| client.connection_id == connection_id) + else { + return false; + }; + if profile.id != public_key || !verify(&client.key, &profile_bytes(&profile), signature) { + return false; + } + client.profile = profile.clone(); + } + + let recipients: Vec<_> = clients + .values() + .filter(|client| client.friends.iter().any(|friend| friend == public_key)) + .map(|client| client.sender.clone()) + .collect(); + drop(clients); + + let message = Message::Text( + serde_json::to_string(&ServerMessage::FriendProfileUpdated { profile }) + .unwrap() + .into(), + ); + for recipient in recipients { + let _ = recipient.send(message.clone()).await; + } + true +} + +async fn update_friends( + clients: &Clients, + public_key: &str, + connection_id: Uuid, + friends: Vec, + signature: &str, ) -> bool { let mut clients = clients.lock().await; let Some(client) = clients @@ -135,13 +201,31 @@ async fn update_profile( else { return false; }; - if profile.id != public_key || !verify(&client.key, &profile_bytes(&profile), signature) { + if !verify(&client.key, &friends_bytes(&friends), signature) { return false; } - client.profile = profile; + client.friends = friends; true } +async fn friend_profiles( + clients: &Clients, + public_key: &str, + connection_id: Uuid, +) -> Option> { + let clients = clients.lock().await; + let client = clients + .get(public_key) + .filter(|client| client.connection_id == connection_id)?; + let mut profiles: Vec<_> = client + .friends + .iter() + .filter_map(|friend| clients.get(friend).map(|client| client.profile.clone())) + .collect(); + profiles.sort_by(|a, b| a.id.cmp(&b.id)); + Some(profiles) +} + async fn remove(clients: &Clients, public_key: &str, connection_id: Uuid) { let mut clients = clients.lock().await; if clients @@ -189,7 +273,8 @@ mod tests { id: URL_SAFE_NO_PAD.encode(signing_key.verifying_key().to_bytes()), display_name: "Wind".to_owned(), }; - let registration = register_bytes("challenge", &profile); + let friends = vec!["friend".to_owned()]; + let registration = register_bytes("challenge", &profile, &friends); let signature = URL_SAFE_NO_PAD.encode(signing_key.sign(®istration).to_bytes()); assert!(verify( @@ -199,7 +284,12 @@ mod tests { )); assert!(!verify( &signing_key.verifying_key(), - ®ister_bytes("another challenge", &profile), + ®ister_bytes("another challenge", &profile, &friends), + &signature, + )); + assert!(!verify( + &signing_key.verifying_key(), + ®ister_bytes("challenge", &profile, &["another-friend".to_owned()]), &signature, )); @@ -219,19 +309,37 @@ mod tests { let public_key = URL_SAFE_NO_PAD.encode(signing_key.verifying_key().to_bytes()); let connection_id = Uuid::new_v4(); let (sender, _receiver) = mpsc::channel(1); + let (friend_sender, mut friend_receiver) = mpsc::channel(1); let clients = Clients::default(); - clients.lock().await.insert( - public_key.clone(), - Client { - connection_id, - key: signing_key.verifying_key(), - profile: Profile { - id: public_key.clone(), - display_name: "Old".to_owned(), + { + let mut clients = clients.lock().await; + clients.insert( + public_key.clone(), + Client { + connection_id, + key: signing_key.verifying_key(), + profile: Profile { + id: public_key.clone(), + display_name: "Old".to_owned(), + }, + friends: Vec::new(), + sender, }, - sender, - }, - ); + ); + clients.insert( + "friend".to_owned(), + Client { + connection_id: Uuid::new_v4(), + key: signing_key.verifying_key(), + profile: Profile { + id: "friend".to_owned(), + display_name: "Friend".to_owned(), + }, + friends: vec![public_key.clone()], + sender: friend_sender, + }, + ); + } let profile = Profile { id: public_key.clone(), display_name: "New".to_owned(), @@ -239,7 +347,49 @@ mod tests { let signature = URL_SAFE_NO_PAD.encode(signing_key.sign(&profile_bytes(&profile)).to_bytes()); - assert!(update_profile(&clients, &public_key, connection_id, profile, &signature).await); + assert!(update_profile(&clients, &public_key, connection_id, profile, &signature,).await); + assert_eq!( + clients + .lock() + .await + .get(&public_key) + .unwrap() + .profile + .display_name, + "New" + ); + let announcement = friend_receiver + .recv() + .await + .expect("friend receives update"); + let Message::Text(announcement) = announcement else { + panic!("expected text announcement"); + }; + let ServerMessage::FriendProfileUpdated { profile } = + serde_json::from_str(&announcement).expect("decode announcement") + else { + panic!("expected friend profile update"); + }; + assert_eq!(profile.id, public_key); + assert_eq!(profile.display_name, "New"); + + let stale_profile = Profile { + id: public_key.clone(), + display_name: "Stale".to_owned(), + }; + let stale_signature = + URL_SAFE_NO_PAD.encode(signing_key.sign(&profile_bytes(&stale_profile)).to_bytes()); + assert!( + !update_profile( + &clients, + &public_key, + Uuid::new_v4(), + stale_profile, + &stale_signature, + ) + .await + ); + assert!(friend_receiver.try_recv().is_err()); assert_eq!( clients .lock() @@ -251,4 +401,102 @@ mod tests { "New" ); } + + #[tokio::test] + async fn authenticated_friend_update_changes_only_ids() { + let signing_key = SigningKey::from_bytes(&[7; 32]); + let public_key = URL_SAFE_NO_PAD.encode(signing_key.verifying_key().to_bytes()); + let connection_id = Uuid::new_v4(); + let (sender, _receiver) = mpsc::channel(1); + let clients = Clients::default(); + clients.lock().await.insert( + public_key.clone(), + Client { + connection_id, + key: signing_key.verifying_key(), + profile: Profile { + id: public_key.clone(), + display_name: "Wind".to_owned(), + }, + friends: Vec::new(), + sender, + }, + ); + let friends = vec!["friend-a".to_owned(), "friend-b".to_owned()]; + let signature = + URL_SAFE_NO_PAD.encode(signing_key.sign(&friends_bytes(&friends)).to_bytes()); + + assert!( + update_friends( + &clients, + &public_key, + connection_id, + friends.clone(), + &signature, + ) + .await + ); + assert_eq!( + clients.lock().await.get(&public_key).unwrap().friends, + friends + ); + } + + #[tokio::test] + async fn profile_sync_returns_only_connected_friends_for_the_current_session() { + let signing_key = SigningKey::from_bytes(&[7; 32]); + let connection_id = Uuid::new_v4(); + let (sender, _receiver) = mpsc::channel(1); + let clients = Clients::default(); + let mut registry = clients.lock().await; + registry.insert( + "requester".to_owned(), + Client { + connection_id, + key: signing_key.verifying_key(), + profile: Profile { + id: "requester".to_owned(), + display_name: "Requester".to_owned(), + }, + friends: vec![ + "friend-b".to_owned(), + "offline".to_owned(), + "friend-a".to_owned(), + ], + sender: sender.clone(), + }, + ); + for (id, display_name) in [("friend-a", "Alice"), ("friend-b", "Bob")] { + registry.insert( + id.to_owned(), + Client { + connection_id: Uuid::new_v4(), + key: signing_key.verifying_key(), + profile: Profile { + id: id.to_owned(), + display_name: display_name.to_owned(), + }, + friends: Vec::new(), + sender: sender.clone(), + }, + ); + } + drop(registry); + + let profiles = friend_profiles(&clients, "requester", connection_id) + .await + .expect("current session"); + assert_eq!( + profiles + .iter() + .map(|profile| (profile.id.as_str(), profile.display_name.as_str())) + .collect::>(), + [("friend-a", "Alice"), ("friend-b", "Bob")] + ); + assert!( + friend_profiles(&clients, "requester", Uuid::new_v4()) + .await + .is_none() + ); + } } diff --git a/src-tauri/src/friends/mod.rs b/src-tauri/src/friends/mod.rs index 73a403c..00ad284 100644 --- a/src-tauri/src/friends/mod.rs +++ b/src-tauri/src/friends/mod.rs @@ -11,7 +11,7 @@ pub struct FriendsChanged { pub friends: Vec, } -async fn all(database: &AppDatabase) -> Result, sqlx::Error> { +pub(crate) async fn all(database: &AppDatabase) -> Result, sqlx::Error> { sqlx::query_as::<_, User>( "SELECT id, display_name FROM friends ORDER BY display_name COLLATE NOCASE, id", ) @@ -27,6 +27,48 @@ async fn emit_changed(handle: &AppHandle, database: &AppDatabase) -> Result<(), .map_err(db::command_error) } +async fn update_display_names( + database: &AppDatabase, + profiles: &[wyd_common::Profile], +) -> Result { + let mut transaction = database.pool().begin().await?; + let mut changed = false; + for profile in profiles { + let result = sqlx::query( + "UPDATE friends SET display_name = ?1 WHERE id = ?2 AND display_name != ?1", + ) + .bind(&profile.display_name) + .bind(&profile.id) + .execute(&mut *transaction) + .await?; + changed |= result.rows_affected() > 0; + } + transaction.commit().await?; + Ok(changed) +} + +pub(crate) async fn apply_profile_update( + handle: &AppHandle, + database: &AppDatabase, + profile: wyd_common::Profile, +) -> Result { + apply_profile_sync(handle, database, vec![profile]).await +} + +pub(crate) async fn apply_profile_sync( + handle: &AppHandle, + database: &AppDatabase, + profiles: Vec, +) -> Result { + let changed = update_display_names(database, &profiles) + .await + .map_err(db::command_error)?; + if changed { + emit_changed(handle, database).await?; + } + Ok(changed) +} + #[tauri::command] #[specta::specta] pub async fn create_friend( @@ -97,3 +139,70 @@ pub async fn delete_friend( Ok(changed) } + +#[cfg(test)] +mod tests { + use sqlx::sqlite::SqlitePoolOptions; + + use super::*; + + async fn database() -> AppDatabase { + let pool = SqlitePoolOptions::new() + .max_connections(1) + .connect("sqlite::memory:") + .await + .expect("connect to in-memory SQLite"); + sqlx::migrate!("./migrations") + .run(&pool) + .await + .expect("run database migrations"); + AppDatabase::new(pool) + } + + #[tokio::test] + async fn profile_update_changes_only_an_existing_friend() { + let database = database().await; + sqlx::query("INSERT INTO friends (id, display_name) VALUES (?1, ?2)") + .bind("friend-id") + .bind("Old") + .execute(database.pool()) + .await + .unwrap(); + + assert!( + update_display_names( + &database, + &[ + wyd_common::Profile { + id: "friend-id".to_owned(), + display_name: "New".to_owned(), + }, + wyd_common::Profile { + id: "missing".to_owned(), + display_name: "Name".to_owned(), + }, + ], + ) + .await + .unwrap() + ); + assert!( + !update_display_names( + &database, + &[wyd_common::Profile { + id: "friend-id".to_owned(), + display_name: "New".to_owned(), + }], + ) + .await + .unwrap() + ); + assert_eq!( + all(&database).await.unwrap(), + [User { + id: "friend-id".to_owned(), + display_name: "New".to_owned(), + }] + ); + } +} diff --git a/src-tauri/src/network/mod.rs b/src-tauri/src/network/mod.rs index 652a4d7..1e8c3a8 100644 --- a/src-tauri/src/network/mod.rs +++ b/src-tauri/src/network/mod.rs @@ -11,10 +11,12 @@ use tauri_specta::Event; use tokio::sync::{mpsc, watch}; use tokio_tungstenite::tungstenite::Message; use wyd_common::{ - ClientMessage, Profile, ServerMessage, message_bytes, profile_bytes, register_bytes, + ClientMessage, Profile, ServerMessage, friends_bytes, message_bytes, profile_bytes, + register_bytes, }; use crate::db::AppDatabase; +use crate::friends::{self, FriendsChanged}; use crate::keypair::AppKeypair; use crate::remotes::{self, Remote, RemotesChanged}; @@ -54,6 +56,7 @@ pub struct Network { connections: Mutex>, statuses: Statuses, profile: watch::Sender, + friends: watch::Sender>, keypair: AppKeypair, next_generation: AtomicU64, } @@ -114,6 +117,7 @@ impl Network { remote.clone(), generation, self.profile.subscribe(), + self.friends.subscribe(), self.keypair.clone(), receiver, )); @@ -138,10 +142,13 @@ pub async fn init(handle: &AppHandle) -> Result<(), Box> let keypair = handle.state::().inner().clone(); let profile = crate::profile::get(&database, keypair.public_key()).await?; let (profile, _) = watch::channel(profile); + let friends = friend_ids(friends::all(&database).await?, keypair.public_key()); + let (friends, _) = watch::channel(friends); let network = Network { connections: Mutex::new(HashMap::new()), statuses: Statuses::default(), profile, + friends, keypair, next_generation: AtomicU64::new(1), }; @@ -155,6 +162,19 @@ pub async fn init(handle: &AppHandle) -> Result<(), Box> eprintln!("failed to synchronize remote connections: {error}"); } }); + + let listener_handle = handle.clone(); + FriendsChanged::listen(handle, move |event| { + let network = listener_handle.state::(); + let ids = friend_ids(event.payload.friends, network.keypair.public_key()); + network.friends.send_if_modified(|current| { + if *current == ids { + return false; + } + *current = ids; + true + }); + }); Ok(()) } @@ -164,6 +184,7 @@ async fn run( remote: Remote, generation: u64, mut profiles: watch::Receiver, + mut friends: watch::Receiver>, keypair: AppKeypair, mut outgoing: mpsc::Receiver, ) { @@ -181,6 +202,7 @@ async fn run( &remote, generation, &mut profiles, + &mut friends, &keypair, &mut outgoing, ) @@ -205,6 +227,7 @@ async fn connect( remote: &Remote, generation: u64, profiles: &mut watch::Receiver, + friends: &mut watch::Receiver>, keypair: &AppKeypair, outgoing: &mut mpsc::Receiver, ) -> Result<(), Box> { @@ -222,11 +245,17 @@ async fn connect( id: current.id, display_name: current.display_name, }; + let registration_friends = friends.borrow_and_update().clone(); send( &mut writer, &ClientMessage::Register { - signature: keypair.sign(®ister_bytes(&challenge, ®istration_profile)), + signature: keypair.sign(®ister_bytes( + &challenge, + ®istration_profile, + ®istration_friends, + )), profile: registration_profile, + friends: registration_friends, }, ) .await?; @@ -234,6 +263,7 @@ async fn connect( if !matches!(recv(&mut reader).await?, ServerMessage::Registered) { return Err("server rejected registration".into()); } + send(&mut writer, &ClientMessage::SyncFriendProfiles).await?; changed( handle, statuses, @@ -263,7 +293,30 @@ async fn connect( profile, }).await?; } + changed = friends.changed() => { + changed.map_err(|_| "friends sender closed")?; + let friends = friends.borrow_and_update().clone(); + send(&mut writer, &ClientMessage::FriendsUpdated { + signature: keypair.sign(&friends_bytes(&friends)), + friends, + }).await?; + } message = reader.next() => match message.ok_or("server closed the socket")?? { + Message::Text(text) => match serde_json::from_str(&text)? { + ServerMessage::FriendProfileUpdated { profile } => { + let database = handle.state::(); + if let Err(error) = friends::apply_profile_update(handle, &database, profile).await { + eprintln!("failed to update friend profile: {error}"); + } + } + ServerMessage::FriendProfiles { profiles } => { + let database = handle.state::(); + if let Err(error) = friends::apply_profile_sync(handle, &database, profiles).await { + eprintln!("failed to synchronize friend profiles: {error}"); + } + } + _ => {} + }, Message::Ping(data) => writer.send(Message::Pong(data)).await?, Message::Close(_) => return Ok(()), _ => {} @@ -402,3 +455,40 @@ fn url(remote: &Remote) -> String { .unwrap_or_default(); format!("{scheme}://{address}{port}/v1/ws") } + +fn friend_ids(friends: Vec, own_id: &str) -> Vec { + let mut ids: Vec<_> = friends + .into_iter() + .map(|friend| friend.id) + .filter(|id| id != own_id) + .collect(); + ids.sort_unstable(); + ids.dedup(); + ids +} + +#[cfg(test)] +mod tests { + use super::friend_ids; + use crate::user::User; + + #[test] + fn friend_ids_discards_display_names_and_normalizes_ids() { + let friends = vec![ + User { + id: "friend-b".to_owned(), + display_name: "Old name".to_owned(), + }, + User { + id: "self".to_owned(), + display_name: "Me".to_owned(), + }, + User { + id: "friend-a".to_owned(), + display_name: "Any name".to_owned(), + }, + ]; + + assert_eq!(friend_ids(friends, "self"), ["friend-a", "friend-b"]); + } +}