diff --git a/server/Cargo.lock b/server/Cargo.lock index 55af365..9d0bf6f 100644 --- a/server/Cargo.lock +++ b/server/Cargo.lock @@ -17,6 +17,15 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f26201604c87b1e01bd3d98f8d5d9a8fcbb815e8cedb41ffccbeb4bf593a35fe" +[[package]] +name = "aho-corasick" +version = "1.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e60d3430d3a69478ad0993f19238d2df97c507009a52b3c10addcd7f6bcb916" +dependencies = [ + "memchr", +] + [[package]] name = "async-trait" version = "0.1.80" @@ -42,6 +51,7 @@ checksum = "3a6c9af12842a67734c9a2e355436e5d03b22383ed60cf13cd0c18fbfe3dcbcf" dependencies = [ "async-trait", "axum-core", + "base64 0.21.7", "bytes", "futures-util", "http", @@ -60,8 +70,10 @@ dependencies = [ "serde_json", "serde_path_to_error", "serde_urlencoded", + "sha1", "sync_wrapper 1.0.1", "tokio", + "tokio-tungstenite", "tower", "tower-layer", "tower-service", @@ -104,6 +116,12 @@ dependencies = [ "rustc-demangle", ] +[[package]] +name = "base64" +version = "0.21.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d297deb1925b89f2ccc13d7635fa0714f12c87adce1c75356b39ca9b7178567" + [[package]] name = "base64" version = "0.22.1" @@ -116,6 +134,21 @@ version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf4b9d6a944f767f8e5e0db018570623c85f3d925ac718db4e06d0187adb21c1" +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + +[[package]] +name = "byteorder" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" + [[package]] name = "bytes" version = "1.6.0" @@ -134,6 +167,15 @@ version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "baf1de4339761588bc0619e3cbc0120ee582ebb74b53b4efbf79117bd2da40fd" +[[package]] +name = "cpufeatures" +version = "0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53fe5e26ff1b7aef8bca9c6080520cfb8d9333c7568e1829cef191a9723e5504" +dependencies = [ + "libc", +] + [[package]] name = "crossbeam-deque" version = "0.8.5" @@ -159,6 +201,32 @@ version = "0.8.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "22ec99545bb0ed0ea7bb9b8e1e9122ea386ff8a48c0922e43f36d45ab09e0e80" +[[package]] +name = "crypto-common" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1bfb12502f3fc46cca1bb51ac28df9d618d813cdc3d2f25b9fe775a34af26bb3" +dependencies = [ + "generic-array", + "typenum", +] + +[[package]] +name = "data-encoding" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e962a19be5cfc3f3bf6dd8f61eb50107f356ad6270fbb3ed41476571db78be5" + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", +] + [[package]] name = "either" version = "1.12.0" @@ -195,6 +263,12 @@ version = "0.3.30" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dfc6580bb841c5a68e9ef15c77ccc837b40a7504914d52e47b8b0e9bbda25a1d" +[[package]] +name = "futures-sink" +version = "0.3.30" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9fb8e00e87438d937621c1c6269e53f536c14d3fbd6a042bb24879e57d474fb5" + [[package]] name = "futures-task" version = "0.3.30" @@ -208,9 +282,32 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3d6401deb83407ab3da39eba7e33987a73c3df0c82b4bb5813ee871c19c41d48" dependencies = [ "futures-core", + "futures-sink", "futures-task", "pin-project-lite", "pin-utils", + "slab", +] + +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + +[[package]] +name = "getrandom" +version = "0.2.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4567c8db10ae91089c99af84c68c38da3ec2f087c3f82960bcdbf3656b6f4d7" +dependencies = [ + "cfg-if", + "libc", + "wasi", ] [[package]] @@ -305,12 +402,28 @@ dependencies = [ "tokio", ] +[[package]] +name = "idna" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "634d9b1461af396cad843f47fdba5597a4f9e6ddd4bfb6ff5d85028c25cb12f6" +dependencies = [ + "unicode-bidi", + "unicode-normalization", +] + [[package]] name = "itoa" version = "1.0.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "49f1f14873335454500d59611f1cf4a4b0f786f9ac11f4312a78e4cf2566695b" +[[package]] +name = "lazy_static" +version = "1.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2abad23fbc42b3700f2f279844dc832adb2b2eb069b2df918f455c4e18cc646" + [[package]] name = "libc" version = "0.2.155" @@ -333,6 +446,15 @@ version = "0.4.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "90ed8c1e510134f979dbc4f070f87d4313098b704861a105fe34231c70a3901c" +[[package]] +name = "matchers" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8263075bb86c5a1b1427b5ae862e8889656f126e9f77c484496e8b47cf5c5558" +dependencies = [ + "regex-automata 0.1.10", +] + [[package]] name = "matchit" version = "0.7.3" @@ -371,6 +493,16 @@ dependencies = [ "windows-sys 0.48.0", ] +[[package]] +name = "nu-ansi-term" +version = "0.46.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77a8165726e8236064dbb45459242600304b42a5ea24ee2948e18e023bf7ba84" +dependencies = [ + "overload", + "winapi", +] + [[package]] name = "num_cpus" version = "1.16.0" @@ -421,6 +553,12 @@ dependencies = [ "tls_codec", ] +[[package]] +name = "overload" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b15813163c1d831bf4a13c3610c05c0d03b39feb07f7e09fa234dac9b15aaf39" + [[package]] name = "parking_lot" version = "0.12.3" @@ -482,6 +620,12 @@ version = "0.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184" +[[package]] +name = "ppv-lite86" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5b40af805b3121feab8a3c29f04d8ad262fa8e0561883e7653e024ae4479e6de" + [[package]] name = "proc-macro2" version = "1.0.84" @@ -500,6 +644,36 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "rand" +version = "0.8.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "34af8d1a0e25924bc5b7c43c079c942339d8f0a8b57c39049bef581b46327404" +dependencies = [ + "libc", + "rand_chacha", + "rand_core", +] + +[[package]] +name = "rand_chacha" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88" +dependencies = [ + "ppv-lite86", + "rand_core", +] + +[[package]] +name = "rand_core" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c" +dependencies = [ + "getrandom", +] + [[package]] name = "rayon" version = "1.10.0" @@ -529,6 +703,50 @@ dependencies = [ "bitflags", ] +[[package]] +name = "regex" +version = "1.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c117dbdfde9c8308975b6a18d71f3f385c89461f7b3fb054288ecf2a2058ba4c" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata 0.4.6", + "regex-syntax 0.8.3", +] + +[[package]] +name = "regex-automata" +version = "0.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c230d73fb8d8c1b9c0b3135c5142a8acee3a0558fb8db5cf1cb65f8d7862132" +dependencies = [ + "regex-syntax 0.6.29", +] + +[[package]] +name = "regex-automata" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "86b83b8b9847f9bf95ef68afb0b8e6cdb80f498442f5179a29fad448fcc1eaea" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax 0.8.3", +] + +[[package]] +name = "regex-syntax" +version = "0.6.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f162c6dd7b008981e4d40210aca20b4bd0f9b60ca9271061b07f78537722f2e1" + +[[package]] +name = "regex-syntax" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "adad44e29e4c806119491a7f06f03de4d1af22c3a680dd47f1e6e179439d1f56" + [[package]] name = "rustc-demangle" version = "0.1.24" @@ -611,10 +829,35 @@ name = "server" version = "0.1.0" dependencies = [ "axum", - "base64", + "base64 0.22.1", "openmls", + "serde", + "serde_json", "tokio", "tower", + "tower-http", + "tracing", + "tracing-subscriber", +] + +[[package]] +name = "sha1" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +dependencies = [ + "cfg-if", + "cpufeatures", + "digest", +] + +[[package]] +name = "sharded-slab" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6" +dependencies = [ + "lazy_static", ] [[package]] @@ -626,6 +869,15 @@ dependencies = [ "libc", ] +[[package]] +name = "slab" +version = "0.4.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f92a496fb766b417c996b9c5e57daf2f7ad3b0bebe1ccfca4856390e3d3bb67" +dependencies = [ + "autocfg", +] + [[package]] name = "smallvec" version = "1.13.2" @@ -685,6 +937,31 @@ dependencies = [ "syn", ] +[[package]] +name = "thread_local" +version = "1.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b9ef9bad013ada3808854ceac7b46812a6465ba368859a37e2100283d2d719c" +dependencies = [ + "cfg-if", + "once_cell", +] + +[[package]] +name = "tinyvec" +version = "1.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "87cc5ceb3875bb20c2890005a4e226a4651264a5c75edb2421b52861a0a0cb50" +dependencies = [ + "tinyvec_macros", +] + +[[package]] +name = "tinyvec_macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" + [[package]] name = "tls_codec" version = "0.3.0" @@ -737,6 +1014,18 @@ dependencies = [ "syn", ] +[[package]] +name = "tokio-tungstenite" +version = "0.21.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c83b561d025642014097b66e6c1bb422783339e0909e4429cde4749d1990bc38" +dependencies = [ + "futures-util", + "log", + "tokio", + "tungstenite", +] + [[package]] name = "tower" version = "0.4.13" @@ -753,6 +1042,22 @@ dependencies = [ "tracing", ] +[[package]] +name = "tower-http" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e9cd434a998747dd2c4276bc96ee2e0c7a2eadf3cae88e52be55a05fa9053f5" +dependencies = [ + "bitflags", + "bytes", + "http", + "http-body", + "http-body-util", + "pin-project-lite", + "tower-layer", + "tower-service", +] + [[package]] name = "tower-layer" version = "0.3.2" @@ -773,9 +1078,21 @@ checksum = "c3523ab5a71916ccf420eebdf5521fcef02141234bbc0b8a49f2fdc4544364ef" dependencies = [ "log", "pin-project-lite", + "tracing-attributes", "tracing-core", ] +[[package]] +name = "tracing-attributes" +version = "0.1.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "34704c8d6ebcbc939824180af020566b01a7c01f80641264eba0999f6c2b6be7" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "tracing-core" version = "0.1.32" @@ -783,20 +1100,141 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c06d3da6113f116aaee68e4d601191614c9053067f9ab7f6edbcb161237daa54" dependencies = [ "once_cell", + "valuable", +] + +[[package]] +name = "tracing-log" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee855f1f400bd0e5c02d150ae5de3840039a3f54b025156404e34c23c03f47c3" +dependencies = [ + "log", + "once_cell", + "tracing-core", +] + +[[package]] +name = "tracing-subscriber" +version = "0.3.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ad0f048c97dbd9faa9b7df56362b8ebcaa52adb06b498c050d2f4e32f90a7a8b" +dependencies = [ + "matchers", + "nu-ansi-term", + "once_cell", + "regex", + "sharded-slab", + "smallvec", + "thread_local", + "tracing", + "tracing-core", + "tracing-log", +] + +[[package]] +name = "tungstenite" +version = "0.21.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ef1a641ea34f399a848dea702823bbecfb4c486f911735368f1f137cb8257e1" +dependencies = [ + "byteorder", + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand", + "sha1", + "thiserror", + "url", + "utf-8", ] +[[package]] +name = "typenum" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42ff0bf0c66b8238c6f3b578df37d0b7848e55df8577b3f74f92a69acceeb825" + +[[package]] +name = "unicode-bidi" +version = "0.3.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08f95100a766bf4f8f28f90d77e0a5461bbdb219042e7679bebe79004fed8d75" + [[package]] name = "unicode-ident" version = "1.0.12" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3354b9ac3fae1ff6755cb6db53683adb661634f67557942dea4facebec0fee4b" +[[package]] +name = "unicode-normalization" +version = "0.1.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a56d1686db2308d901306f92a263857ef59ea39678a5458e7cb17f01415101f5" +dependencies = [ + "tinyvec", +] + +[[package]] +name = "url" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "31e6302e3bb753d46e83516cae55ae196fc0c309407cf11ab35cc51a4c2a4633" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", +] + +[[package]] +name = "utf-8" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" + +[[package]] +name = "valuable" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "830b7e5d4d90034032940e4ace0d9a9a057e7a45cd94e6c007832e39edb82f6d" + +[[package]] +name = "version_check" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "49874b5167b65d7193b8aba1567f5c7d93d001cafc34600cee003eda787e483f" + [[package]] name = "wasi" version = "0.11.0+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9c8d87e72b64a3b4db28d11ce29237c246188f4f51057d65a7eab63b7987e423" +[[package]] +name = "winapi" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419" +dependencies = [ + "winapi-i686-pc-windows-gnu", + "winapi-x86_64-pc-windows-gnu", +] + +[[package]] +name = "winapi-i686-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6" + +[[package]] +name = "winapi-x86_64-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" + [[package]] name = "windows-sys" version = "0.48.0" diff --git a/server/Cargo.toml b/server/Cargo.toml index f600aa9..975bfa7 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -4,8 +4,13 @@ version = "0.1.0" edition = "2021" [dependencies] -axum = "0.7.5" +axum = { version = "0.7.5", features = ["ws"] } base64 = "0.22.1" openmls = "0.5.0" +serde = "1.0.203" +serde_json = "1.0.117" tokio = { version = "1.37.0", features = ["full"] } tower = "0.4.13" +tower-http = { version = "0.5.2", features = ["cors"] } +tracing = "0.1.40" +tracing-subscriber = { version = "0.3.18", features = ["env-filter"] } diff --git a/server/src/key_package.rs b/server/src/key_package.rs index fa3554a..60f7d07 100644 --- a/server/src/key_package.rs +++ b/server/src/key_package.rs @@ -1,14 +1,9 @@ use axum::{ async_trait, - body::{Body, Bytes}, + body::Bytes, extract::{FromRequest, Request}, - http::{ - header::{HeaderValue, USER_AGENT}, - StatusCode, - }, + http::StatusCode, response::{IntoResponse, Response}, - routing::get, - Router, }; use openmls::prelude::*; pub(crate) struct KeyPackage(pub(crate) KeyPackageIn); @@ -22,10 +17,7 @@ where type Rejection = Response; async fn from_request(request: Request, state: &S) -> Result { - // let bytes = Bytes::from_request(req, state) - // .await - // .map_err(IntoResponse::into_response)?; - + //TODO stream bytes directly into KeyPackageIn::tls_deserialize without buffering everything first let bytes = Bytes::from_request(request, state) .await .map_err(IntoResponse::into_response)?; @@ -34,7 +26,7 @@ where let package = openmls::key_packages::KeyPackageIn::tls_deserialize(&mut bytes); match package { Ok(package) => Ok(KeyPackage(package)), - //TODO log error? + //TODO log error if it doesn't contain PII Err(_) => Err(StatusCode::BAD_REQUEST.into_response()), } } diff --git a/server/src/main.rs b/server/src/main.rs index e1c1258..edfcd83 100644 --- a/server/src/main.rs +++ b/server/src/main.rs @@ -1,24 +1,119 @@ -use axum::{response::Html, routing::get, Router}; +use std::{collections::HashMap, sync::Arc}; + +use axum::{ + body::Body, + extract::{ws::WebSocket, Path, State, WebSocketUpgrade}, + http::{HeaderValue, Method, StatusCode}, + response::{IntoResponse, Response}, + routing::get, + Json, Router, +}; +use base64::prelude::*; use key_package::KeyPackage; +use mls_message::MlsMessage; +use openmls::prelude::*; +use openmls::{key_packages::KeyPackageIn, prelude::TlsSerializeTrait}; +use tokio::sync::Mutex; +use tower_http::cors::CorsLayer; +use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt}; +use user_actor::UserActorHandle; + mod key_package; +mod mls_message; +mod user_actor; +mod websocket_actor; + +#[derive(Clone)] +struct AppState { + key_packages_by_identity: Arc>>, + user_actors: Arc>>, +} #[tokio::main] async fn main() { - let app = Router::new().route("/packages", get(get_key_packages).post(create_key_package)); + tracing_subscriber::registry() + .with( + tracing_subscriber::EnvFilter::try_from_default_env() + .unwrap_or_else(|_| "server=debug,tower_http=debug".into()), + ) + .with(tracing_subscriber::fmt::layer()) + .init(); + + let app = Router::new() + .route( + "/packages", + get(get_key_package_identities).post(create_key_package), + ) + .route("/packages/:identity", get(get_key_package)) + .route("/ws", get(websocket_handler)) + .layer( + CorsLayer::new() + .allow_origin("*".parse::().unwrap()) + // .allow_origin("localhost:1420".parse::().unwrap()) + // .allow_origin("localhost:1421".parse::().unwrap()) + .allow_methods([Method::GET]), + ) + .with_state(AppState { + key_packages_by_identity: Default::default(), + user_actors: Default::default(), + }); let listener = tokio::net::TcpListener::bind("127.0.0.1:3000") .await .unwrap(); - println!("listening on {}", listener.local_addr().unwrap()); + tracing::debug!("listening on {}", listener.local_addr().unwrap()); axum::serve(listener, app).await.unwrap(); } -async fn create_key_package(KeyPackage(package): KeyPackage) -> Result<(), ()> { - print!("Received key package {:?}", package); +async fn create_key_package( + State(state): State, + KeyPackage(package): KeyPackage, +) -> Result<(), StatusCode> { + tracing::debug!("Received key package"); + let mut key_packages = state.key_packages_by_identity.lock().await; + + let credential = package.unverified_credential().credential; + let identity = credential.identity(); + + let Ok(identity) = std::str::from_utf8(identity) else { + return Err(StatusCode::BAD_REQUEST); + }; + + key_packages.insert(identity.to_string(), package); Ok(()) } -async fn get_key_packages() -> Result<(), ()> { - Ok(()) +async fn get_key_package_identities(State(state): State) -> impl IntoResponse { + let key_packages = state.key_packages_by_identity.lock().await; + let identities = key_packages.keys().cloned().collect::>(); + + Json(identities) +} + +async fn get_key_package( + State(state): State, + Path(identity): Path, +) -> Result, StatusCode> { + let packages = state.key_packages_by_identity.lock().await; + let Some(package) = packages.get(&identity) else { + return Err(StatusCode::NOT_FOUND); + }; + + package + .tls_serialize_detached() + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR) +} + +async fn websocket_handler( + websocket: WebSocketUpgrade, + state: State, +) -> impl IntoResponse { + websocket.on_upgrade(move |socket| create_actor(socket, state)) +} + +// 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); } diff --git a/server/src/mls_message.rs b/server/src/mls_message.rs new file mode 100644 index 0000000..d2a95ab --- /dev/null +++ b/server/src/mls_message.rs @@ -0,0 +1,34 @@ +use axum::{ + async_trait, + body::Bytes, + extract::{FromRequest, Request}, + http::StatusCode, + response::{IntoResponse, Response}, +}; +use openmls::prelude::*; + +pub(crate) struct MlsMessage(pub(crate) MlsMessageIn); + +#[async_trait] +impl FromRequest for MlsMessage +where + Bytes: FromRequest, + S: Send + Sync, +{ + type Rejection = Response; + + async fn from_request(request: Request, state: &S) -> Result { + //TODO stream bytes directly into without buffering everything first + let bytes = Bytes::from_request(request, state) + .await + .map_err(IntoResponse::into_response)?; + + let mut bytes = bytes.as_ref(); + let package = MlsMessageIn::tls_deserialize(&mut bytes); + match package { + Ok(package) => Ok(MlsMessage(package)), + //TODO log error if it doesn't contain PII + Err(_) => Err(StatusCode::BAD_REQUEST.into_response()), + } + } +} diff --git a/server/src/user_actor.rs b/server/src/user_actor.rs new file mode 100644 index 0000000..9ab6ab9 --- /dev/null +++ b/server/src/user_actor.rs @@ -0,0 +1,136 @@ +use std::sync::Arc; + +use axum::{ + body::Bytes, + extract::ws::{Message, WebSocket}, +}; +use openmls::framing::MlsMessageIn; +use tokio::sync::{ + mpsc::{self, error::SendError}, + Mutex, +}; + +struct UserActor { + receiver: mpsc::Receiver, + websocket: WebSocket, + other_actors: Arc>>, +} +enum UserActorMessage { + /// Instruct the actor to send a message to the user represented by the actor + SendMessage(Vec), +} + +enum Instruction { + Continue, + Stop, +} + +impl UserActor { + fn new( + websocket: WebSocket, + receiver: mpsc::Receiver, + other_actors: Arc>>, + ) -> Self { + UserActor { + receiver, + websocket, + other_actors, + } + } + + async fn handle_websocket_message(&mut self, message: Message) -> Instruction { + if let Message::Close(close) = message { + tracing::debug!("Received close message: {:?}", close); + return Instruction::Stop; + } + + let Message::Binary(binary) = message else { + tracing::warn!("Received unexpected message type"); + return Instruction::Continue; + }; + + let mut others = self.other_actors.lock().await; + let mut dead_actors = Vec::new(); + for (index, other) in others.iter().enumerate() { + //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 + if let Err(error) = result { + dead_actors.push(index); + tracing::error!("Error sending actor message: {:?}", error); + } + } + + for index in dead_actors.into_iter() { + tracing::debug!("Cleaning up deceased actor remains at index: {:?}", index); + others.remove(index); + } + + Instruction::Continue + } + + async fn handle_message(&mut self, message: UserActorMessage) -> Instruction { + match message { + UserActorMessage::SendMessage(message) => { + tracing::debug!("Sending message: {:?}", message); + let result = self.websocket.send(Message::Binary(message)).await; + if let Err(error) = result { + tracing::error!("Error sending message: {:?}", error); + return Instruction::Stop; + } + + Instruction::Continue + } + } + } +} + +async fn run_my_actor(mut actor: UserActor) { + tracing::debug!("Actor started"); + loop { + tokio::select! { + Some(message) = actor.receiver.recv() => { + let result = actor.handle_message(message).await; + if let Instruction::Stop = result { + break; + } + }, + // Stop actor on error + Some(Ok(message)) = actor.websocket.recv() => { + let result = actor.handle_websocket_message(message).await; + if let Instruction::Stop = result { + break; + } + }, + else => break, + } + } + + tracing::debug!("Actor stopped"); +} + +pub(crate) struct UserActorHandle { + sender: mpsc::Sender, +} + +impl UserActorHandle { + pub(crate) fn new( + websocket: WebSocket, + other_actors: Arc>>, + ) -> Self { + let (sender, receiver) = mpsc::channel(8); + let actor = UserActor::new(websocket, receiver, other_actors); + tokio::spawn(run_my_actor(actor)); + Self { sender } + } + + /// Errors when actor stopped receiving messages meaning the channel is closed and the actor is deceased + pub(crate) async fn send_message( + &self, + message: Vec, + ) -> Result<(), impl std::error::Error> { + self.sender + .send(UserActorMessage::SendMessage(message)) + .await + } +} diff --git a/server/src/websocket_actor.rs b/server/src/websocket_actor.rs new file mode 100644 index 0000000..54f0bb0 --- /dev/null +++ b/server/src/websocket_actor.rs @@ -0,0 +1,13 @@ +use axum::extract::ws::WebSocket; + +struct WebSocketActor { + websocket: WebSocket, +} + +impl WebSocketActor { + pub(crate) fn new(websocket: WebSocket) -> Self { + Self { websocket } + } +} + +async fn run_websocket_actor(mut actor: WebSocketActor) {} diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index dfa98bb..d0814cc 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -2186,6 +2186,7 @@ dependencies = [ "tauri-plugin-shell", "thiserror", "tls_codec 0.3.0", + "tokio", ] [[package]] @@ -4387,9 +4388,9 @@ dependencies = [ [[package]] name = "tokio" -version = "1.37.0" +version = "1.38.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1adbebffeca75fcfd058afa480fb6c0b81e165a0323f9c9d39c9697e37c46787" +checksum = "ba4f4a02a7a80d6f274636f0aa95c7e383b912d41fe721a31f29e29698585a4a" dependencies = [ "backtrace", "bytes", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 1b24fe8..83e1581 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -27,3 +27,4 @@ thiserror = "1.0.61" base64 = "0.22.1" reqwest = "0.12.4" tls_codec = "0.3" +tokio = { version = "1.38.0", features = ["sync"] } diff --git a/src-tauri/capabilities/i-guess-i-need-an-identifier-but-idk-what-to-put-here.json b/src-tauri/capabilities/i-guess-i-need-an-identifier-but-idk-what-to-put-here.json new file mode 100644 index 0000000..4300c20 --- /dev/null +++ b/src-tauri/capabilities/i-guess-i-need-an-identifier-but-idk-what-to-put-here.json @@ -0,0 +1,7 @@ +{ + "identifier": "i-guess-i-need-an-identifier-but-idk-what-to-put-here", + "description": "", + "local": true, + "windows": ["main"], + "permissions": ["event:default", "event:allow-listen"] +} diff --git a/src-tauri/player2.tauri.config.json b/src-tauri/player2.tauri.config.json new file mode 100644 index 0000000..3c29b96 --- /dev/null +++ b/src-tauri/player2.tauri.config.json @@ -0,0 +1,7 @@ +{ + "build": { + "devUrl": "http://localhost:1421", + "frontendDist": "../dist", + "beforeDevCommand": "pnpm dev --port 1421" + } +} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 359dae4..cf4a750 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -1,12 +1,17 @@ use base64::prelude::*; use openmls::prelude::*; use openmls_basic_credential::SignatureKeyPair; -use openmls_rust_crypto::{MemoryKeyStore, MemoryKeyStoreError, OpenMlsRustCrypto}; +use openmls_rust_crypto::{MemoryKeyStoreError, OpenMlsRustCrypto}; use reqwest::{Client, Method}; use serde::Serialize; -use std::sync::{Mutex, PoisonError}; -use tauri::{Manager, State}; +use std::{ + collections::HashMap, + io::Read, + sync::{Arc, PoisonError}, +}; +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 = @@ -18,16 +23,14 @@ struct User { } struct AppState { - backend: OpenMlsRustCrypto, - user: Mutex>, - groups: Mutex>, + backend: Arc, + user: Arc>>, + groups: Arc>>, client: Client, } #[derive(Error, Debug, Serialize)] enum CreateUserError { - #[error("Could not access state")] - PoisonError(), #[error("User already exists")] UserExists, #[error("Error creating credentials for user")] @@ -52,15 +55,9 @@ enum CreateUserError { ), } -impl From> for CreateUserError { - fn from(_: PoisonError) -> Self { - CreateUserError::PoisonError() - } -} - #[tauri::command] -fn create_user(name: &str, state: State) -> Result<(), CreateUserError> { - let mut state = state.user.lock()?; +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); } @@ -96,15 +93,13 @@ impl From> for IsAuthenticatedError { } #[tauri::command] -fn is_authenticated(state: State) -> Result { - let state = state.user.lock()?; +async fn is_authenticated(state: State<'_, AppState>) -> Result { + let state = state.user.lock().await; Ok(state.is_some()) } #[derive(Error, Debug, Serialize)] enum CreateGroupError { - #[error("Could not access state")] - PoisonError, #[error("No user is signed in")] NoUserError, #[error("Error creating group")] @@ -113,8 +108,6 @@ enum CreateGroupError { #[serde(skip)] NewGroupError, ), - #[error("Could not access groups")] - GroupsPoisonError, } #[derive(Error, Debug, Serialize)] @@ -170,23 +163,15 @@ async fn advertise_key_package( #[derive(Error, Debug, Serialize)] enum AdvertiseError { - #[error("Could not access state")] - PoisonError, #[error("No user is signed in")] NoUserError, #[error("Error advertising key package")] AdvertiseKeyPackageError(#[from] AdvertiseKeyPackageError), } -impl From> for AdvertiseError { - fn from(_: PoisonError) -> Self { - AdvertiseError::PoisonError - } -} - #[tauri::command] async fn advertise(state: State<'_, AppState>) -> Result<(), AdvertiseError> { - let user = state.user.lock().map_err(|_| AdvertiseError::PoisonError)?; + let user = state.user.lock().await; let Some(user) = user.as_ref() else { return Err(AdvertiseError::NoUserError); }; @@ -198,15 +183,196 @@ async fn advertise(state: State<'_, AppState>) -> Result<(), AdvertiseError> { &state.client, ) .await?; + + Ok(()) +} +#[derive(Error, Debug, Serialize)] +enum GetPackageError { + #[error("Error getting package from server")] + RequestError( + #[from] + #[serde(skip)] + reqwest::Error, + ), + #[error("Error deserializing package")] + DeserializeError( + #[from] + #[serde(skip)] + tls_codec::Error, + ), +} + +async fn get_package(id: &str, client: &Client) -> Result { + let response = client + .get(&format!("http://localhost:3000/packages/{}", id)) + .send() + .await? + .error_for_status()?; + + //TODO directly stream bytes into deserializer without first buffering + let bytes = response.bytes().await?; + let package = KeyPackageIn::tls_deserialize(&mut bytes.as_ref())?; + + Ok(package) +} + +#[derive(Error, Debug, Serialize)] +enum SendMessageError { + #[error("Error sending message")] + RequestError( + #[from] + #[serde(skip)] + reqwest::Error, + ), + #[error("Error serializing message")] + SerializeError( + #[from] + #[serde(skip)] + tls_codec::Error, + ), +} + +async fn send_message( + recipient: String, + message: MlsMessageOut, + client: &Client, +) -> Result<(), SendMessageError> { + let message = message.tls_serialize_detached()?; + let response = client + .post(&format!("http://localhost:3000/messages/{}", recipient)) + .body(message) + .send() + .await?; + + response.error_for_status()?; + Ok(()) +} + +#[derive(Error, Debug, Serialize)] +enum InvitePackageError { + #[error("No user is signed in")] + NoUserError, + #[error("Group not found")] + GroupNotFound, + #[error("Error getting package from server")] + GetPackageError( + #[from] + #[serde(skip)] + GetPackageError, + ), + #[error("Error validating package")] + ValidatePackageError( + #[from] + #[serde(skip)] + KeyPackageVerifyError, + ), + + #[error("Error adding member to group")] + AddMemberError( + #[from] + #[serde(skip)] + AddMembersError, + ), + + #[error("Error serializing welcome message")] + SerializeWelcomeError( + #[from] + #[serde(skip)] + tls_codec::Error, + ), +} + +#[tauri::command] +async fn invite_package( + group_id: &str, + package_id: &str, + state: State<'_, AppState>, +) -> Result, InvitePackageError> { + let user = state.user.lock().await; + let Some(user) = user.as_ref() else { + return Err(InvitePackageError::NoUserError); + }; + + let mut groups = state.groups.lock().await; + let Some(group) = groups.get_mut(group_id) else { + return Err(InvitePackageError::GroupNotFound); + }; + + let package = get_package(package_id, &state.client).await?; + + let backend = state.backend.crypto(); + let package = package.validate(backend, ProtocolVersion::default())?; + + let (mls_message_out, welcome_out, group_information) = + group.add_members(state.backend.as_ref(), &user.signature_key, &[package])?; + + // Return welcome message to frontent to send over websockets to other clients + let data = welcome_out.tls_serialize_detached()?; + + Ok(data) +} + +#[derive(Error, Debug, Serialize)] +enum ReceiveMessageError { + #[error("Error deserializing message")] + DeserializeError( + #[from] + #[serde(skip)] + tls_codec::Error, + ), + #[error("Error joining group")] + JoinGroupError( + #[from] + #[serde(skip)] + WelcomeError, + ), + #[error("Error emitting event")] + EmitError( + #[from] + #[serde(skip)] + tauri::Error, + ), +} + +const JOIN_GROUP_EVENT: &str = "join_group"; + +#[derive(Serialize, Clone)] +struct JoinGroupEvent { + group_id: String, +} +#[tauri::command] +async fn process_message( + data: Vec, + state: State<'_, AppState>, + app: AppHandle, +) -> 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 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); + + app.emit(JOIN_GROUP_EVENT, JoinGroupEvent { group_id: id })?; + return Ok(()); + } + todo!() } #[tauri::command] -fn create_group(state: State) -> Result { - let user = state - .user - .lock() - .map_err(|_| CreateGroupError::PoisonError)?; +async fn create_group(state: State<'_, AppState>) -> Result { + let user = state.user.lock().await; let Some(user) = user.as_ref() else { return Err(CreateGroupError::NoUserError); }; @@ -216,7 +382,7 @@ fn create_group(state: State) -> Result { .build(); let group = MlsGroup::new( - &state.backend, + &*state.backend, &user.signature_key, &group_configuration, user.credential.clone(), @@ -225,44 +391,29 @@ fn create_group(state: State) -> Result { let id = group.group_id().as_slice(); let id = BASE64_URL_SAFE_NO_PAD.encode(id); - let mut groups = state - .groups - .lock() - .map_err(|_| CreateGroupError::GroupsPoisonError)?; - - groups.push(group); + let mut groups = state.groups.lock().await; + groups.insert(id.clone(), group); Ok(id) } -#[derive(Error, Debug, Serialize)] -enum GetGroupsError { - #[error("Could not access state")] - PoisonError, -} - +//TODO check if I can contribute support for Infallible error to tauri #[tauri::command] -fn get_groups(state: State) -> Result, GetGroupsError> { - let groups = state - .groups - .lock() - .map_err(|_| GetGroupsError::PoisonError)?; - - Ok(groups - .iter() - .map(|groups| BASE64_URL_SAFE_NO_PAD.encode(groups.group_id().as_slice())) - .collect()) - - // todo!() +async fn get_groups(state: State<'_, AppState>) -> Result, ()> { + let groups = state.groups.lock().await; + + let ids = groups.keys().cloned().collect(); + + Ok(ids) } #[cfg_attr(mobile, tauri::mobile_entry_point)] pub fn run() { let client = Client::new(); let state = AppState { - backend: OpenMlsRustCrypto::default(), - user: Mutex::new(None), - groups: Mutex::new(Vec::new()), + backend: OpenMlsRustCrypto::default().into(), + user: Arc::default(), + groups: Arc::default(), client, }; @@ -284,6 +435,7 @@ pub fn run() { create_group, create_user, get_groups, + invite_package, ]) .run(tauri::generate_context!()) .expect("error while running tauri application"); diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 3b017d8..8c04e0f 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -8,7 +8,8 @@ "beforeBuildCommand": "pnpm build", "frontendDist": "../dist" }, - "app": {"windows": [ + "app": { + "windows": [ { "title": "mealt", "width": 800, @@ -16,7 +17,8 @@ } ], "security": { - "csp": null + "csp": null, + "capabilities": ["i-guess-i-need-an-identifier-but-idk-what-to-put-here"] } }, "bundle": { diff --git a/src/App.tsx b/src/App.tsx index b015730..a9a4144 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -1,5 +1,6 @@ import { invoke } from "@tauri-apps/api/core"; import { For, Show, createResource, createSignal } from "solid-js"; +import { useAppState } from "./AppContext"; async function createUser(name: string) { await invoke("create_user", { name }); @@ -14,9 +15,7 @@ function App() { const [isAuthenticatedResource, { refetch: refetchIsAuthenticated }] = createResource(isAuthenticated); - const [groups, { mutate: mutateGroups }] = createResource( - async () => (await invoke("get_groups")) as string[] - ); + const { groups, setGroups } = useAppState(); function handleSubmit(event: SubmitEvent) { event.preventDefault(); @@ -29,7 +28,11 @@ function App() { async function handleCreateGroup() { const id = await createGroup(); - mutateGroups((groups) => (groups === undefined ? [id] : [...groups, id])); + setGroups((groups) => (groups === undefined ? [id] : [...groups, id])); + } + + async function handleAdvertise() { + await invoke("advertise"); } return ( @@ -44,6 +47,7 @@ function App() { +
    {(id) => ( diff --git a/src/AppContext.tsx b/src/AppContext.tsx new file mode 100644 index 0000000..fd4b2f1 --- /dev/null +++ b/src/AppContext.tsx @@ -0,0 +1,62 @@ +import { invoke } from "@tauri-apps/api/core"; +import { listen } from "@tauri-apps/api/event"; +import { + Accessor, + JSX, + Resource, + Setter, + createContext, + createResource, + onCleanup, + useContext, +} from "solid-js"; + +const socket = new WebSocket("ws://localhost:3000/ws"); +const [groups, { mutate: setGroups }] = createResource( + async () => (await invoke("get_groups")) as string[] +); +type AppState = { + socket: WebSocket; + groups: Resource; + setGroups: Setter; +}; +const state = { socket, groups, setGroups } satisfies AppState; +const AppContext = createContext(state); + +export function SocketProvider(properties: { children: JSX.Element }) { + return ( + + {properties.children} + + ); +} + +export function useWebSocket(onmessage?: (event: MessageEvent) => any) { + const { socket: webSocket } = useContext(AppContext); + + if (onmessage) { + webSocket.addEventListener("message", onmessage); + + onCleanup(() => { + webSocket.removeEventListener("message", onmessage); + }); + } + + return (data: string | ArrayBufferLike | Blob | ArrayBufferView) => + webSocket.send(data); +} + +export const useAppState = () => useContext(AppContext); + +listen("join_group", (event) => { + if ( + typeof event.payload !== "object" || + event.payload === null || + !("group_id" in event.payload) || + typeof event.payload.group_id !== "string" + ) + throw new Error("Unexpected join group event payload"); + + const group = event.payload.group_id; + setGroups((groups) => (groups === undefined ? [group] : [...groups, group])); +}); diff --git a/src/index.tsx b/src/index.tsx index 89e881f..95c59f0 100644 --- a/src/index.tsx +++ b/src/index.tsx @@ -5,13 +5,16 @@ import { Route, Router } from "@solidjs/router"; import "./styles.css"; import App from "./App"; import Groups from "./routes/Groups"; +import { SocketProvider } from "./AppContext"; render( () => ( - - - - + + + + + + ), document.body ); diff --git a/src/routes/Groups.tsx b/src/routes/Groups.tsx index d36c945..a177663 100644 --- a/src/routes/Groups.tsx +++ b/src/routes/Groups.tsx @@ -1,10 +1,51 @@ import { useParams } from "@solidjs/router"; +import { invoke } from "@tauri-apps/api/core"; +import { For, createResource } from "solid-js"; +import { useWebSocket } from "../AppContext"; + +async function getPackagesIndex() { + const response = await fetch(`http://localhost:3000/packages`); + + if (!response.ok) throw new Error("Could not fetch packages"); + + return (await response.json()) as string[]; +} export default function Groups() { const parameters = useParams(); + const sendMessage = useWebSocket(); + const groupId = () => parameters.id; + + const [packages] = createResource(getPackagesIndex); + + async function invitePackage(id: string) { + if (!groupId()) return; + + const message = (await invoke("invite_package", { + groupId: groupId(), + packageId: id, + })) as number[]; + + console.debug("invite", { message }); + + const data = Uint8Array.from(message); + sendMessage(data); + } return (
    -

    Group {parameters.id}

    +

    Group {groupId()}

    +

    Packages to invite

    +
      + + {(id) => ( +
    1. + +
    2. + )} +
      +
    ); }