diff --git a/.gitignore b/.gitignore index 76add87..5deb29a 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,3 @@ node_modules -dist \ No newline at end of file +dist +.idea/ \ No newline at end of file diff --git a/server/src/main.rs b/server/src/main.rs index edfcd83..4a9d924 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -45,7 +45,7 @@ async fn main() { get(get_key_package_identities).post(create_key_package), ) .route("/packages/:identity", get(get_key_package)) - .route("/ws", get(websocket_handler)) + .route("/:identity/messages", get(websocket_handler)) .layer( CorsLayer::new() .allow_origin("*".parse::().unwrap()) @@ -106,14 +106,16 @@ async fn get_key_package( } async fn websocket_handler( + Path(identity): Path, websocket: WebSocketUpgrade, state: State, ) -> impl IntoResponse { - websocket.on_upgrade(move |socket| create_actor(socket, state)) + websocket.on_upgrade(move |socket| create_actor(socket, state, identity)) } // 2/3e, duck2duck encryption, melt -async fn create_actor(stream: WebSocket, State(state): State) { - let actor = UserActorHandle::new(stream, state.user_actors.clone()); - state.user_actors.lock().await.push(actor); +async fn create_actor(stream: WebSocket, State(state): State, identity: String) { + let actor = UserActorHandle::new(identity, stream, state.user_actors.clone()); + let mut actors_guild = state.user_actors.lock().await; + actors_guild.push(actor); } diff --git a/server/src/user_actor.rs b/server/src/user_actor.rs index 9ab6ab9..a4421ee 100644 --- a/server/src/user_actor.rs +++ b/server/src/user_actor.rs @@ -11,6 +11,7 @@ use tokio::sync::{ }; struct UserActor { + id: String, receiver: mpsc::Receiver, websocket: WebSocket, other_actors: Arc>>, @@ -27,11 +28,13 @@ enum Instruction { impl UserActor { fn new( + id: String, websocket: WebSocket, receiver: mpsc::Receiver, other_actors: Arc>>, ) -> Self { UserActor { + id, receiver, websocket, other_actors, @@ -52,6 +55,9 @@ impl UserActor { let mut others = self.other_actors.lock().await; let mut dead_actors = Vec::new(); for (index, other) in others.iter().enumerate() { + if other.id == self.id { + continue; + } //TODO use shared reference instead to avoid cloning of possibly large messages let result = other.send_message(binary.clone()).await; // Erros when channel is closed @@ -87,6 +93,7 @@ impl UserActor { async fn run_my_actor(mut actor: UserActor) { tracing::debug!("Actor started"); + loop { tokio::select! { Some(message) = actor.receiver.recv() => { @@ -110,18 +117,22 @@ async fn run_my_actor(mut actor: UserActor) { } pub(crate) struct UserActorHandle { + /// A simple unique identifier for the actor to distinguish it from others + /// The actor and actor handle share the same id + id: String, sender: mpsc::Sender, } impl UserActorHandle { pub(crate) fn new( + id: String, websocket: WebSocket, other_actors: Arc>>, ) -> Self { let (sender, receiver) = mpsc::channel(8); - let actor = UserActor::new(websocket, receiver, other_actors); + let actor = UserActor::new(id.clone(), websocket, receiver, other_actors); tokio::spawn(run_my_actor(actor)); - Self { sender } + Self { id, sender } } /// Errors when actor stopped receiving messages meaning the channel is closed and the actor is deceased diff --git a/src-tauri/src/command/create_user.rs b/src-tauri/src/command/create_user.rs new file mode 100644 index 0000000..842f2d1 --- /dev/null +++ b/src-tauri/src/command/create_user.rs @@ -0,0 +1,59 @@ +use openmls::credentials::{Credential, CredentialType, CredentialWithKey}; +use openmls::prelude::CredentialError; +use openmls_basic_credential::SignatureKeyPair; +use openmls_traits::types::CryptoError; +use serde::Serialize; +use tauri::State; +use thiserror::Error; +use crate::{AdvertiseKeyPackageError, AppState, CIPHERSUITE, User}; + +#[derive(Error, Debug, Serialize)] +pub(crate) enum CreateUserError { + #[error("User already exists")] + UserExists, + #[error("Error creating credentials for user")] + CredentialsError( + #[from] + #[serde(skip)] + CredentialError, + ), + + #[error("Error creating signature key pair")] + SignatureKeyPairError( + #[from] + #[serde(skip)] + CryptoError, + ), + + #[error("Error advertising key package on server")] + AdvertiseKeyPackageError( + #[from] + #[serde(skip)] + AdvertiseKeyPackageError, + ), +} + +#[tauri::command] +pub(crate) async fn create_user(name: &str, state: State<'_, AppState>) -> Result<(), CreateUserError> { + let mut state = state.user.lock().await; + if state.is_some() { + return Err(CreateUserError::UserExists); + } + + let credential = Credential::new(name.into(), CredentialType::Basic)?; + let signature_key_pair = SignatureKeyPair::new(CIPHERSUITE.signature_algorithm())?; + + let credential = CredentialWithKey { + credential, + signature_key: signature_key_pair.public().into(), + }; + + let user = User { + credential, + signature_key: signature_key_pair, + }; + + state.replace(user); + + Ok(()) +} diff --git a/src-tauri/src/command/mod.rs b/src-tauri/src/command/mod.rs new file mode 100644 index 0000000..8a25578 --- /dev/null +++ b/src-tauri/src/command/mod.rs @@ -0,0 +1,2 @@ + mod create_user; + pub use create_user::*; \ No newline at end of file diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 2573d19..4b82a0a 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -1,3 +1,5 @@ +mod command; + use base64::prelude::*; use openmls::prelude::*; use openmls_basic_credential::SignatureKeyPair; @@ -12,8 +14,10 @@ use std::{ use tauri::{AppHandle, Manager, State}; use thiserror::Error; use tokio::sync::Mutex; + // Disable dead code warnings for this file #[allow(dead_code)] + pub(crate) const CIPHERSUITE: Ciphersuite = Ciphersuite::MLS_128_DHKEMX25519_AES128GCM_SHA256_Ed25519; @@ -29,57 +33,6 @@ struct AppState { client: Client, } -#[derive(Error, Debug, Serialize)] -enum CreateUserError { - #[error("User already exists")] - UserExists, - #[error("Error creating credentials for user")] - CredentialsError( - #[from] - #[serde(skip)] - CredentialError, - ), - - #[error("Error creating signature key pair")] - SignatureKeyPairError( - #[from] - #[serde(skip)] - CryptoError, - ), - - #[error("Error advertising key package on server")] - AdvertiseKeyPackageError( - #[from] - #[serde(skip)] - AdvertiseKeyPackageError, - ), -} - -#[tauri::command] -async fn create_user(name: &str, state: State<'_, AppState>) -> Result<(), CreateUserError> { - let mut state = state.user.lock().await; - if state.is_some() { - return Err(CreateUserError::UserExists); - } - - let credential = Credential::new(name.into(), CredentialType::Basic)?; - let signature_key_pair = SignatureKeyPair::new(CIPHERSUITE.signature_algorithm())?; - - let credential = CredentialWithKey { - credential, - signature_key: signature_key_pair.public().into(), - }; - - let user = User { - credential, - signature_key: signature_key_pair, - }; - - state.replace(user); - - Ok(()) -} - #[derive(Error, Debug, Serialize)] enum IsAuthenticatedError { #[error("Could not access state")] @@ -294,14 +247,14 @@ async fn invite_package( let backend = state.backend.crypto(); let package = package.validate(backend, ProtocolVersion::default())?; - let (mls_message_out, welcome_out, group_information) = + let (_mls_message_out, welcome_out, _group_information) = group.add_members(state.backend.as_ref(), &user.signature_key, &[package])?; //TODO check if commit needs to be synchronized with others // Merge pending commit that adds the new member group.merge_pending_commit(state.backend.as_ref())?; - // Return welcome message to frontent to send over websockets to other clients + // Return welcome message to frontend to send over websockets to other clients let data = welcome_out.tls_serialize_detached()?; Ok(data) @@ -327,6 +280,16 @@ enum ReceiveMessageError { #[serde(skip)] tauri::Error, ), + + //TODO Private message errors + #[error("Group not found")] + GroupNotFound, + #[error("Error processing message")] + ProcessMessageError( + #[from] + #[serde(skip)] + ProcessMessageError, + ), } const JOIN_GROUP_EVENT: &str = "join_group"; @@ -344,26 +307,52 @@ async fn process_message( ) -> Result<(), ReceiveMessageError> { let message = MlsMessageIn::tls_deserialize(&mut data.as_slice())?; - if let MlsMessageInBody::Welcome(welcome) = message.extract() { - // Create group from welcome message - let group = MlsGroup::new_from_welcome( - state.backend.as_ref(), - &MlsGroupConfig::default(), - welcome, - None, - )?; + let extract = message.extract(); + match extract { + MlsMessageInBody::Welcome(welcome) => { + // Create group from welcome message + let group = MlsGroup::new_from_welcome( + state.backend.as_ref(), + &MlsGroupConfig::default(), + welcome, + None, + )?; + + let id = group.group_id(); + let id = BASE64_URL_SAFE_NO_PAD.encode(id.as_slice()); - let id = group.group_id(); - let id = BASE64_URL_SAFE_NO_PAD.encode(id.as_slice()); + let mut groups = state.groups.lock().await; + groups.insert(id.clone(), group); - let mut groups = state.groups.lock().await; - groups.insert(id.clone(), group); + app.emit(JOIN_GROUP_EVENT, JoinGroupEvent { group_id: id })?; + Ok(()) + }, - app.emit(JOIN_GROUP_EVENT, JoinGroupEvent { group_id: id })?; - return Ok(()); + MlsMessageInBody::PrivateMessage(message) => { + // Process the message + let protocol_message : ProtocolMessage = message.into(); + + let id = protocol_message.group_id(); + let id = BASE64_URL_SAFE_NO_PAD.encode(id.as_slice()); + + let mut groups = state.groups.lock().await; + let Some(group) = groups.get_mut(&id) else { + return Err(ReceiveMessageError::GroupNotFound); + }; + + let processed_message = group.process_message(state.backend.as_ref(), protocol_message)?; + match processed_message.into_content() { + ProcessedMessageContent::ApplicationMessage(application_message) => { + //TODO send message to frontend + println!("Application message: {:?}", application_message); + }, + _ => unimplemented!("Message processed but that type is not implemented yet"), + } + Ok(()) + }, + _ => unimplemented!("Processing messages is not implemented yet"), } - unimplemented!("Processing messages is not implemented yet"); } #[tauri::command] @@ -393,7 +382,6 @@ async fn create_group(state: State<'_, AppState>) -> Result) -> Result, ()> { let groups = state.groups.lock().await; @@ -403,6 +391,70 @@ async fn get_groups(state: State<'_, AppState>) -> Result, ()> { Ok(ids) } +#[derive(Error, Debug, Serialize)] +enum GetIdentityError { + #[error("No user is signed in")] + NoUserError, +} + +#[tauri::command] +async fn get_identity(state: State<'_, AppState>) -> Result { + let user = state.user.lock().await; + let Some(user) = user.as_ref() else { + return Err(GetIdentityError::NoUserError); + }; + + let id = user.credential.credential.identity(); + let id = BASE64_URL_SAFE_NO_PAD.encode(id); + Ok(id) +} + +#[derive(Error, Debug, Serialize)] +enum CreateMessageError { + #[error("No user is signed in")] + NoUserError, + #[error("Group not found")] + GroupNotFound, + #[error("Error creating message")] + CreateMessageError( + #[from] + #[serde(skip)] + openmls::group::CreateMessageError, + ), + #[error("Error serializing message")] + SerializeMessageError( + #[from] + #[serde(skip)] + tls_codec::Error, + ), +} + +#[tauri::command] +async fn create_message( + state: State<'_, AppState>, + group_id: &str, + message: &str, +) -> Result, CreateMessageError> { + let user = state.user.lock().await; + let Some(user) = user.as_ref() else { + return Err(CreateMessageError::NoUserError); + }; + + let mut groups = state.groups.lock().await; + let Some(group) = groups.get_mut(group_id) else { + return Err(CreateMessageError::GroupNotFound); + }; + + let message = group.create_message( + state.backend.as_ref(), + &user.signature_key, + message.as_bytes(), + )?; + + let data = message.tls_serialize_detached()?; + Ok(data) +} + #[cfg_attr(mobile, tauri::mobile_entry_point)] pub fn run() { let client = Client::new(); @@ -427,10 +479,12 @@ pub fn run() { .plugin(tauri_plugin_shell::init()) .invoke_handler(tauri::generate_handler![ advertise, - is_authenticated, create_group, - create_user, + create_message, + command::create_user, + is_authenticated, get_groups, + get_identity, invite_package, process_message, ]) diff --git a/src/AppContext.tsx b/src/AppContext.tsx index a28ed02..f05ac9f 100644 --- a/src/AppContext.tsx +++ b/src/AppContext.tsx @@ -1,31 +1,77 @@ import { invoke } from "@tauri-apps/api/core"; import { listen } from "@tauri-apps/api/event"; import { + Accessor, JSX, Resource, Setter, createContext, + createEffect, createResource, + createSignal, onCleanup, useContext, } from "solid-js"; -const socket = new WebSocket("ws://localhost:3000/ws"); -socket.addEventListener("message", async (event) => { - if (typeof event.data === "string") - throw new Error("Unexpected string as message event payload"); - - await invoke("process_message", event.data); -}); const [groups, { mutate: setGroups }] = createResource( async () => (await invoke("get_groups")) as string[] ); + +const [identity, { mutate: setIdentity }] = createResource( + async () => + (await invoke("get_identity").catch((error) => { + console.warn( + "Could not get identity. But this might be expected if this is the first time the app is run", + error + ); + return undefined; + })) as string +); + +async function handleMessage(event: MessageEvent) { + // Might remove redundant check later + if (typeof event.data === "string") + throw new Error("Unexpected string as message event payload"); + + const data = event.data; + if (!(data instanceof Blob)) + throw new Error("Unexpected non-blob as message event payload"); + + //TODO find out how to pass the data as binary to tauri without going through serde + const buffer = await data.arrayBuffer(); + const array = [...new Uint8Array(buffer)]; + + await invoke("process_message", { data: array }); +} +const [socket, setSocket] = createSignal(); +createEffect((previous) => { + const id = identity(); + if (id === undefined) return; + + previous?.removeEventListener("message", handleMessage); + previous?.close(); + + const newSocket = new WebSocket(`ws://localhost:3000/${id}/messages`); + newSocket.addEventListener("message", handleMessage); + setSocket(newSocket); + return newSocket; +}, socket()); + type AppState = { - socket: WebSocket; + identity: Resource; + setIdentity: Setter; + socket: Accessor; groups: Resource; setGroups: Setter; }; -const state = { socket, groups, setGroups } satisfies AppState; + +const state = { + identity, + setIdentity, + socket, + groups, + setGroups, +} satisfies AppState; const AppContext = createContext(state); export function SocketProvider(properties: { children: JSX.Element }) { @@ -37,18 +83,19 @@ export function SocketProvider(properties: { children: JSX.Element }) { } export function useWebSocket(onmessage?: (event: MessageEvent) => any) { - const { socket: webSocket } = useContext(AppContext); + const { socket } = useContext(AppContext); - if (onmessage) { - webSocket.addEventListener("message", onmessage); + createEffect((previous) => { + const current = socket(); + if (onmessage === undefined || current === undefined) return current; - onCleanup(() => { - webSocket.removeEventListener("message", onmessage); - }); - } + previous?.removeEventListener("message", onmessage); + current.addEventListener("message", onmessage); + return current; + }, socket()); return (data: string | ArrayBufferLike | Blob | ArrayBufferView) => - webSocket.send(data); + socket()?.send(data); } export const useAppState = () => useContext(AppContext); diff --git a/src/index.tsx b/src/index.tsx index 01e062b..2a11b77 100644 --- a/src/index.tsx +++ b/src/index.tsx @@ -3,7 +3,7 @@ import { render } from "solid-js/web"; import { Route, Router } from "@solidjs/router"; import "./styles.css"; -import App from "./App"; +import Home from "./routes/Home"; import Group from "./routes/Group"; import { SocketProvider } from "./AppContext"; @@ -11,7 +11,7 @@ render( () => ( - + diff --git a/src/routes/Group.tsx b/src/routes/Group.tsx index 2f53417..3168d81 100644 --- a/src/routes/Group.tsx +++ b/src/routes/Group.tsx @@ -1,7 +1,7 @@ import { useParams } from "@solidjs/router"; import { invoke } from "@tauri-apps/api/core"; import { For, createResource } from "solid-js"; -import { useWebSocket } from "../AppContext"; +import { useAppState, useWebSocket } from "../AppContext"; async function getPackagesIndex() { const response = await fetch(`http://localhost:3000/packages`); @@ -17,6 +17,7 @@ export default function Group() { const groupId = () => parameters.id; const [packages] = createResource(getPackagesIndex); + const { identity } = useAppState(); async function invitePackage(id: string) { if (!groupId()) return; @@ -31,8 +32,26 @@ export default function Group() { const data = Uint8Array.from(message); sendMessage(data); } + + async function handleMessageSubmit(event: SubmitEvent) { + event.preventDefault(); + + // @ts-ignore + const message = event.target.message.value; + + const data = (await invoke("create_message", { + groupId: groupId(), + message, + })) as number[]; + + const buffer = Uint8Array.from(data); + + sendMessage(buffer); + } + return (
+

Your identity is {identity()}

Group {groupId()}

Packages to invite

    @@ -46,6 +65,13 @@ export default function Group() { )}
+ +

Messages

+
+ + + +
); } diff --git a/src/App.tsx b/src/routes/Home.tsx similarity index 87% rename from src/App.tsx rename to src/routes/Home.tsx index a9a4144..544c486 100644 --- a/src/App.tsx +++ b/src/routes/Home.tsx @@ -1,6 +1,6 @@ import { invoke } from "@tauri-apps/api/core"; import { For, Show, createResource, createSignal } from "solid-js"; -import { useAppState } from "./AppContext"; +import { useAppState } from "../AppContext"; async function createUser(name: string) { await invoke("create_user", { name }); @@ -11,11 +11,11 @@ const isAuthenticated = async () => const createGroup = async () => (await invoke("create_group")) as string; -function App() { +function Home() { const [isAuthenticatedResource, { refetch: refetchIsAuthenticated }] = createResource(isAuthenticated); - const { groups, setGroups } = useAppState(); + const { groups, setGroups, identity, setIdentity } = useAppState(); function handleSubmit(event: SubmitEvent) { event.preventDefault(); @@ -23,6 +23,7 @@ function App() { const name = event.target.name.value; createUser(name); refetchIsAuthenticated(); + setIdentity(name); } async function handleCreateGroup() { @@ -46,6 +47,7 @@ function App() { +

Your identity is {identity()}

    @@ -62,4 +64,4 @@ function App() { ); } -export default App; +export default Home;