diff --git a/Cargo.lock b/Cargo.lock index ebc11c1..f19e7c9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -84,131 +84,6 @@ dependencies = [ "syn 2.0.119", ] -[[package]] -name = "async-channel" -version = "1.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "81953c529336010edd6d8e358f886d9581267795c61b19475b71314bffa46d35" -dependencies = [ - "concurrent-queue", - "event-listener 2.5.3", - "futures-core", -] - -[[package]] -name = "async-channel" -version = "2.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "924ed96dd52d1b75e9c1a3e6275715fd320f5f9439fb5a4a11fa51f4221158d2" -dependencies = [ - "concurrent-queue", - "event-listener-strategy", - "futures-core", - "pin-project-lite", -] - -[[package]] -name = "async-executor" -version = "1.14.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c96bf972d85afc50bf5ab8fe2d54d1586b4e0b46c97c50a0c9e71e2f7bcd812a" -dependencies = [ - "async-task", - "concurrent-queue", - "fastrand", - "futures-lite", - "pin-project-lite", - "slab", -] - -[[package]] -name = "async-global-executor" -version = "2.4.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "05b1b633a2115cd122d73b955eadd9916c18c8f510ec9cd1686404c60ad1c29c" -dependencies = [ - "async-channel 2.5.0", - "async-executor", - "async-io", - "async-lock", - "blocking", - "futures-lite", - "once_cell", -] - -[[package]] -name = "async-io" -version = "2.6.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "456b8a8feb6f42d237746d4b3e9a178494627745c3c56c6ea55d92ba50d026fc" -dependencies = [ - "autocfg", - "cfg-if", - "concurrent-queue", - "futures-io", - "futures-lite", - "parking", - "polling", - "rustix", - "slab", - "windows-sys 0.61.2", -] - -[[package]] -name = "async-lock" -version = "3.4.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "290f7f2596bd5b78a9fec8088ccd89180d7f9f55b94b0576823bbbdc72ee8311" -dependencies = [ - "event-listener 5.4.1", - "event-listener-strategy", - "pin-project-lite", -] - -[[package]] -name = "async-native-tls" -version = "0.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9343dc5acf07e79ff82d0c37899f079db3534d99f189a1837c8e549c99405bec" -dependencies = [ - "futures-util", - "native-tls", - "thiserror 1.0.69", - "url", -] - -[[package]] -name = "async-std" -version = "1.13.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2c8e079a4ab67ae52b7403632e4618815d6db36d2a010cfe41b02c1b1578f93b" -dependencies = [ - "async-channel 1.9.0", - "async-global-executor", - "async-io", - "async-lock", - "crossbeam-utils", - "futures-channel", - "futures-core", - "futures-io", - "futures-lite", - "gloo-timers", - "kv-log-macro", - "log", - "memchr", - "once_cell", - "pin-project-lite", - "pin-utils", - "slab", - "wasm-bindgen-futures", -] - -[[package]] -name = "async-task" -version = "4.7.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b75356056920673b02621b35afd0f7dda9306d03c79a30f5c56c44cf256e3de" - [[package]] name = "async-trait" version = "0.1.91" @@ -217,25 +92,7 @@ checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec" dependencies = [ "proc-macro2", "quote", - "syn 3.0.1", -] - -[[package]] -name = "async-tungstenite" -version = "0.29.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ef0f7efedeac57d9b26170f72965ecfd31473ca52ca7a64e925b0b6f5f079886" -dependencies = [ - "async-native-tls", - "async-std", - "atomic-waker", - "futures-core", - "futures-io", - "futures-task", - "futures-util", - "log", - "pin-project-lite", - "tungstenite 0.26.2", + "syn 3.0.2", ] [[package]] @@ -334,19 +191,6 @@ dependencies = [ "objc2", ] -[[package]] -name = "blocking" -version = "1.6.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e83f8d02be6967315521be875afa792a316e28d57b5a2d401897e2a7921b7f21" -dependencies = [ - "async-channel 2.5.0", - "async-task", - "futures-io", - "futures-lite", - "piper", -] - [[package]] name = "bumpalo" version = "3.20.3" @@ -475,15 +319,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "concurrent-queue" -version = "2.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4ca0197aee26d1ae37445ee532fefce43251d24cc7c166799f4d46817f1d3973" -dependencies = [ - "crossbeam-utils", -] - [[package]] name = "const-oid" version = "0.9.6" @@ -784,33 +619,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "event-listener" -version = "2.5.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0206175f82b8d6bf6652ff7d71a1e27fd2e4efde587fd368662814d6ec1d9ce0" - -[[package]] -name = "event-listener" -version = "5.4.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e13b66accf52311f30a0db42147dadea9850cb48cd070028831ae5f5d4b856ab" -dependencies = [ - "concurrent-queue", - "parking", - "pin-project-lite", -] - -[[package]] -name = "event-listener-strategy" -version = "0.5.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8be9f3dfaaffdae2972880079a491a1a8bb7cbed0b8dd7a347f668b4150a3b93" -dependencies = [ - "event-listener 5.4.1", - "pin-project-lite", -] - [[package]] name = "fastbloom" version = "0.17.0" @@ -853,21 +661,6 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb" -[[package]] -name = "foreign-types" -version = "0.3.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f6f339eb8adc052cd2ca78910fda869aefa38d22d5cb648e6485e4d3fc06f3b1" -dependencies = [ - "foreign-types-shared", -] - -[[package]] -name = "foreign-types-shared" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "00b0228411908ca8685dba7fc2cdd70ec9990a6e753e89b6ac91a84c40fbaf4b" - [[package]] name = "form_urlencoded" version = "1.2.2" @@ -931,19 +724,6 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4577ecaa3c4f96589d473f679a71b596316f6641bc350038b962a5daf0085d7a" -[[package]] -name = "futures-lite" -version = "2.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f78e10609fe0e0b3f4157ffab1876319b5b0db102a2c60dc4626306dc46b44ad" -dependencies = [ - "fastrand", - "futures-core", - "futures-io", - "parking", - "pin-project-lite", -] - [[package]] name = "futures-macro" version = "0.3.33" @@ -1033,18 +813,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "gloo-timers" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bbb143cf96099802033e0d4f4963b19fd2e0b728bcf076cd9cf7f6634f092994" -dependencies = [ - "futures-channel", - "futures-core", - "js-sys", - "wasm-bindgen", -] - [[package]] name = "h3" version = "0.0.8" @@ -1125,12 +893,6 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" -[[package]] -name = "hermit-abi" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" - [[package]] name = "hkdf" version = "0.13.0" @@ -1215,9 +977,9 @@ dependencies = [ [[package]] name = "hyper" -version = "1.10.1" +version = "1.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498" +checksum = "d22053281f852e11534f5198498373cbb59295120a20771d90f7ed1897490a72" dependencies = [ "atomic-waker", "bytes", @@ -1234,18 +996,18 @@ dependencies = [ ] [[package]] -name = "hyper-tls" -version = "0.6.0" +name = "hyper-rustls" +version = "0.27.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "70206fc6890eaca9fde8a0bf71caa2ddfc9fe045ac9e5c70df101a7dbde866e0" +checksum = "33ca68d021ef39cf6463ab54c1d0f5daf03377b70561305bb89a8f83aab66e0f" dependencies = [ - "bytes", - "http-body-util", + "http", "hyper", "hyper-util", - "native-tls", + "rustls", + "rustls-native-certs", "tokio", - "tokio-native-tls", + "tokio-rustls", "tower-service", ] @@ -1542,12 +1304,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "json" -version = "0.12.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "078e285eafdfb6c4b434e0d31e8cfcb5115b651496faca5749b88fafd4f23bfd" - [[package]] name = "jsonwebtoken" version = "10.4.0" @@ -1582,15 +1338,6 @@ dependencies = [ "cpufeatures 0.3.0", ] -[[package]] -name = "kv-log-macro" -version = "1.0.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0de8b303297635ad57c9f5059fd9cee7a47f8e8daa09df0fcd07dd39fb22977f" -dependencies = [ - "log", -] - [[package]] name = "lazy_static" version = "1.5.0" @@ -1633,7 +1380,6 @@ version = "0.5.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e827d3444235ddccf360fc1d4737034e86ee2ea2c0a696f148ecc08e7249bfd3" dependencies = [ - "async-tungstenite", "base64 0.21.7", "device-info", "hmac 0.12.1", @@ -1649,6 +1395,7 @@ dependencies = [ "prost", "rand 0.9.5", "reqwest", + "rustls-native-certs", "scopeguard", "serde", "serde_json", @@ -1656,6 +1403,7 @@ dependencies = [ "signature 2.2.0", "thiserror 2.0.19", "tokio", + "tokio-rustls", "tokio-tungstenite", "url", ] @@ -1705,9 +1453,6 @@ name = "log" version = "0.4.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" -dependencies = [ - "value-bag", -] [[package]] name = "lru-cache" @@ -1958,23 +1703,6 @@ version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d87ecb2933e8aeadb3e3a02b828fed80a7528047e68b4f424523a0981a3a084" -[[package]] -name = "native-tls" -version = "0.2.18" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "465500e14ea162429d264d44189adc38b199b62b1c21eea9f69e4b73cb03bbf2" -dependencies = [ - "libc", - "log", - "openssl", - "openssl-probe", - "openssl-sys", - "schannel", - "security-framework", - "security-framework-sys", - "tempfile", -] - [[package]] name = "nix" version = "0.31.3" @@ -2213,17 +1941,14 @@ dependencies = [ "base64 0.22.1", "dashmap", "dotenv", - "futures", - "json", "livekit-api", "livekit-protocol", "log", "mtp", "mtp-transport", "once_cell", - "rand 0.8.7", + "rand 0.10.2", "rustls", - "serde", "serde_json", "strum", "strum_macros", @@ -2245,49 +1970,12 @@ version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" -[[package]] -name = "openssl" -version = "0.10.81" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "77823a27f0babb03091cb9ed9ef80af3b39dbc82f97e8fa530374b7dafd87a45" -dependencies = [ - "bitflags", - "cfg-if", - "foreign-types", - "libc", - "openssl-macros", - "openssl-sys", -] - -[[package]] -name = "openssl-macros" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a948666b637a0f465e8564c73e89d4dde00d72d4d473cc972f390fc3dcee7d9c" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.119", -] - [[package]] name = "openssl-probe" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" -[[package]] -name = "openssl-sys" -version = "0.9.117" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b47e7e6bb2c38cd930d25a23b40fa52e068c10e85f3e03a7f5ba5aaca5713695" -dependencies = [ - "cc", - "libc", - "pkg-config", - "vcpkg", -] - [[package]] name = "os_info" version = "3.15.0" @@ -2304,12 +1992,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "parking" -version = "2.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba" - [[package]] name = "parking_lot" version = "0.12.5" @@ -2411,23 +2093,6 @@ version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" -[[package]] -name = "pin-utils" -version = "0.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184" - -[[package]] -name = "piper" -version = "0.2.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c835479a4443ded371d6c535cbfd8d31ad92c5d23ae9770a61bc155e4992a3c1" -dependencies = [ - "atomic-waker", - "fastrand", - "futures-io", -] - [[package]] name = "pkcs8" version = "0.10.2" @@ -2454,20 +2119,6 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" -[[package]] -name = "polling" -version = "3.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d0e4f59085d47d8241c88ead0f274e8a0cb551f3625263c05eb8dd897c34218" -dependencies = [ - "cfg-if", - "concurrent-queue", - "hermit-abi", - "pin-project-lite", - "rustix", - "windows-sys 0.61.2", -] - [[package]] name = "poly1305" version = "0.8.0" @@ -2813,20 +2464,22 @@ dependencies = [ "http-body", "http-body-util", "hyper", - "hyper-tls", + "hyper-rustls", "hyper-util", "js-sys", "log", - "native-tls", "percent-encoding", "pin-project-lite", + "quinn", + "rustls", + "rustls-native-certs", "rustls-pki-types", "serde", "serde_json", "serde_urlencoded", "sync_wrapper", "tokio", - "tokio-native-tls", + "tokio-rustls", "tower", "tower-http", "tower-service", @@ -3062,14 +2715,14 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn 3.0.1", + "syn 3.0.2", ] [[package]] name = "serde_json" -version = "1.0.150" +version = "1.0.151" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" +checksum = "c841b55ecdae098c80dcae9cf767f6f8a0c2cdb3416bbef72181df4d0fe73f14" dependencies = [ "itoa", "memchr", @@ -3305,9 +2958,9 @@ dependencies = [ [[package]] name = "syn" -version = "3.0.1" +version = "3.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5edbec4ed188954a10c12c038215f8ce7606b2d5c973cd8dc43e8795065c5f2f" +checksum = "a207d6d6a2b7fc470b80443726053f18a2481b7e1eee970597051596567987a3" dependencies = [ "proc-macro2", "quote", @@ -3384,14 +3037,14 @@ checksum = "43cbfe0cf76104d42a574802844187e84a305e531ed54455f11fbde0f10541cd" dependencies = [ "proc-macro2", "quote", - "syn 3.0.1", + "syn 3.0.2", ] [[package]] name = "time" -version = "0.3.53" +version = "0.3.54" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "18dfaaeddcb932337b5e7866ee7d0ce9b76d2fd092997146f187ec09b4558a50" +checksum = "3e1d5e639ff6bab73cb6885cc7e7b1de96c3f32c68ec55f3952614bec1092244" dependencies = [ "deranged", "num-conv", @@ -3409,9 +3062,9 @@ checksum = "9e1c906769ad99c88eaa54e728060edef082f8e358ff32030cb7c7d315e81109" [[package]] name = "time-macros" -version = "0.2.31" +version = "0.2.32" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c431b87111666e491a90baa837f914fb45cd5dc3c268591b0220ff5057f2085f" +checksum = "7e689342a48d2ea927c87ea50cabf8594854bf940e9310208848d680d668ed85" dependencies = [ "num-conv", "time-core", @@ -3444,9 +3097,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.53.0" +version = "1.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d988bcd52dbe076d3d46903332f58c912b87a2c49b1428419a5845154762ffee" +checksum = "202caea871b69668250d242070849eb495be178ed697a3e98aebce5bc81a0bed" dependencies = [ "bytes", "libc", @@ -3471,12 +3124,12 @@ dependencies = [ ] [[package]] -name = "tokio-native-tls" -version = "0.3.1" +name = "tokio-rustls" +version = "0.26.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bbae76ab933c85776efabc971569dd6119c580d8f5d448769dec1764bf796ef2" +checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" dependencies = [ - "native-tls", + "rustls", "tokio", ] @@ -3499,10 +3152,12 @@ checksum = "8f72a05e828585856dacd553fba484c242c46e391fb0e58917c942ee9202915c" dependencies = [ "futures-util", "log", - "native-tls", + "rustls", + "rustls-native-certs", + "rustls-pki-types", "tokio", - "tokio-native-tls", - "tungstenite 0.29.0", + "tokio-rustls", + "tungstenite", ] [[package]] @@ -3647,25 +3302,6 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" -[[package]] -name = "tungstenite" -version = "0.26.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4793cb5e56680ecbb1d843515b23b6de9a75eb04b66643e256a396d43be33c13" -dependencies = [ - "bytes", - "data-encoding", - "http", - "httparse", - "log", - "native-tls", - "rand 0.9.5", - "sha1", - "thiserror 2.0.19", - "url", - "utf-8", -] - [[package]] name = "tungstenite" version = "0.29.0" @@ -3677,8 +3313,9 @@ dependencies = [ "http", "httparse", "log", - "native-tls", "rand 0.9.5", + "rustls", + "rustls-pki-types", "sha1", "thiserror 2.0.19", "url", @@ -3751,12 +3388,6 @@ dependencies = [ "serde", ] -[[package]] -name = "utf-8" -version = "0.7.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" - [[package]] name = "utf8_iter" version = "1.0.4" @@ -3774,18 +3405,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "value-bag" -version = "1.13.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5dd4ec1eb1d240636e354a30110a1dfcb37047169a4d9bd6d9d3469df574b5c4" - -[[package]] -name = "vcpkg" -version = "0.2.15" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "accd4ea62f7bb7a82fe23066fb0957d48ef677f6eeb8215f372f52e48bb32426" - [[package]] name = "version_check" version = "0.9.5" @@ -4253,18 +3872,18 @@ dependencies = [ [[package]] name = "zerocopy" -version = "0.8.54" +version = "0.8.55" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b7cbbc0a705a0fd05cc3676525980d2bf5a9bc4adac6d6475209a7887cf59d19" +checksum = "b5a105cd7b140f6eeec8acff2ea38135d3cab283ada58540f629fe51e46696eb" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.54" +version = "0.8.55" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e2e817b7b52d0c7358d3246da9d69935ebb18116b2b102b4230dac079b4862f5" +checksum = "0fe976fb70c78cd64cccfe3a6fc142244e8a77b70959b30faf9d0ac37ee228eb" dependencies = [ "proc-macro2", "quote", diff --git a/Cargo.toml b/Cargo.toml index 32beb7e..2c37c3e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -12,28 +12,25 @@ mtp = { git = "https://git.methanium.net/Methanium/mtp.git", features = [ ] } mtp-transport = { git = "https://git.methanium.net/Methanium/mtp.git" } -ansi_term = "*" -uuid = { version = "*", features = ["v4"] } +ansi_term = "0.12.1" +uuid = { version = "1.24.0", features = ["v4"] } base64 = "0.22.1" -dashmap = "*" -futures = "*" +dashmap = "6.2.1" once_cell = "1.21.4" -rand = "0.8" +rand = "0.10.2" rustls = { version = "0.23.42", default-features = false, features = [ "std", "tls12", "aws-lc-rs", "prefer-post-quantum", ] } -tokio = { version = "*", features = ["full"] } +tokio = { version = "1.53.0", features = ["full"] } log = "0.4" dotenv = "0.15.0" strum = "0.28.0" strum_macros = "0.28.0" -livekit-api = { version = "0.5.6", features = ["native-tls"] } +livekit-api = { version = "0.5.6", features = ["rustls-tls-native-roots"] } livekit-protocol = "0.7.10" thiserror = "2.0.19" trust-dns-resolver = "0.23.2" -serde = "1.0.229" -serde_json = "1.0.150" -json = "0.12.4" +serde_json = "1.0.151" diff --git a/src/anonymous_clients/anonymous_client_connection.rs b/src/anonymous_clients/anonymous_client_connection.rs index 750f628..22f86e4 100644 --- a/src/anonymous_clients/anonymous_client_connection.rs +++ b/src/anonymous_clients/anonymous_client_connection.rs @@ -6,15 +6,16 @@ use tokio::sync::RwLock; use uuid::Uuid; use crate::anonymous_clients::anonymous_manager::{self, generate_username}; -use crate::calls::{call_group::call_invite_secret_from_cv, call_manager}; +use crate::app_state::AppState; +use crate::calls::call_group::call_invite_secret_from_cv; use crate::data::user::UserStatus; -use crate::omega::omega_connection::{OmegaConnection, get_omega_connection}; use crate::rho::connection::{GeneralConnection, MtpReceiver, MtpSender}; -use crate::rho::rho_manager; +use crate::util::data_type_id; use crate::util::logger::PrintType; use crate::{log_cv_in, log_cv_out, log_out}; pub struct AnonymousClientConnection { + pub state: Arc, user_id: u64, pub sender: Arc, @@ -25,12 +26,14 @@ pub struct AnonymousClientConnection { pub user_name: Arc>, pub display_name: Arc>, pub avatar: Arc>, + message_slots: Arc, } impl AnonymousClientConnection { pub async fn from_general(general: Arc, user_id: u64) -> Arc { let username: String = generate_username(); Arc::new(Self { + state: general.state.clone(), user_id: user_id, ping: Arc::new(RwLock::new(0)), @@ -41,6 +44,7 @@ impl AnonymousClientConnection { user_name: Arc::new(RwLock::new(username.to_lowercase())), display_name: Arc::new(RwLock::new(username)), avatar: Arc::new(RwLock::new(String::new())), + message_slots: Arc::new(tokio::sync::Semaphore::new(32)), }) } pub fn start(self: Arc) { @@ -102,7 +106,11 @@ impl AnonymousClientConnection { /// Handle incoming message from client pub async fn handle_message(self: Arc, cv: CommunicationValue) { + let Ok(permit) = self.message_slots.clone().acquire_owned().await else { + return; + }; tokio::spawn(async move { + let _permit = permit; if cv.is_type(CommunicationType::Ping) { self.handle_ping(cv).await; return; @@ -113,7 +121,7 @@ impl AnonymousClientConnection { let call_id = Uuid::parse_str(cv.get_data(DataType::CallId).as_str().unwrap_or("")) .unwrap_or(Uuid::new_v4()); - let call = if let Some(call) = call_manager::get_call(call_id).await { + let call = if let Some(call) = self.state.call_manager.get_call(call_id).await { if call.is_anonymous().await { call } else { @@ -135,7 +143,10 @@ impl AnonymousClientConnection { let mut invited = Vec::new(); for call_invitee in call.members.read().await.clone() { - let call_invitee_cv = get_omega_connection() + let Ok(call_invitee_cv) = self + .state + .omega + .clone() .await_response( &CommunicationValue::new(CommunicationType::GetUserData) .add_typed_default( @@ -145,7 +156,9 @@ impl AnonymousClientConnection { Some(Duration::from_secs(2)), ) .await - .unwrap(); + else { + continue; + }; let mut json_invitee = Vec::new(); let _ = json_invitee.push(( DataType::UserId, @@ -167,18 +180,21 @@ impl AnonymousClientConnection { let _ = invited.push(DataValue::Container( json_invitee .iter() - .map(|(k, v)| (k.to_id(&TypeMap::latest()), v.clone())) + .map(|(k, v)| (data_type_id(*k, &TypeMap::latest()), v.clone())) .collect(), )); } - let token = call.create_anonymous_token(self.get_user_id()).await; + let token = match call.create_anonymous_token(self.get_user_id()).await { + Ok(Some(token)) => token, + Ok(None) | Err(_) => return, + }; let mut serialized = Vec::new(); let _ = serialized.push((DataType::CallId, DataValue::Str(call_id.to_string()))); let _ = serialized.push((DataType::CallInvited, DataValue::Array(invited.clone()))); let _ = serialized.push((DataType::CallMembers, DataValue::Array(invited))); - let _ = serialized.push((DataType::CallToken, DataValue::Str(token.unwrap()))); + let _ = serialized.push((DataType::CallToken, DataValue::Str(token))); self.clone() .send_message( &&CommunicationValue::new(CommunicationType::IdentificationResponse) @@ -204,7 +220,9 @@ impl AnonymousClientConnection { DataValue::Container( serialized .iter() - .map(|(k, v)| (k.to_id(&TypeMap::latest()), v.clone())) + .map(|(k, v)| { + (data_type_id(*k, &TypeMap::latest()), v.clone()) + }) .collect(), ), ), @@ -304,7 +322,10 @@ impl AnonymousClientConnection { async fn handle_omega_forward(self: Arc, cv: CommunicationValue) { let client_for_closure = self.clone(); tokio::spawn(async move { - let response_cv = get_omega_connection() + let response_cv = self + .state + .omega + .clone() .await_response(&cv.with_sender(self.user_id), Some(Duration::from_secs(20))) .await; if let Ok(response_cv) = response_cv { @@ -333,7 +354,9 @@ impl AnonymousClientConnection { async fn handle_client_changed(self: Arc, cv: CommunicationValue) { if let DataValue::Str(status_str) = cv.get_data(DataType::UserState) { let user_status = UserStatus::from_str(&status_str).unwrap_or(UserStatus::user_online); - OmegaConnection::client_changed(self.user_id as i64, self.user_id as i64, user_status) + self.state + .omega + .client_changed(self.user_id as i64, self.user_id as i64, user_status) .await; } } @@ -371,23 +394,29 @@ impl AnonymousClientConnection { return; } }; - let invited = - call_manager::add_invite(call_id, self.user_id, receiver_id as u64, secret.clone()) - .await; + let invited = self + .state + .call_manager + .add_invite(call_id, self.user_id, receiver_id as u64, secret.clone()) + .await; if !invited { self.send_error_response(&cv.get_id(), CommunicationType::ErrorInvalidCallId) .await; return; } - if !call_manager::should_forward_invite(self.user_id, receiver_id as u64) { + if !self + .state + .call_manager + .should_forward_invite(self.user_id, receiver_id as u64) + { let response = CommunicationValue::new(CommunicationType::Success).with_id(cv.get_id()); self.send_message(&response).await; return; } // Find target RhoConnection - let target_rho = match rho_manager::get_rho_con_for_user(receiver_id).await { + let target_rho = match self.state.rho.get_for_user(receiver_id).await { Some(rho) => rho, _ => { // Get sender user ID @@ -406,7 +435,7 @@ impl AnonymousClientConnection { DataValue::Str("call_invite".to_string()), ); - let omega_conn = get_omega_connection(); + let omega_conn = self.state.omega.clone(); // Send fire-and-forget, don't await to avoid blocking tokio::spawn(async move { let _ = omega_conn.send_message(&push_cv).await; @@ -467,22 +496,39 @@ impl AnonymousClientConnection { } }; - if let Some(token) = call_manager::get_call_token(user_id, call_id).await { - let response = CommunicationValue::new(CommunicationType::CallToken) - .with_id(cv.get_id()) - .with_receiver(user_id) - .add_typed_default(DataType::CallToken, DataValue::Str(token.to_string())); - self.send_message(&response).await; - } else { - let error_cv = CommunicationValue::new(CommunicationType::ErrorNoCallId) - .with_id(cv.get_id()) - .add_typed_default(DataType::CallId, DataValue::Str(call_id.to_string())); - self.send_message(&error_cv).await; - return; + match self + .state + .call_manager + .get_call_token(user_id, call_id) + .await + { + Ok(token) => { + let response = CommunicationValue::new(CommunicationType::CallToken) + .with_id(cv.get_id()) + .with_receiver(user_id) + .add_typed_default(DataType::CallToken, DataValue::Str(token)); + self.send_message(&response).await; + } + Err(error) => { + log::warn!( + "Unable to create anonymous call token for {}: {}", + call_id, + error + ); + let error_cv = CommunicationValue::new(CommunicationType::ErrorNoCallId) + .with_id(cv.get_id()) + .add_typed_default(DataType::CallId, DataValue::Str(call_id.to_string())); + self.send_message(&error_cv).await; + } } } async fn handle_call_timeout_user(self: Arc, cv: CommunicationValue) { - let call_id = Uuid::from_str(cv.get_data(DataType::CallId).as_str().unwrap_or("")).unwrap(); + let Ok(call_id) = Uuid::from_str(cv.get_data(DataType::CallId).as_str().unwrap_or("")) + else { + self.send_error_response(&cv.get_id(), CommunicationType::ErrorInvalidCallId) + .await; + return; + }; let user_id = cv .get_data(DataType::UserId) .as_signed_number() @@ -492,39 +538,46 @@ impl AnonymousClientConnection { .as_signed_number() .unwrap_or(0); - let call = call_manager::get_call(call_id).await; - if let Some(call) = call { - if call - .get_caller(self.get_user_id()) - .await - .unwrap() - .has_admin() - { - call.get_caller(user_id as u64) - .await - .unwrap() - .set_timeout(untill.try_into().unwrap()) - .await; + let Some(call) = self.state.call_manager.get_call(call_id).await else { + self.send_error_response(&cv.get_id(), CommunicationType::ErrorNotFound) + .await; + return; + }; + let Some(caller) = call.get_caller(self.get_user_id()).await else { + self.send_error_response(&cv.get_id(), CommunicationType::ErrorInvalidUserId) + .await; + return; + }; + if caller.has_admin() { + if let Some(target) = call.get_caller(user_id as u64).await { + target.set_timeout(untill as i64).await; } } } async fn handle_call_disconnect_user(self: Arc, cv: CommunicationValue) { - let call_id = Uuid::from_str(cv.get_data(DataType::CallId).as_str().unwrap_or("")).unwrap(); + let Ok(call_id) = Uuid::from_str(cv.get_data(DataType::CallId).as_str().unwrap_or("")) + else { + self.send_error_response(&cv.get_id(), CommunicationType::ErrorInvalidCallId) + .await; + return; + }; let user_id = cv .get_data(DataType::UserId) .as_signed_number() .unwrap_or(0); - let call = call_manager::get_call(call_id).await; - if let Some(call) = call { - if call - .get_caller(self.get_user_id()) - .await - .unwrap() - .has_admin() - { - call.remove_caller(user_id as u64).await; - } + let Some(call) = self.state.call_manager.get_call(call_id).await else { + self.send_error_response(&cv.get_id(), CommunicationType::ErrorNotFound) + .await; + return; + }; + let Some(caller) = call.get_caller(self.get_user_id()).await else { + self.send_error_response(&cv.get_id(), CommunicationType::ErrorInvalidUserId) + .await; + return; + }; + if caller.has_admin() { + call.remove_caller(user_id as u64).await; } } @@ -586,6 +639,7 @@ impl AnonymousClientConnection { impl Clone for AnonymousClientConnection { fn clone(&self) -> Self { Self { + state: self.state.clone(), sender: Arc::clone(&self.sender), receiver: Arc::clone(&self.receiver), user_id: self.user_id, @@ -595,6 +649,7 @@ impl Clone for AnonymousClientConnection { user_name: Arc::clone(&self.user_name), display_name: Arc::clone(&self.display_name), avatar: Arc::clone(&self.avatar), + message_slots: Arc::clone(&self.message_slots), } } } diff --git a/src/anonymous_clients/anonymous_manager.rs b/src/anonymous_clients/anonymous_manager.rs index 8330ecf..41703bf 100644 --- a/src/anonymous_clients/anonymous_manager.rs +++ b/src/anonymous_clients/anonymous_manager.rs @@ -1,7 +1,6 @@ use dashmap::DashMap; use once_cell::sync::Lazy; -use rand::Rng; -use rand::seq::SliceRandom; +use rand::prelude::{IndexedRandom, RngExt}; use std::sync::Arc; use crate::anonymous_clients::anonymous_client_connection::AnonymousClientConnection; @@ -45,11 +44,11 @@ pub fn generate_username() -> String { let adjectives = ["Swift", "Clever", "Brave", "Sneaky", "Fierce"]; let nouns = ["Tiger", "Eagle", "Shark", "Wolf", "Dragon"]; - let mut rng = rand::thread_rng(); + let mut rng = rand::rng(); let adj = adjectives.choose(&mut rng).unwrap(); let noun = nouns.choose(&mut rng).unwrap(); - let number: u16 = rng.gen_range(0..10000); + let number: u16 = rng.random_range(0..10000); format!("{}{}{}", adj, noun, number) } diff --git a/src/app_state.rs b/src/app_state.rs new file mode 100644 index 0000000..0b8fb5a --- /dev/null +++ b/src/app_state.rs @@ -0,0 +1,50 @@ +use std::sync::Arc; + +use mtp::crypto::Keyring; + +use crate::{ + calls::{call_manager::CallManager, call_util::LiveKitService}, + config::Config, + omega::omega_connection::OmegaConnection, + rho::rho_manager::RhoManager, + services::call_state::CallStateAggregator, +}; + +/* + * Holds startup dependencies so listener and connection constructors can take + * one explicit handle while the remaining manager migrations are completed. + */ +pub struct AppState { + pub config: Config, + pub keyring: Keyring, + pub omega: Arc, + pub rho: Arc, + pub call_manager: Arc, + pub call_state_aggregator: Arc, + pub livekit: Arc, +} + +impl AppState { + pub fn new( + config: Config, + keyring: Keyring, + omega: Arc, + rho: Arc, + call_manager: Arc, + livekit: Arc, + ) -> Arc { + Arc::new(Self { + config, + keyring, + omega, + rho, + call_state_aggregator: Arc::new(CallStateAggregator::new(call_manager.clone())), + call_manager, + livekit, + }) + } + + pub fn keyring_for_host(&self) -> Result { + Keyring::from_bytes(&self.keyring.to_bytes()).map_err(|error| error.to_string()) + } +} diff --git a/src/calls/call_group.rs b/src/calls/call_group.rs index 6c1fb1d..ef03b0f 100755 --- a/src/calls/call_group.rs +++ b/src/calls/call_group.rs @@ -1,13 +1,14 @@ use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue, TypeMap}; use serde_json::Map; -use std::{collections::BTreeMap, env, sync::Arc, time::Duration}; +use std::{collections::BTreeMap, sync::Arc, time::Duration}; use tokio::sync::RwLock; use uuid::Uuid; use crate::{ - calls::{call_util, caller::Caller}, - omega::omega_connection::get_omega_connection, + calls::{call_util::LiveKitService, caller::Caller, error::CallError}, + omega::omega_connection::OmegaConnection, + util::data_type_id, }; pub struct CallGroup { @@ -17,6 +18,7 @@ pub struct CallGroup { pub anonymous_joining: RwLock, pub short_link: RwLock>, pub secrets: RwLock>, + livekit: Arc, } #[derive(Clone, Debug, Eq, PartialEq)] @@ -32,27 +34,27 @@ impl CallSecretEnvelope { pub fn from_data_value(value: &DataValue) -> Option { let tm = TypeMap::latest(); let secret_id = value - .get_field(DataType::SecretId.to_id(&tm))? + .get_field(data_type_id(DataType::SecretId, &tm))? .as_str()? .to_string(); let version_number = value - .get_field(DataType::VersionNumber.to_id(&tm))? + .get_field(data_type_id(DataType::VersionNumber, &tm))? .as_signed_number() .and_then(|n| i64::try_from(n).ok()) .or_else(|| { value - .get_field(DataType::VersionNumber.to_id(&tm))? + .get_field(data_type_id(DataType::VersionNumber, &tm))? .as_number() .and_then(|n| i64::try_from(n).ok()) })?; let encrypted_secret = value - .get_field(DataType::EncryptedSecret.to_id(&tm))? + .get_field(data_type_id(DataType::EncryptedSecret, &tm))? .as_bytes()?; let kem_ciphertext = value - .get_field(DataType::KemCiphertext.to_id(&tm))? + .get_field(data_type_id(DataType::KemCiphertext, &tm))? .as_bytes()?; let wrapping_scheme = value - .get_field(DataType::WrappingScheme.to_id(&tm))? + .get_field(data_type_id(DataType::WrappingScheme, &tm))? .as_str()? .to_string(); @@ -69,23 +71,23 @@ impl CallSecretEnvelope { let tm = TypeMap::latest(); let mut map: BTreeMap = BTreeMap::new(); map.insert( - DataType::SecretId.to_id(&tm), + data_type_id(DataType::SecretId, &tm), DataValue::Str(self.secret_id.clone()), ); map.insert( - DataType::VersionNumber.to_id(&tm), + data_type_id(DataType::VersionNumber, &tm), DataValue::SignedNumber(self.version_number.into()), ); map.insert( - DataType::EncryptedSecret.to_id(&tm), + data_type_id(DataType::EncryptedSecret, &tm), DataValue::Bytes(self.encrypted_secret.clone()), ); map.insert( - DataType::KemCiphertext.to_id(&tm), + data_type_id(DataType::KemCiphertext, &tm), DataValue::Bytes(self.kem_ciphertext.clone()), ); map.insert( - DataType::WrappingScheme.to_id(&tm), + data_type_id(DataType::WrappingScheme, &tm), DataValue::Str(self.wrapping_scheme.clone()), ); DataValue::container_from_map(&map) @@ -98,6 +100,14 @@ pub fn call_invite_secret_from_cv(cv: &CommunicationValue) -> Option) -> Self { + Self::new_with_service(call_id, user, Arc::new(LiveKitService::new(None))) + } + + pub fn new_with_service( + call_id: Uuid, + user: Arc, + livekit: Arc, + ) -> Self { CallGroup { call_id, members: RwLock::new(vec![user]), @@ -105,6 +115,7 @@ impl CallGroup { anonymous_joining: RwLock::new(false), short_link: RwLock::new(None), secrets: RwLock::new(BTreeMap::new()), + livekit, } } @@ -126,7 +137,7 @@ impl CallGroup { } pub async fn update_admins(&self) { - let call_metadata = match call_util::get_room_metadata(self.call_id).await { + let call_metadata = match self.livekit.get_room_metadata(self.call_id).await { Ok(metadata) => metadata, Err(_) => "{}".to_string(), }; @@ -150,17 +161,31 @@ impl CallGroup { call_metadata.insert("admins".to_string(), serde_json::json!(admin_ids)); - let _ = call_util::set_room_metadata( - self.call_id, - serde_json::Value::Object(call_metadata).to_string(), - ) - .await; + if let Err(error) = self + .livekit + .set_room_metadata( + self.call_id, + serde_json::Value::Object(call_metadata).to_string(), + ) + .await + { + log::warn!( + "Unable to update administrators for call {}: {}", + self.call_id, + error + ); + } } - pub async fn set_anonymous_joining(&self, enable: bool) { + pub async fn set_anonymous_joining( + &self, + enable: bool, + omega: &OmegaConnection, + omikron_id: u64, + ) { *self.anonymous_joining.write().await = enable; - let call_metadata = match call_util::get_room_metadata(self.call_id).await { + let call_metadata = match self.livekit.get_room_metadata(self.call_id).await { Ok(metadata) => metadata, Err(_) => "{}".to_string(), }; @@ -175,22 +200,27 @@ impl CallGroup { call_metadata.insert("anonymous_joining".to_string(), serde_json::json!(enable)); - let _ = call_util::set_room_metadata( - self.call_id, - serde_json::Value::Object(call_metadata).to_string(), - ) - .await; + if let Err(error) = self + .livekit + .set_room_metadata( + self.call_id, + serde_json::Value::Object(call_metadata).to_string(), + ) + .await + { + log::warn!( + "Unable to update anonymous access for call {}: {}", + self.call_id, + error + ); + } if self.short_link.read().await.is_none() { let long_link = format!( "https://app.tensamin.net/call/anonymous?call_id={}&omikron_id={}", - self.call_id, - env::var("ID") - .unwrap_or("0".to_string()) - .parse::() - .unwrap_or(0), + self.call_id, omikron_id, ); - let response_cv = get_omega_connection() + let response_cv = omega .await_response( &CommunicationValue::new(CommunicationType::ShortenLink) .add_typed_default(DataType::Link, DataValue::Str(long_link)), @@ -198,33 +228,33 @@ impl CallGroup { ) .await; if let Ok(response) = response_cv { - *self.short_link.write().await = Some( - response - .get_data(DataType::Link) - .as_str() - .unwrap() - .to_string(), - ); - log::info!( - "Shortened link for call {} is {}", - self.call_id, - self.short_link.read().await.as_ref().unwrap() - ); + if let Some(link) = response.get_data(DataType::Link).as_str() { + *self.short_link.write().await = Some(link.to_string()); + log::info!("Shortened link for call {}", self.call_id); + } } } } - pub async fn create_anonymous_token(&self, user_id: u64) -> Option { + pub async fn create_anonymous_token(&self, user_id: u64) -> Result, CallError> { if self.is_anonymous().await { - if let Ok(token) = call_util::create_token(user_id, self.call_id, false) { - return Some(token); - } + return self + .livekit + .create_token(user_id, self.call_id, false) + .map(Some); } - None + Ok(None) } pub async fn remove_caller(&self, user_id: u64) { - let _ = call_util::remove_participant(self.call_id, user_id).await; + if let Err(error) = self.livekit.remove_participant(self.call_id, user_id).await { + log::warn!( + "Unable to remove user {} from call {}: {}", + user_id, + self.call_id, + error + ); + } self.members .write() .await @@ -263,19 +293,19 @@ mod tests { let tm = TypeMap::latest(); let mut map: BTreeMap = BTreeMap::new(); map.insert( - DataType::SecretId.to_id(&tm), + data_type_id(DataType::SecretId, &tm), DataValue::Str("call:test:main".to_string()), ); map.insert( - DataType::VersionNumber.to_id(&tm), + data_type_id(DataType::VersionNumber, &tm), DataValue::SignedNumber(1), ); map.insert( - DataType::EncryptedSecret.to_id(&tm), + data_type_id(DataType::EncryptedSecret, &tm), DataValue::Bytes(vec![1, 2, 3]), ); map.insert( - DataType::WrappingScheme.to_id(&tm), + data_type_id(DataType::WrappingScheme, &tm), DataValue::Str("mtp-call-secret-kem-chacha20poly1305-hkdf-sha256-v1".to_string()), ); diff --git a/src/calls/call_manager.rs b/src/calls/call_manager.rs index 19dad95..93e5789 100644 --- a/src/calls/call_manager.rs +++ b/src/calls/call_manager.rs @@ -1,106 +1,151 @@ use dashmap::DashMap; -use once_cell::sync::Lazy; use std::sync::Arc; use uuid::Uuid; use crate::calls::{ call_group::{CallGroup, CallSecretEnvelope}, - call_util, + call_util::LiveKitService, caller::Caller, + error::CallError, }; -pub static CALL_GROUPS: Lazy>> = Lazy::new(|| DashMap::new()); -#[allow(dead_code)] -pub async fn get_call_invites(user_id: u64) -> Vec> { - let mut callers = Vec::new(); - for (_, cg) in CALL_GROUPS.clone().into_iter() { - let members = cg.members.read().await; - for member in members.iter() { - if member.user_id == user_id { - callers.push(member.clone()); - } +pub struct CallManager { + pub(crate) groups: DashMap>, + pub livekit: Arc, +} + +impl Default for CallManager { + fn default() -> Self { + Self { + groups: DashMap::new(), + livekit: Arc::new(LiveKitService::new(None)), } } - callers } -pub async fn get_call(call_id: Uuid) -> Option> { - if let Some(b) = CALL_GROUPS.get(&call_id) { - Some(b.clone()) - } else { - None +impl CallManager { + pub fn new(livekit: Arc) -> Self { + Self { + groups: DashMap::new(), + livekit, + } } -} -pub async fn get_call_groups(user_id: u64) -> Vec> { - let mut call_groups = Vec::new(); - for (_, cg) in CALL_GROUPS.clone().into_iter() { - let is_member = { + #[allow(dead_code)] + pub async fn get_call_invites(&self, user_id: u64) -> Vec> { + let mut callers = Vec::new(); + let call_groups = self + .groups + .iter() + .map(|entry| entry.value().clone()) + .collect::>(); + for cg in call_groups { let members = cg.members.read().await; - members.iter().any(|m| m.user_id == user_id) - }; - - if is_member { - call_groups.push(cg.clone()); + for member in members.iter() { + if member.user_id == user_id { + callers.push(member.clone()); + } + } } - } - call_groups -} - -pub async fn get_call_token(user_id: u64, call_id: Uuid) -> Option { - if let Some(cg) = CALL_GROUPS.get(&call_id) { - let mut members = cg.members.write().await; - - if let Some(member) = members.iter().find(|m| m.user_id == user_id) { - return Some(member.create_token()); - } - - let new_caller = Arc::new(Caller::new(user_id, call_id, false)); - let token = new_caller.create_token(); - - members.push(new_caller); - - return Some(token); + callers } - let caller = Arc::new(Caller::new(user_id, call_id, true)); - let call_group = Arc::new(CallGroup::new(call_id, caller.clone())); + pub async fn get_call(&self, call_id: Uuid) -> Option> { + if let Some(b) = self.groups.get(&call_id) { + Some(b.clone()) + } else { + None + } + } - CALL_GROUPS.insert(call_id, call_group.clone()); + pub async fn get_call_groups(&self, user_id: u64) -> Vec> { + let mut call_groups = Vec::new(); + let tracked_groups = self + .groups + .iter() + .map(|entry| entry.value().clone()) + .collect::>(); + for cg in tracked_groups { + let is_member = { + let members = cg.members.read().await; + members.iter().any(|m| m.user_id == user_id) + }; - let _ = call_util::create_room(call_id).await; - call_group.update_admins().await; + if is_member { + call_groups.push(cg.clone()); + } + } + call_groups + } - Some(caller.create_token()) -} + pub async fn get_call_token(&self, user_id: u64, call_id: Uuid) -> Result { + if let Some(cg) = self.groups.get(&call_id) { + let mut members = cg.members.write().await; -pub async fn add_invite( - call_id: Uuid, - inviter_id: u64, - invitee_id: u64, - secret: CallSecretEnvelope, -) -> bool { - if let Some(cg) = CALL_GROUPS.get(&call_id) { - let mut members = cg.members.write().await; - - let is_inviter_member = members.iter().any(|m| m.user_id == inviter_id); - - if is_inviter_member { - if !members.iter().any(|m| m.user_id == invitee_id) { - members.push(Arc::new(Caller::new(invitee_id, call_id, false))); + if let Some(member) = members.iter().find(|m| m.user_id == user_id) { + return self.livekit.create_token( + member.user_id, + member.call_id, + member.has_admin(), + ); } - let mut secrets = cg.secrets.write().await; - secrets.insert(invitee_id, secret); + let new_caller = Arc::new(Caller::new(user_id, call_id, false)); + let token = self.livekit.create_token( + new_caller.user_id, + new_caller.call_id, + new_caller.has_admin(), + )?; - return true; + members.push(new_caller); + + return Ok(token); } - } - false -} -pub fn should_forward_invite(inviter_id: u64, invitee_id: u64) -> bool { - inviter_id != invitee_id + let caller = Arc::new(Caller::new(user_id, call_id, true)); + let call_group = Arc::new(CallGroup::new_with_service( + call_id, + caller.clone(), + self.livekit.clone(), + )); + + self.livekit.create_room(call_id).await?; + self.groups.insert(call_id, call_group.clone()); + call_group.update_admins().await; + + self.livekit + .create_token(caller.user_id, caller.call_id, caller.has_admin()) + } + + pub async fn add_invite( + &self, + call_id: Uuid, + inviter_id: u64, + invitee_id: u64, + secret: CallSecretEnvelope, + ) -> bool { + if let Some(cg) = self.groups.get(&call_id) { + let mut members = cg.members.write().await; + + let is_inviter_member = members.iter().any(|m| m.user_id == inviter_id); + + if is_inviter_member { + if !members.iter().any(|m| m.user_id == invitee_id) { + members.push(Arc::new(Caller::new(invitee_id, call_id, false))); + } + + let mut secrets = cg.secrets.write().await; + secrets.insert(invitee_id, secret); + + return true; + } + } + false + } + + pub fn should_forward_invite(&self, inviter_id: u64, invitee_id: u64) -> bool { + inviter_id != invitee_id + } } #[cfg(test)] @@ -126,19 +171,22 @@ mod tests { call_id, Arc::new(Caller::new(sender_id, call_id, true)), )); - CALL_GROUPS.insert(call_id, group.clone()); + let manager = CallManager::default(); + manager.groups.insert(call_id, group.clone()); let receiver_secret = envelope("receiver"); - assert!(add_invite(call_id, sender_id, receiver_id, receiver_secret.clone(),).await); + assert!( + manager + .add_invite(call_id, sender_id, receiver_id, receiver_secret.clone(),) + .await + ); assert_eq!( group.get_secret_for_user(receiver_id).await, Some(receiver_secret.clone()) ); assert_eq!(group.get_secret_for_user(sender_id).await, None); - - CALL_GROUPS.remove(&call_id); } #[tokio::test] @@ -149,11 +197,16 @@ mod tests { call_id, Arc::new(Caller::new(sender_id, call_id, true)), )); - CALL_GROUPS.insert(call_id, group.clone()); + let manager = CallManager::default(); + manager.groups.insert(call_id, group.clone()); let secret = envelope("self"); - assert!(add_invite(call_id, sender_id, sender_id, secret.clone()).await); + assert!( + manager + .add_invite(call_id, sender_id, sender_id, secret.clone()) + .await + ); assert_eq!(group.get_secret_for_user(sender_id).await, Some(secret)); assert_eq!( @@ -166,13 +219,12 @@ mod tests { .count(), 1 ); - - CALL_GROUPS.remove(&call_id); } #[test] fn self_invites_are_not_forwarded() { - assert!(!should_forward_invite(44, 44)); - assert!(should_forward_invite(44, 55)); + let manager = CallManager::default(); + assert!(!manager.should_forward_invite(44, 44)); + assert!(manager.should_forward_invite(44, 55)); } } diff --git a/src/calls/call_util.rs b/src/calls/call_util.rs index eb976f0..8476c8c 100644 --- a/src/calls/call_util.rs +++ b/src/calls/call_util.rs @@ -1,139 +1,175 @@ -use livekit_api::services::room::CreateRoomOptions; +use std::{str::FromStr, sync::Arc, time::Duration}; + use livekit_api::{ access_token::{self}, - services::room::RoomClient, + services::room::{CreateRoomOptions, RoomClient}, }; use livekit_protocol::Room; -use std::env; -use std::str::FromStr; -use std::time::Duration; use uuid::Uuid; -use crate::{calls::call_manager::CALL_GROUPS, log, log_err, util::logger::PrintType}; +use crate::{ + calls::{call_manager::CallManager, error::CallError}, + config::LiveKitConfig, + log, log_err, + util::logger::PrintType, +}; -pub fn get_livekit() -> Result<(String, String, String), ()> { - let hostname = match env::var("LIVEKI_HOSTNAME") { - Ok(secret) => secret, - Err(_) => { - log_err!(0, PrintType::General, "LIVEKI_HOSTNAME not set!"); - return Err(()); - } - }; - let api_key = match env::var("LIVEKIT_API_KEY") { - Ok(key) => key, - Err(_) => { - log_err!(0, PrintType::General, "LIVEKIT_API_KEY not set!"); - return Err(()); - } - }; - let api_secret = match env::var("LIVEKIT_API_SECRET") { - Ok(secret) => secret, - Err(_) => { - log_err!(0, PrintType::General, "LIVEKIT_API_SECRET not set!"); - return Err(()); - } - }; - Ok((hostname, api_key, api_secret)) +const LIVEKIT_REQUEST_TIMEOUT: Duration = Duration::from_secs(10); + +pub struct LiveKitService { + config: Option, } -pub async fn create_room(call_id: Uuid) -> Result<(), ()> { - let (hostname, api_key, api_secret) = get_livekit()?; - let room_service = RoomClient::with_api_key(&hostname, &api_key, &api_secret); +impl LiveKitService { + pub fn new(config: Option) -> Self { + Self { config } + } - let options = CreateRoomOptions::default(); + fn livekit_config(&self) -> Result<&LiveKitConfig, CallError> { + self.config.as_ref().ok_or(CallError::NotConfigured) + } - room_service - .create_room(&call_id.to_string(), options) + fn room_client(&self) -> Result { + let config = self.livekit_config()?; + Ok(RoomClient::with_api_key( + &config.hostname, + &config.api_key, + &config.api_secret, + )) + } + + pub async fn create_room(&self, call_id: Uuid) -> Result<(), CallError> { + tokio::time::timeout( + LIVEKIT_REQUEST_TIMEOUT, + self.room_client()? + .create_room(&call_id.to_string(), CreateRoomOptions::default()), + ) .await - .map_err(|_| ())?; - - Ok(()) -} - -pub fn create_token(user_id: u64, call_id: Uuid, has_admin: bool) -> Result { - let (_, api_key, api_secret) = get_livekit()?; - - let token = access_token::AccessToken::with_api_key(&api_key, &api_secret) - .with_identity(&user_id.to_string()) - .with_grants(access_token::VideoGrants { - room_join: true, - can_update_own_metadata: true, - room_admin: has_admin, - room: call_id.to_string(), - ..Default::default() - }) - .to_jwt(); - if let Ok(token) = token { - Ok(token) - } else { - Err(()) + .map_err(|_| CallError::RequestTimedOut { call_id })? + .map_err(|error| CallError::RoomCreationFailed { + call_id, + detail: error.to_string(), + })?; + Ok(()) } -} -#[allow(dead_code)] -pub async fn get_room(call_id: Uuid) -> Result<(RoomClient, Room), ()> { - if let Ok((hostname, api_key, api_secret)) = get_livekit() { - let room_service = RoomClient::with_api_key(&hostname, &api_key, &api_secret); - let rooms = room_service.list_rooms(Vec::new()).await; - if let Ok(rooms) = rooms { - for room in rooms { - if room.name == call_id.to_string() { - return Ok((room_service, room)); + + pub fn create_token( + &self, + user_id: u64, + call_id: Uuid, + has_admin: bool, + ) -> Result { + let config = self.livekit_config()?; + access_token::AccessToken::with_api_key(&config.api_key, &config.api_secret) + .with_identity(&user_id.to_string()) + .with_grants(access_token::VideoGrants { + room_join: true, + can_update_own_metadata: true, + room_admin: has_admin, + room: call_id.to_string(), + ..Default::default() + }) + .to_jwt() + .map_err(|error| CallError::TokenCreationFailed { + detail: error.to_string(), + }) + } + + #[allow(dead_code)] + pub async fn get_room(&self, call_id: Uuid) -> Result<(RoomClient, Room), CallError> { + let room_service = self.room_client()?; + let rooms = + tokio::time::timeout(LIVEKIT_REQUEST_TIMEOUT, room_service.list_rooms(Vec::new())) + .await + .map_err(|_| CallError::RequestTimedOut { call_id })? + .map_err(|error| CallError::RoomLookupFailed { + detail: error.to_string(), + })?; + + rooms + .into_iter() + .find(|room| room.name == call_id.to_string()) + .map(|room| (room_service, room)) + .ok_or(CallError::RoomNotFound { call_id }) + } + + pub async fn remove_participant(&self, call_id: Uuid, user_id: u64) -> Result<(), CallError> { + tokio::time::timeout( + LIVEKIT_REQUEST_TIMEOUT, + self.room_client()? + .remove_participant(&call_id.to_string(), &user_id.to_string()), + ) + .await + .map_err(|_| CallError::RequestTimedOut { call_id })? + .map_err(|error| CallError::ParticipantRemovalFailed { + call_id, + user_id, + detail: error.to_string(), + })?; + Ok(()) + } + + #[allow(dead_code)] + pub async fn get_room_metadata(&self, call_id: Uuid) -> Result { + let (_, room) = self.get_room(call_id).await?; + Ok(room.metadata) + } + + pub async fn set_room_metadata( + &self, + call_id: Uuid, + metadata: String, + ) -> Result<(), CallError> { + tokio::time::timeout( + LIVEKIT_REQUEST_TIMEOUT, + self.room_client()? + .update_room_metadata(&call_id.to_string(), &metadata), + ) + .await + .map_err(|_| CallError::RequestTimedOut { call_id })? + .map_err(|error| CallError::MetadataUpdateFailed { + call_id, + detail: error.to_string(), + })?; + Ok(()) + } + + pub fn garbage_collect_calls(self: Arc, manager: Arc) { + tokio::spawn(async move { + loop { + match self.room_client() { + Ok(room_service) => clean_calls(&manager, room_service).await, + Err(CallError::NotConfigured) => return, + Err(error) => log_err!(0, PrintType::Call, "Call cleanup skipped: {error}"), } + tokio::time::sleep(Duration::from_secs(2)).await; } - } + }); } - return Err(()); } -pub async fn remove_participant(call_id: Uuid, user_id: u64) -> Result<(), ()> { - if let Ok((hostname, api_key, api_secret)) = get_livekit() { - let room_service = RoomClient::with_api_key(&hostname, &api_key, &api_secret); - if let Ok(_) = room_service - .remove_participant(&call_id.to_string(), &user_id.to_string()) +pub async fn clean_calls(manager: &CallManager, room_service: RoomClient) { + let rooms = + match tokio::time::timeout(LIVEKIT_REQUEST_TIMEOUT, room_service.list_rooms(Vec::new())) .await { - return Ok(()); - } - } - return Err(()); -} - -#[allow(dead_code)] -pub async fn get_room_metadata(call_id: Uuid) -> Result { - if let Ok((_, room)) = get_room(call_id).await { - Ok(room.metadata) - } else { - Err(()) - } -} - -pub async fn set_room_metadata(call_id: Uuid, metadata: String) -> Result<(), ()> { - let (hostname, api_key, api_secret) = get_livekit()?; - let room_service = RoomClient::with_api_key(&hostname, &api_key, &api_secret); - - room_service - .update_room_metadata(&call_id.to_string(), &metadata) - .await - .map_err(|_| ())?; - - Ok(()) -} - -pub fn garbage_collect_calls() { - tokio::spawn(async move { - loop { - if let Ok((hostname, api_key, api_secret)) = get_livekit() { - let room_service = RoomClient::with_api_key(&hostname, &api_key, &api_secret); - clean_calls(room_service).await; + Ok(Ok(rooms)) => rooms, + Ok(Err(error)) => { + log_err!( + 0, + PrintType::Call, + "Unable to list LiveKit rooms during cleanup: {error}" + ); + return; } - tokio::time::sleep(Duration::from_secs(2)).await; - } - }); -} -pub async fn clean_calls(room_service: RoomClient) { - let rooms = room_service.list_rooms(Vec::new()).await.unwrap(); - let mut call_ids: Vec = Vec::new(); - let mut no_users: Vec = Vec::new(); + Err(_) => { + log_err!(0, PrintType::Call, "LiveKit room cleanup timed out"); + return; + } + }; + + let mut call_ids = Vec::new(); + let mut no_users = Vec::new(); for room in rooms { if let Ok(id) = Uuid::from_str(&room.name) { if room.num_participants == 0 { @@ -142,23 +178,35 @@ pub async fn clean_calls(room_service: RoomClient) { call_ids.push(id); } } - let size_pre = CALL_GROUPS.len(); - for (id, _) in CALL_GROUPS.clone().into_iter() { + + let size_pre = manager.groups.len(); + let tracked_calls = manager + .groups + .iter() + .map(|entry| *entry.key()) + .collect::>(); + for id in tracked_calls { if !call_ids.contains(&id) { - CALL_GROUPS.remove(&id); + manager.groups.remove(&id); } } - for (_, cg) in CALL_GROUPS.clone().into_iter() { - *cg.show.write().await = !no_users.contains(&cg.call_id); + + let call_groups = manager + .groups + .iter() + .map(|entry| entry.value().clone()) + .collect::>(); + for call_group in call_groups { + *call_group.show.write().await = !no_users.contains(&call_group.call_id); } - let size_post = CALL_GROUPS.len(); - if size_pre - size_post != 0 { + let size_post = manager.groups.len(); + if size_pre != size_post { log!( 0, PrintType::Call, "Cleaned {} calls, {} remaining", - size_pre - size_post, + size_pre.saturating_sub(size_post), size_post ); } diff --git a/src/calls/caller.rs b/src/calls/caller.rs index 74e0c3b..6a69566 100644 --- a/src/calls/caller.rs +++ b/src/calls/caller.rs @@ -3,7 +3,7 @@ use std::time::{SystemTime, UNIX_EPOCH}; use tokio::sync::RwLock; use uuid::Uuid; -use crate::calls::call_util; +use crate::calls::{call_util::LiveKitService, error::CallError}; pub struct Caller { pub user_id: u64, @@ -39,11 +39,7 @@ impl Caller { pub async fn set_timeout(&self, timeout: i64) { *self.timeout.write().await = timeout; } - pub fn create_token(&self) -> String { - if let Ok(token) = call_util::create_token(self.user_id, self.call_id, self.has_admin()) { - token - } else { - String::new() - } + pub fn create_token(&self, livekit: &LiveKitService) -> Result { + livekit.create_token(self.user_id, self.call_id, self.has_admin()) } } diff --git a/src/calls/error.rs b/src/calls/error.rs new file mode 100644 index 0000000..ddae470 --- /dev/null +++ b/src/calls/error.rs @@ -0,0 +1,25 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum CallError { + #[error("LiveKit is not configured")] + NotConfigured, + #[error("LiveKit request for room {call_id} timed out")] + RequestTimedOut { call_id: uuid::Uuid }, + #[error("failed to create room {call_id}: {detail}")] + RoomCreationFailed { call_id: uuid::Uuid, detail: String }, + #[error("room {call_id} was not found")] + RoomNotFound { call_id: uuid::Uuid }, + #[error("failed to list rooms: {detail}")] + RoomLookupFailed { detail: String }, + #[error("failed to remove participant {user_id} from room {call_id}: {detail}")] + ParticipantRemovalFailed { + call_id: uuid::Uuid, + user_id: u64, + detail: String, + }, + #[error("failed to update room {call_id} metadata: {detail}")] + MetadataUpdateFailed { call_id: uuid::Uuid, detail: String }, + #[error("failed to create access token: {detail}")] + TokenCreationFailed { detail: String }, +} diff --git a/src/calls/mod.rs b/src/calls/mod.rs index 8413354..1ba3975 100644 --- a/src/calls/mod.rs +++ b/src/calls/mod.rs @@ -2,3 +2,4 @@ pub mod call_group; pub mod call_manager; pub mod call_util; pub mod caller; +pub mod error; diff --git a/src/config.rs b/src/config.rs new file mode 100644 index 0000000..7c10773 --- /dev/null +++ b/src/config.rs @@ -0,0 +1,133 @@ +use std::env; + +use thiserror::Error; + +const DEFAULT_RHO_PORT: u16 = 443; +const DEFAULT_OMEGA_HOST: &str = "tensamin.net"; +const DEFAULT_OMEGA_PORT: u16 = 9187; + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct LiveKitConfig { + pub hostname: String, + pub api_key: String, + pub api_secret: String, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub struct Config { + pub rho_port: u16, + pub omega_host: String, + pub omega_port: u16, + pub omikron_id: u64, + pub livekit: Option, +} + +#[derive(Debug, Error, Eq, PartialEq)] +pub enum ConfigError { + #[error("{name} must be a valid {kind}")] + InvalidValue { + name: &'static str, + kind: &'static str, + }, + #[error("LIVEKIT_HOSTNAME, LIVEKIT_API_KEY, and LIVEKIT_API_SECRET must be set together")] + IncompleteLiveKitCredentials, +} + +impl Config { + pub fn from_environment() -> Result { + let rho_port = parse_or_default("RHO_PORT", DEFAULT_RHO_PORT)?; + let omega_port = parse_or_default("OMEGA_PORT", DEFAULT_OMEGA_PORT)?; + let omikron_id = parse_or_default("ID", 0_u64)?; + let omega_host = env::var("OMEGA_HOST") + .unwrap_or_else(|_| DEFAULT_OMEGA_HOST.to_string()) + .trim() + .to_string(); + + if omega_host.is_empty() { + return Err(ConfigError::InvalidValue { + name: "OMEGA_HOST", + kind: "non-empty host name", + }); + } + + Ok(Self { + rho_port, + omega_host, + omega_port, + omikron_id, + livekit: livekit_from_environment()?, + }) + } +} + +fn parse_or_default(name: &'static str, default: T) -> Result +where + T: std::str::FromStr, +{ + match env::var(name) { + Ok(value) => value.trim().parse().map_err(|_| ConfigError::InvalidValue { + name, + kind: "number", + }), + Err(_) => Ok(default), + } +} + +fn livekit_from_environment() -> Result, ConfigError> { + livekit_from_values( + env::var("LIVEKIT_HOSTNAME").ok(), + env::var("LIVEKIT_API_KEY").ok(), + env::var("LIVEKIT_API_SECRET").ok(), + ) +} + +fn livekit_from_values( + hostname: Option, + api_key: Option, + api_secret: Option, +) -> Result, ConfigError> { + match (hostname, api_key, api_secret) { + (None, None, None) => Ok(None), + (Some(hostname), Some(api_key), Some(api_secret)) + if !hostname.trim().is_empty() + && !api_key.trim().is_empty() + && !api_secret.trim().is_empty() => + { + Ok(Some(LiveKitConfig { + hostname, + api_key, + api_secret, + })) + } + _ => Err(ConfigError::IncompleteLiveKitCredentials), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn livekit_credentials_must_be_complete() { + assert_eq!( + livekit_from_values(Some("host".to_string()), Some("key".to_string()), None), + Err(ConfigError::IncompleteLiveKitCredentials) + ); + } + + #[test] + fn complete_livekit_credentials_are_loaded() { + assert_eq!( + livekit_from_values( + Some("https://livekit.example".to_string()), + Some("key".to_string()), + Some("secret".to_string()), + ), + Ok(Some(LiveKitConfig { + hostname: "https://livekit.example".to_string(), + api_key: "key".to_string(), + api_secret: "secret".to_string(), + })) + ); + } +} diff --git a/src/main.rs b/src/main.rs index 19e1c10..760e37c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,11 +1,13 @@ mod anonymous_clients; +mod app_state; mod calls; +mod config; mod data; mod omega; mod rho; +mod services; mod util; -use std::env; use std::path::PathBuf; use dotenv::dotenv; @@ -19,29 +21,30 @@ use mtp::crypto::Keyring; use mtp::files::{load_keyring_raw, save_keyring_raw, save_public_key_bundle}; use crate::{ - calls::call_util::garbage_collect_calls, omega::omega_connection::get_omega_connection, - rho::server::start, util::logger::startup, + app_state::AppState, + calls::{call_manager::CallManager, call_util::LiveKitService}, + config::Config, + omega::omega_connection::{OmegaConnection, start_task_cleanup_loop}, + rho::rho_manager::RhoManager, + rho::server::start, + util::logger::startup, }; const KEYRING_PATH: &str = "./omikron.mk"; const PUBLIC_KEY_PATH: &str = "./omikron.mpkb"; -static KEYRING: Lazy = Lazy::new(|| { - load_keyring_raw(KEYRING_PATH).unwrap_or_else(|_| { - let kr = Keyring::generate(); - save_keyring_raw(&kr, KEYRING_PATH).expect("Failed to save generated keyring"); - save_public_key_bundle(&kr.public_key_bundle(), PUBLIC_KEY_PATH) - .expect("Failed to save generated public key bundle"); - eprintln!("Generated new keyring at {}", KEYRING_PATH); - kr - }) -}); - -pub fn get_keyring() -> &'static Keyring { - &KEYRING -} -pub fn load_keyring() -> Keyring { - Keyring::from_bytes(&KEYRING.to_bytes()).unwrap() +fn load_keyring() -> Result { + match load_keyring_raw(KEYRING_PATH) { + Ok(keyring) => Ok(keyring), + Err(_) => { + let kr = Keyring::generate(); + save_keyring_raw(&kr, KEYRING_PATH).map_err(|error| error.to_string())?; + save_public_key_bundle(&kr.public_key_bundle(), PUBLIC_KEY_PATH) + .map_err(|error| error.to_string())?; + eprintln!("Generated new keyring at {}", KEYRING_PATH); + Ok(kr) + } + } } #[tokio::main] @@ -53,18 +56,51 @@ async fn main() { dotenv().ok(); startup(); - let rho_port = env::var("RHO_PORT") - .ok() - .and_then(|s| s.parse().ok()) - .unwrap_or(443); + let config = match Config::from_environment() { + Ok(config) => config, + Err(error) => { + eprintln!("Invalid configuration: {error}"); + return; + } + }; - get_omega_connection(); + let keyring = match load_keyring() { + Ok(keyring) => keyring, + Err(error) => { + eprintln!("Unable to load keyring: {error}"); + return; + } + }; + + let omega_keyring = match keyring_for_omega(&keyring) { + Ok(keyring) => keyring, + Err(error) => { + eprintln!("Unable to copy keyring for Omega: {error}"); + return; + } + }; + let rho = std::sync::Arc::new(RhoManager::new()); + let omega = std::sync::Arc::new(OmegaConnection::from_config( + &config, + omega_keyring, + rho.clone(), + )); + omega.clone().start().await; + start_task_cleanup_loop(omega.clone()); + let livekit = std::sync::Arc::new(LiveKitService::new(config.livekit.clone())); + let call_manager = std::sync::Arc::new(CallManager::new(livekit.clone())); + let state = AppState::new(config.clone(), keyring, omega, rho, call_manager, livekit); tokio::spawn(async move { - if let Err(e) = start(rho_port).await { + if let Err(e) = start(state).await { log_err!(0, util::logger::PrintType::General, "{}", e); } }); - garbage_collect_calls(); - tokio::signal::ctrl_c().await.unwrap(); + if let Err(error) = tokio::signal::ctrl_c().await { + eprintln!("Unable to wait for shutdown signal: {error}"); + } +} + +fn keyring_for_omega(keyring: &Keyring) -> Result { + Keyring::from_bytes(&keyring.to_bytes()).map_err(|error| error.to_string()) } diff --git a/src/omega/omega_connection.rs b/src/omega/omega_connection.rs index a54dd28..bb10ed2 100644 --- a/src/omega/omega_connection.rs +++ b/src/omega/omega_connection.rs @@ -1,8 +1,6 @@ use crate::{ - data::user::UserStatus, - load_keyring, log_cv_in, log_cv_out, log_err, log_in, - rho::rho_manager::{self, RHO_CONNECTIONS, connection_count}, - util::logger::PrintType, + config::Config, data::user::UserStatus, log_cv_in, log_cv_out, log_err, log_in, + rho::rho_manager::RhoManager, util::logger::PrintType, }; use dashmap::DashMap; use mtp::client::{Client, Receiver, Sender}; @@ -13,8 +11,7 @@ use mtp::{ host::{Policy, SendMode}, }; use mtp_transport::ConnectionHandle; -use once_cell::sync::Lazy; -use std::{collections::HashMap, env, sync::Arc, time::Duration}; +use std::{collections::HashMap, sync::Arc, time::Duration}; use tokio::{ sync::{Mutex, RwLock, mpsc, watch}, task::JoinHandle, @@ -23,32 +20,15 @@ use tokio::{ use uuid::Uuid; -// ============================================================================ -// Configuration -// ============================================================================ - -const OMEGA_HOST_DEFAULT: &str = "tensamin.net"; -const OMEGA_PORT_DEFAULT: u16 = 9187; - -fn omega_host_and_port() -> (String, u16) { - let host = env::var("OMEGA_HOST") - .map(|s| s.trim().to_string()) - .unwrap_or_else(|_| OMEGA_HOST_DEFAULT.to_string()); - - let port = env::var("OMEGA_PORT") - .ok() - .and_then(|s| s.trim().parse().ok()) - .unwrap_or(OMEGA_PORT_DEFAULT); - - (host, port) -} - const RECONNECT_DELAY: Duration = Duration::from_secs(5); const MAX_RECONNECT_DELAY: Duration = Duration::from_secs(300); const CONNECTION_TIMEOUT: Duration = Duration::from_secs(10); const HEARTBEAT_INTERVAL: Duration = Duration::from_secs(5); const TASK_CLEANUP_INTERVAL: Duration = Duration::from_secs(60); const TASK_MAX_AGE: Duration = Duration::from_secs(60); +const MAX_CONCURRENT_REQUESTS: usize = 128; +const CIRCUIT_BREAKER_FAILURE_THRESHOLD: u32 = 3; +const CIRCUIT_BREAKER_COOLDOWN: Duration = Duration::from_secs(30); // ============================================================================ // Waiting Task System @@ -59,13 +39,40 @@ pub struct WaitingTask { pub inserted_at: Instant, } -pub static WAITING_TASKS: Lazy> = Lazy::new(DashMap::new); +struct CircuitBreaker { + consecutive_failures: u32, + open_until: Option, +} -pub fn start_task_cleanup_loop() { - tokio::spawn(async { +impl CircuitBreaker { + fn allow_request(&mut self) -> bool { + if self.open_until.is_some_and(|until| until > Instant::now()) { + return false; + } + self.open_until = None; + true + } + + fn record_success(&mut self) { + self.consecutive_failures = 0; + self.open_until = None; + } + + fn record_failure(&mut self) { + self.consecutive_failures += 1; + if self.consecutive_failures >= CIRCUIT_BREAKER_FAILURE_THRESHOLD { + self.open_until = Some(Instant::now() + CIRCUIT_BREAKER_COOLDOWN); + } + } +} + +pub fn start_task_cleanup_loop(omega: Arc) { + tokio::spawn(async move { loop { sleep(TASK_CLEANUP_INTERVAL).await; - WAITING_TASKS.retain(|_, v| v.inserted_at.elapsed() < TASK_MAX_AGE); + omega + .waiting_tasks + .retain(|_, v| v.inserted_at.elapsed() < TASK_MAX_AGE); } }); } @@ -109,6 +116,7 @@ pub struct OmegaConnection { connection_loop_handle: Arc>>>, host: String, port: u16, + omikron_id: u64, last_ping: Arc>, heartbeat_handle: Arc>>>, message_send_times: Arc>>, @@ -116,31 +124,42 @@ pub struct OmegaConnection { shutdown_tx: Arc>>>, // Track if we should reconnect on close reconnect_on_close: Arc>, + request_slots: tokio::sync::Semaphore, + circuit_breaker: Mutex, + keyring: mtp::crypto::Keyring, + rho: Arc, + waiting_tasks: DashMap, } impl OmegaConnection { - pub fn new() -> Self { - let (host, port) = omega_host_and_port(); - Self::with_host(&host, port) - } - - // omikron and omega are both served over publicly trusted (fullchain) certs on - // the same domain, so the client uses system root trust - no pinning needed. - pub fn with_host(host: &str, port: u16) -> Self { + pub fn from_config( + config: &Config, + keyring: mtp::crypto::Keyring, + rho: Arc, + ) -> Self { let (shutdown_tx, _) = watch::channel(false); OmegaConnection { state: Arc::new(RwLock::new(ConnectionState::Disconnected)), sender: Arc::new(RwLock::new(None)), connection_loop_handle: Arc::new(Mutex::new(None)), - host: host.to_string(), - port, + host: config.omega_host.clone(), + port: config.omega_port, + omikron_id: config.omikron_id, last_ping: Arc::new(Mutex::new(-1)), heartbeat_handle: Arc::new(Mutex::new(None)), message_send_times: Arc::new(Mutex::new(HashMap::new())), connection_id: Uuid::new_v4(), shutdown_tx: Arc::new(Mutex::new(Some(shutdown_tx))), reconnect_on_close: Arc::new(RwLock::new(true)), + request_slots: tokio::sync::Semaphore::new(MAX_CONCURRENT_REQUESTS), + circuit_breaker: Mutex::new(CircuitBreaker { + consecutive_failures: 0, + open_until: None, + }), + keyring, + rho, + waiting_tasks: DashMap::new(), } } @@ -190,7 +209,7 @@ impl OmegaConnection { // Close sender if connected if let Some(sender) = self.sender.read().await.as_ref() { - sender.close(); + sender.close().await; } *self.state.write().await = ConnectionState::Disconnected; @@ -257,12 +276,7 @@ impl OmegaConnection { *self.state.write().await = ConnectionState::Connecting; let client_config = ClientConfig::new(format!("https://{}:{}", self.host, self.port)) - .with_client_id( - env::var("ID") - .unwrap_or("0".to_string()) - .parse::() - .unwrap_or(0), - ) + .with_client_id(self.omikron_id) .with_policy( Policy::default() .with_send_mode(SendMode::SingleStreamPerMessage) @@ -282,7 +296,7 @@ impl OmegaConnection { let host_public_key = load_public_key_bundle("./omega.mpkb") .map_err(|e| format!("Failed to load omega.mpkb: {}", e))?; - let mut connection = Client::auth_connect(client_config, &load_keyring(), &host_public_key) + let mut connection = Client::auth_connect(client_config, &self.keyring, &host_public_key) .await .map_err(|e| format!("Connection failed: {}", e))?; @@ -349,13 +363,14 @@ impl OmegaConnection { let mut connected_iota_ids: Vec = Vec::new(); let mut connected_user_ids: Vec = Vec::new(); - let rho_connections_reader = RHO_CONNECTIONS.read().await; + let iota_ids = self.rho.iota_ids().await; + let rho_connections = self.rho.connections().await; - for iota_id in rho_connections_reader.keys() { - connected_iota_ids.push(DataValue::SignedNumber((*iota_id).into())); + for iota_id in iota_ids { + connected_iota_ids.push(DataValue::SignedNumber(iota_id.into())); } - for rho in rho_connections_reader.values() { + for rho in rho_connections { for client_conn in rho.get_client_connections().await { connected_user_ids.push(DataValue::SignedNumber( client_conn.get_user_id().await.into(), @@ -363,14 +378,12 @@ impl OmegaConnection { } } - drop(rho_connections_reader); - let sync_msg = CommunicationValue::new(CommunicationType::SyncClientIotaStatus) .add_typed_default(DataType::IotaIds, DataValue::Array(connected_iota_ids)) .add_typed_default(DataType::UserIds, DataValue::Array(connected_user_ids)) .add_typed_default( DataType::RhoConnections, - DataValue::SignedNumber(connection_count().await as i128), + DataValue::SignedNumber(self.rho.connection_count().await as i128), ); self.send_message(&sync_msg).await; @@ -403,7 +416,7 @@ impl OmegaConnection { } let msg_id = cv.get_id(); - if let Some((_, task)) = WAITING_TASKS.remove(&msg_id) { + if let Some((_, task)) = self.waiting_tasks.remove(&msg_id) { if (task.task)(self.clone(), cv.clone()) { continue; } @@ -411,20 +424,20 @@ impl OmegaConnection { if cv.is_type(CommunicationType::IotaUserData) { if let DataValue::Array(users) = cv.get_data(DataType::UserIds) { - let mut user_ids: Vec = Vec::new(); + let mut user_ids = Vec::new(); for value in users { if let DataValue::SignedNumber(user_id) = value { - user_ids.push(user_id.clone() as u64); + if let Ok(user_id) = i64::try_from(*user_id) { + user_ids.push(user_id); + } } } - let connections = crate::rho::rho_manager::RHO_CONNECTIONS.read().await; if let Some(iota_id) = cv.get_data(DataType::IotaId).as_number() { - if let Some(rho) = connections.get(&(iota_id as i64)) { - rho.get_iota_connection().set_user_ids(user_ids).await; - } + self.rho.replace_users_for_iota(iota_id as i64, user_ids).await; } else { - for rho in connections.values() { - rho.get_iota_connection().set_user_ids(user_ids.clone()).await; + let iota_ids = self.rho.iota_ids().await; + for iota_id in iota_ids { + self.rho.replace_users_for_iota(iota_id, user_ids.clone()).await; } } } @@ -516,7 +529,7 @@ impl OmegaConnection { drop(sender_guard); // Trigger reconnection by closing the connection state if let Some(sender) = self.sender.write().await.take() { - sender.close(); + sender.close().await; } return; } @@ -560,13 +573,38 @@ impl OmegaConnection { &self, cv: &CommunicationValue, timeout_duration: Option, + ) -> Result { + let _permit = self + .request_slots + .acquire() + .await + .map_err(|_| "Omega request queue is unavailable".to_string())?; + + if !self.circuit_breaker.lock().await.allow_request() { + return Err("Omega circuit breaker is open".to_string()); + } + + let result = self.await_response_inner(cv, timeout_duration).await; + let mut breaker = self.circuit_breaker.lock().await; + if result.is_ok() { + breaker.record_success(); + } else { + breaker.record_failure(); + } + result + } + + async fn await_response_inner( + &self, + cv: &CommunicationValue, + timeout_duration: Option, ) -> Result { self.await_connection(timeout_duration).await?; let (tx, mut rx) = mpsc::channel(1); let msg_id = cv.get_id(); - WAITING_TASKS.insert( + self.waiting_tasks.insert( msg_id, WaitingTask { task: Box::new(move |_, response_cv| { @@ -588,7 +626,7 @@ impl OmegaConnection { Ok(Some(response_cv)) => Ok(response_cv), Ok(_) => Err("Channel closed".to_string()), Err(_) => { - WAITING_TASKS.remove(&msg_id); + self.waiting_tasks.remove(&msg_id); Err("Request timed out".to_string()) } } @@ -604,13 +642,13 @@ impl OmegaConnection { } #[allow(dead_code)] - pub async fn close_iota(iota_id: i64) { + pub async fn close_iota(&self, iota_id: i64) { let cv = CommunicationValue::new(CommunicationType::IotaDisconnected) .add_typed_default(DataType::IotaId, DataValue::SignedNumber(iota_id.into())); - OMEGA_CONNECTION.send_message(&cv).await; + self.send_message(&cv).await; } - pub async fn client_changed(_iota_id: i64, user_id: i64, state: UserStatus) { + pub async fn client_changed(&self, _iota_id: i64, user_id: i64, state: UserStatus) { let msg_type = match state { UserStatus::iota_offline => CommunicationType::UserDisconnected, UserStatus::user_offline => CommunicationType::UserDisconnected, @@ -621,10 +659,10 @@ impl OmegaConnection { let cv = CommunicationValue::new(msg_type) .add_typed_default(DataType::UserId, DataValue::SignedNumber(user_id.into())) .add_typed_default(DataType::UserState, DataValue::Str(state.to_string())); - OMEGA_CONNECTION.send_message(&cv).await; + self.send_message(&cv).await; } - pub async fn user_states(user_id: i64, user_ids: Vec) { + pub async fn user_states(&self, user_id: i64, user_ids: Vec) { let user_ids = user_ids .iter() .map(|v| DataValue::SignedNumber((*v).into())) @@ -634,13 +672,15 @@ impl OmegaConnection { .add_typed_default(DataType::UserIds, DataValue::Array(user_ids)); let msg_id = cv.get_id(); - WAITING_TASKS.insert( + let rho_manager = self.rho.clone(); + self.waiting_tasks.insert( msg_id, WaitingTask { task: Box::new( move |_: Arc, response: CommunicationValue| { + let rho_manager = rho_manager.clone(); tokio::spawn(async move { - let rho = rho_manager::get_rho_con_for_user(user_id).await; + let rho = rho_manager.get_for_user(user_id).await; if let Some(rho) = rho { for client in rho.get_client_connections_for_user(user_id).await { client.send_message(&response).await; @@ -654,28 +694,6 @@ impl OmegaConnection { }, ); - OMEGA_CONNECTION.send_message(&cv).await; + self.send_message(&cv).await; } } - -// ============================================================================ -// Global Instance -// ============================================================================ - -static OMEGA_CONNECTION: Lazy> = Lazy::new(|| { - let conn = Arc::new(OmegaConnection::new()); - - // Start the connection manager immediately - let conn_clone = conn.clone(); - tokio::spawn(async move { - conn_clone.start().await; - }); - - start_task_cleanup_loop(); - - conn -}); - -pub fn get_omega_connection() -> Arc { - OMEGA_CONNECTION.clone() -} diff --git a/src/rho/app_connection.rs b/src/rho/app_connection.rs index a4b6285..855104b 100644 --- a/src/rho/app_connection.rs +++ b/src/rho/app_connection.rs @@ -1,7 +1,7 @@ use crate::anonymous_clients::anonymous_manager; -use crate::omega::omega_connection::get_omega_connection; +use crate::app_state::AppState; use crate::rho::connection::{GeneralConnection, MtpReceiver, MtpSender}; -use crate::rho::{rho_connection::RhoConnection, rho_manager}; +use crate::rho::rho_connection::RhoConnection; use crate::util::logger::PrintType; use crate::{log_cv_in, log_cv_out, log_err, log_in, log_out}; use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue}; @@ -11,6 +11,7 @@ use tokio::sync::RwLock; use uuid::Uuid; pub struct AppConnection { + pub state: Arc, pub user_id: u64, pub app_identifier: String, pub app_session: Uuid, @@ -27,6 +28,7 @@ pub struct AppConnection { impl AppConnection { pub async fn from_general(general: Arc, user_id: u64) -> Arc { Arc::new(Self { + state: general.state.clone(), ping: Arc::new(RwLock::new(0)), pub_key: Arc::new(RwLock::new(None)), rho_connection: general.rho_connection.clone(), @@ -154,7 +156,10 @@ impl AppConnection { async fn handle_omega_forward(self: Arc, cv: CommunicationValue) { let app_for_closure = self.clone(); tokio::spawn(async move { - let response_cv = get_omega_connection() + let response_cv = self + .state + .omega + .clone() .await_response(&cv.with_sender(self.user_id), Some(Duration::from_secs(20))) .await; if let Ok(response_cv) = response_cv { @@ -169,7 +174,7 @@ impl AppConnection { if let DataValue::SignedNumber(last_ping) = cv.get_data(DataType::LastPing) { let current = SystemTime::now() .duration_since(UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_millis(); let mut ping_guard = self.ping.write().await; *ping_guard = (current as i128 - *last_ping) as i64; @@ -227,7 +232,10 @@ impl AppConnection { return; } - let load_uuid_response = get_omega_connection() + let load_uuid_response = self + .state + .omega + .clone() .await_response( &CommunicationValue::new(CommunicationType::GetUserData) .with_id(cv.clone().get_id()) @@ -328,7 +336,7 @@ impl AppConnection { /// Handle connection close pub async fn handle_close(&self) { let user_id = self.get_user_id().await; - if let Some(rho_conn) = rho_manager::get_rho_con_for_user(user_id as i64).await { + if let Some(rho_conn) = self.state.rho.get_for_user(user_id as i64).await { rho_conn.close_app_connection(Arc::new(self.clone())).await; } } @@ -338,6 +346,7 @@ impl AppConnection { impl Clone for AppConnection { fn clone(&self) -> Self { Self { + state: self.state.clone(), sender: Arc::clone(&self.sender), receiver: Arc::clone(&self.receiver), user_id: self.user_id, diff --git a/src/rho/client_connection.rs b/src/rho/client_connection.rs index d08f679..615885d 100644 --- a/src/rho/client_connection.rs +++ b/src/rho/client_connection.rs @@ -1,10 +1,10 @@ use crate::anonymous_clients::anonymous_manager; -use crate::calls::{call_group::call_invite_secret_from_cv, call_manager, call_util}; -use crate::omega::omega_connection::get_omega_connection; +use crate::app_state::AppState; +use crate::calls::call_group::call_invite_secret_from_cv; +use crate::data::user::UserStatus; use crate::rho::connection::{GeneralConnection, MtpReceiver, MtpSender}; -use crate::rho::{rho_connection::RhoConnection, rho_manager}; +use crate::rho::rho_connection::RhoConnection; use crate::util::logger::PrintType; -use crate::{data::user::UserStatus, omega::omega_connection::OmegaConnection}; use crate::{log_cv_in, log_cv_out, log_err, log_in, log_out}; use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue}; use std::str::FromStr; @@ -15,6 +15,7 @@ use trust_dns_resolver::TokioAsyncResolver; use uuid::Uuid; pub struct ClientConnection { + pub state: Arc, pub user_id: u64, pub session_id: u64, pub client_version: String, @@ -26,11 +27,13 @@ pub struct ClientConnection { pub rho_connection: Arc>>>, pub interested_users: Arc>>, is_open: Arc>, + message_slots: Arc, } impl ClientConnection { pub async fn from_general(general: Arc, user_id: u64) -> Arc { Arc::new(Self { + state: general.state.clone(), ping: Arc::new(RwLock::new(0)), pub_key: Arc::new(RwLock::new(None)), rho_connection: general.rho_connection.clone(), @@ -41,6 +44,7 @@ impl ClientConnection { user_id: user_id, session_id: general.session_id.read().await.clone(), client_version: general.client_version.read().await.clone(), + message_slots: Arc::new(tokio::sync::Semaphore::new(32)), }) } pub fn start(self: Arc) { @@ -96,7 +100,11 @@ impl ClientConnection { /// Handle incoming message from client pub async fn handle_message(self: Arc, cv: CommunicationValue) { + let Ok(permit) = self.message_slots.clone().acquire_owned().await else { + return; + }; tokio::spawn(async move { + let _permit = permit; if cv.is_type(CommunicationType::Ping) { self.handle_ping(cv).await; return; @@ -270,14 +278,15 @@ impl ClientConnection { } async fn handle_omega_forward(self: Arc, cv: CommunicationValue) { let client_for_closure = self.clone(); - tokio::spawn(async move { - let response_cv = get_omega_connection() - .await_response(&cv.with_sender(self.user_id), Some(Duration::from_secs(20))) - .await; - if let Ok(response_cv) = response_cv { - client_for_closure.send_message(&response_cv).await; - } - }); + let response_cv = self + .state + .omega + .clone() + .await_response(&cv.with_sender(self.user_id), Some(Duration::from_secs(20))) + .await; + if let Ok(response_cv) = response_cv { + client_for_closure.send_message(&response_cv).await; + } } /// Handle ping message @@ -286,7 +295,7 @@ impl ClientConnection { if let DataValue::SignedNumber(last_ping) = cv.get_data(DataType::LastPing) { let current = SystemTime::now() .duration_since(UNIX_EPOCH) - .unwrap() + .unwrap_or_default() .as_millis(); let mut ping_guard = self.ping.write().await; *ping_guard = (current as i128 - *last_ping) as i64; @@ -316,12 +325,14 @@ impl ClientConnection { if let DataValue::Str(status_str) = cv.get_data(DataType::UserState) { let user_status = UserStatus::from_str(&status_str).unwrap_or(UserStatus::user_online); if let Some(rho_conn) = self.get_rho_connection().await { - OmegaConnection::client_changed( - rho_conn.get_iota_id().await as i64, - user_id as i64, - user_status, - ) - .await; + self.state + .omega + .client_changed( + rho_conn.get_iota_id().await as i64, + user_id as i64, + user_status, + ) + .await; } } } @@ -359,23 +370,29 @@ impl ClientConnection { return; } }; - let invited = - call_manager::add_invite(call_id, self.user_id, receiver_id as u64, secret.clone()) - .await; + let invited = self + .state + .call_manager + .add_invite(call_id, self.user_id, receiver_id as u64, secret.clone()) + .await; if !invited { self.send_error_response(cv.get_id(), CommunicationType::ErrorInvalidCallId) .await; return; } - if !call_manager::should_forward_invite(self.user_id, receiver_id as u64) { + if !self + .state + .call_manager + .should_forward_invite(self.user_id, receiver_id as u64) + { let response = CommunicationValue::new(CommunicationType::Success).with_id(cv.get_id()); self.send_message(&response).await; return; } // Find target RhoConnection - let target_rho = match rho_manager::get_rho_con_for_user(receiver_id as i64).await { + let target_rho = match self.state.rho.get_for_user(receiver_id as i64).await { Some(rho) => rho, _ => { // Get sender user ID @@ -394,7 +411,7 @@ impl ClientConnection { DataValue::Str("call_invite".to_string()), ); - let omega_conn = get_omega_connection(); + let omega_conn = self.state.omega.clone(); // Send fire-and-forget, don't await to avoid blocking tokio::spawn(async move { let _ = omega_conn.send_message(&push_cv).await; @@ -455,18 +472,26 @@ impl ClientConnection { } }; - if let Some(token) = call_manager::get_call_token(user_id, call_id).await { - let response = CommunicationValue::new(CommunicationType::CallToken) - .with_id(cv.get_id()) - .with_receiver(user_id as u64) - .add_typed_default(DataType::CallToken, DataValue::Str(token)); - self.send_message(&response).await; - } else { - let error_cv = CommunicationValue::new(CommunicationType::ErrorNoCallId) - .with_id(cv.get_id()) - .add_typed_default(DataType::CallId, DataValue::Str(call_id.to_string())); - self.send_message(&error_cv).await; - return; + match self + .state + .call_manager + .get_call_token(user_id, call_id) + .await + { + Ok(token) => { + let response = CommunicationValue::new(CommunicationType::CallToken) + .with_id(cv.get_id()) + .with_receiver(user_id as u64) + .add_typed_default(DataType::CallToken, DataValue::Str(token)); + self.send_message(&response).await; + } + Err(error) => { + log::warn!("Unable to create call token for {}: {}", call_id, error); + let error_cv = CommunicationValue::new(CommunicationType::ErrorNoCallId) + .with_id(cv.get_id()) + .add_typed_default(DataType::CallId, DataValue::Str(call_id.to_string())); + self.send_message(&error_cv).await; + } } } async fn handle_get_call_data(self: Arc, cv: CommunicationValue) { @@ -488,7 +513,7 @@ impl ClientConnection { } }; - if let Some(call) = call_manager::get_call(call_id).await { + if let Some(call) = self.state.call_manager.get_call(call_id).await { if let Some(_) = call.get_caller(user_id).await { let mut user_ids: Vec = Vec::new(); let members = call.members.read().await.clone(); @@ -535,7 +560,7 @@ impl ClientConnection { .as_signed_number() .unwrap_or(0); - let Some(call) = call_manager::get_call(call_id).await else { + let Some(call) = self.state.call_manager.get_call(call_id).await else { self.send_error_response(cv.get_id(), CommunicationType::ErrorNotFound) .await; return; @@ -547,7 +572,11 @@ impl ClientConnection { return; }; if caller.has_admin() { - let _ = call_util::remove_participant(call_id, user_id as u64).await; + let _ = self + .state + .livekit + .remove_participant(call_id, user_id as u64) + .await; if let Some(target) = call.get_caller(user_id as u64).await { target.set_timeout(untill as i64).await; } @@ -565,7 +594,7 @@ impl ClientConnection { .as_signed_number() .unwrap_or(0); - let Some(call) = call_manager::get_call(call_id).await else { + let Some(call) = self.state.call_manager.get_call(call_id).await else { self.send_error_response(cv.get_id(), CommunicationType::ErrorNotFound) .await; return; @@ -588,13 +617,18 @@ impl ClientConnection { }; let enable = cv.get_data(DataType::Enabled).as_bool().unwrap_or(true); - let call = call_manager::get_call(call_id).await; + let call = self.state.call_manager.get_call(call_id).await; let mut short_link = None; if let Some(call) = call { if let Some(caller) = call.get_caller(self.get_user_id().await).await { if caller.has_admin() { - call.set_anonymous_joining(enable).await; + call.set_anonymous_joining( + enable, + &self.state.omega, + self.state.config.omikron_id, + ) + .await; } } short_link = call.get_short_link().await; @@ -626,11 +660,22 @@ impl ClientConnection { match resolver.txt_lookup(path).await { Ok(txt_lookup) => { if let Some(txt_record) = txt_lookup.iter().next() { - let record_text: String = txt_record + let record_bytes: Vec = txt_record .txt_data() .iter() - .map(|b| String::from_utf8_lossy(b)) + .flat_map(|chunk| chunk.iter().copied()) .collect(); + let record_text = match String::from_utf8(record_bytes) { + Ok(text) => text, + Err(_) => { + self.send_error_response( + cv.get_id(), + CommunicationType::ErrorInvalidData, + ) + .await; + return; + } + }; let response = CommunicationValue::new(CommunicationType::LoadTxtRecord) .with_id(cv.get_id()) @@ -701,7 +746,10 @@ impl ClientConnection { return; } - let load_uuid_response = get_omega_connection() + let load_uuid_response = self + .state + .omega + .clone() .await_response( &CommunicationValue::new(CommunicationType::GetUserData) .with_id(cv.clone().get_id()) @@ -830,7 +878,7 @@ impl ClientConnection { /// Handle connection close pub async fn handle_close(&self) { let user_id = self.get_user_id().await; - if let Some(rho_conn) = rho_manager::get_rho_con_for_user(user_id as i64).await { + if let Some(rho_conn) = self.state.rho.get_for_user(user_id as i64).await { rho_conn .close_client_connection(Arc::new(self.clone())) .await; @@ -842,6 +890,7 @@ impl ClientConnection { impl Clone for ClientConnection { fn clone(&self) -> Self { Self { + state: self.state.clone(), sender: Arc::clone(&self.sender), receiver: Arc::clone(&self.receiver), user_id: self.user_id, @@ -852,6 +901,7 @@ impl Clone for ClientConnection { rho_connection: Arc::clone(&self.rho_connection), interested_users: Arc::clone(&self.interested_users), is_open: Arc::clone(&self.is_open), + message_slots: Arc::clone(&self.message_slots), } } } diff --git a/src/rho/connection.rs b/src/rho/connection.rs index 72d604e..fb99436 100755 --- a/src/rho/connection.rs +++ b/src/rho/connection.rs @@ -1,19 +1,18 @@ -use std::{collections::BTreeMap, collections::HashMap, sync::Arc, time::Duration}; +use std::{sync::Arc, time::Duration}; use tokio::sync::RwLock; use uuid::Uuid; use crate::{ anonymous_clients::anonymous_client_connection::AnonymousClientConnection, - calls::call_manager, + app_state::AppState, log_cv_out, log_err, log_in, log_out, - omega::omega_connection::get_omega_connection, rho::{ app_connection::AppConnection, client_connection::ClientConnection, - iota_connection::IotaConnection, rho_connection::RhoConnection, rho_manager, + iota_connection::IotaConnection, rho_connection::RhoConnection, }, util::logger::PrintType, }; -use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataTypeId, DataValue}; +use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue}; use mtp::host::AuthState; use mtp::webserver::{WebMTPConnection, WebMtpReceiver, WebMtpSender}; @@ -38,6 +37,7 @@ pub enum ConnectionKind { } pub struct GeneralConnection { + pub state: Arc, pub sender: Arc, pub receiver: Arc, @@ -61,7 +61,7 @@ impl GeneralConnection { * to a known, appropriately-authenticated kind is rejected here instead * of being handed off to a connection handler. */ - pub fn new(conn: WebMTPConnection) -> Option> { + pub fn new(conn: WebMTPConnection, state: Arc) -> Option> { let kind = match (conn.description.as_deref(), &conn.auth_state) { (Some("iota"), AuthState::Authenticated) => ConnectionKind::Iota, (Some("client"), AuthState::Authenticated) => ConnectionKind::Client, @@ -70,6 +70,7 @@ impl GeneralConnection { }; Some(Arc::new(Self { + state, sender: Arc::new(conn.sender), receiver: Arc::new(conn.receiver), connection_kind: kind, @@ -89,293 +90,181 @@ impl GeneralConnection { } async fn migrate(self: &Arc) -> bool { - let id = self.id; - match self.connection_kind { - ConnectionKind::Client => { - let notify = CommunicationValue::new(CommunicationType::UserConnected) - .add_typed_default( - DataType::UserId, - DataValue::SignedNumber((id as i64).into()), - ); - get_omega_connection().send_message(¬ify).await; - - let user_id = id as i64; - - let client = ClientConnection::from_general(self.clone(), id).await; - - let mut rho = rho_manager::get_rho_con_for_user(user_id).await; - - if rho.is_none() { - let get_user_msg = CommunicationValue::new(CommunicationType::GetUserData) - .add_typed_default( - DataType::UserId, - DataValue::SignedNumber(user_id.into()), - ); - - if let Ok(user_data_cv) = get_omega_connection() - .await_response(&get_user_msg, Some(Duration::from_secs(20))) - .await - { - if let DataValue::SignedNumber(iota_id) = - user_data_cv.get_data(DataType::IotaId) - { - if let Some(bound_rho) = - rho_manager::bind_user_to_iota(user_id, *iota_id as i64).await - { - bound_rho.bind_user_id(user_id).await; - rho = Some(bound_rho); - } - } - } - } - - *self.rho_connection.write().await = rho.clone(); - - if let Some(rho_conn) = rho { - let session_id = *self.session_id.read().await as i64; - let iota_msg = CommunicationValue::new(CommunicationType::ClientConnected) - .add_typed_default( - DataType::UserId, - DataValue::SignedNumber(user_id.into()), - ) - .add_typed_default( - DataType::SessionId, - DataValue::SignedNumber(session_id.into()), - ); - - if let Ok(resp) = rho_conn - .get_iota_connection() - .clone() - .await_response(&iota_msg, Some(Duration::from_secs(20))) - .await - { - /* - * mtp's handshake already confirmed this login; this - * message is an unsolicited push carrying the initial - * contacts/calls state, not a reply to a client-side - * request id (the old raw hello frame this used to - * reply to no longer reaches application code). - */ - let mut ident_resp = - CommunicationValue::new(CommunicationType::IdentificationResponse); - let tm = mtp::codec::TypeMap::latest(); - for (k, v) in resp.iter_typed_data() { - let value_to_add = if k == Some(DataType::Contacts) { - if let Some(contacts) = v.as_array() { - let call_groups = - call_manager::get_call_groups(user_id as u64).await; - let mut invites: HashMap> = HashMap::new(); - let mut global_calls: Vec = Vec::new(); - - for call in call_groups { - let members = call.members.read().await; - - let call_self = - members.iter().find(|m| m.user_id == user_id as u64); - - if let Some(call_self) = call_self { - let timeout = *call_self.timeout.read().await; - let admin = call_self.has_admin(); - - // List of all members in the call - let member_ids: Vec = members - .iter() - .map(|m| DataValue::SignedNumber(m.user_id.into())) - .collect(); - - let mut base_call_map: BTreeMap = - BTreeMap::new(); - base_call_map.insert( - DataType::CallId.to_id(&tm), - DataValue::Str(call.call_id.to_string()), - ); - base_call_map.insert( - DataType::CallMembers.to_id(&tm), - DataValue::Array(member_ids), - ); - - if timeout > 0 { - base_call_map.insert( - DataType::Timeout.to_id(&tm), - DataValue::SignedNumber( - (timeout as i64).into(), - ), - ); - } - - if admin { - base_call_map.insert( - DataType::HasAdmin.to_id(&tm), - DataValue::Bool(true), - ); - } - - if let Some(secret) = - call.get_secret_for_user(user_id as u64).await - { - base_call_map.insert( - DataType::CallSecret.to_id(&tm), - secret.to_data_value(), - ); - } - - // Add to global calls with only this user's recipient-specific secret. - global_calls.push(DataValue::container_from_map( - &base_call_map, - )); - - for member in members.iter() { - let member_id = member.user_id; - if member_id == user_id as u64 { - continue; - } - - let contact_call_map = base_call_map.clone(); - - invites - .entry(member_id as i64) - .or_insert_with(Vec::new) - .push(DataValue::container_from_map( - &contact_call_map, - )); - } - } - } - - let mut new_contacts: Vec = Vec::new(); - for contact in contacts { - if let Some(mut contact_map) = contact.as_map() { - if let Some(DataValue::SignedNumber(id)) = - contact_map.get(&DataType::UserId.to_id(&tm)) - { - if let Some(call_list) = invites.get(&(*id as i64)) - { - contact_map.insert( - DataType::Calls.to_id(&tm), - DataValue::Array(call_list.clone()), - ); - } - } - new_contacts - .push(DataValue::container_from_map(&contact_map)); - } else { - new_contacts.push(contact.clone()); - } - } - - ident_resp = ident_resp.add_typed_default( - DataType::Calls, - DataValue::Array(global_calls), - ); - DataValue::Array(new_contacts) - } else { - v.clone() - } - } else { - v.clone() - }; - - if let Some(dt) = k { - ident_resp = ident_resp.add_typed_default(dt, value_to_add); - } - } - log_cv_out!(ident_resp); - let _ = self.sender.send(&ident_resp).await; - rho_conn.bind_user_id(user_id).await; - rho_conn.add_client_connection(client.clone()).await; - } - } else { - log_err!( - user_id, - PrintType::Client, - "No RhoConnection found for user {}, client not attached to iota", - id - ); - } - - client.start(); - } - ConnectionKind::Iota => { - log_in!(id as i64, PrintType::Iota, "Iota {} connected", id); - - let notify = CommunicationValue::new(CommunicationType::IotaConnected) - .add_typed_default(DataType::IotaId, DataValue::SignedNumber(id.into())); - get_omega_connection().send_message(¬ify).await; - - let iota = IotaConnection::from_general(self.clone(), id).await; - - let rho = Arc::new(RhoConnection::new(iota.clone(), Vec::new()).await); - - iota.set_rho_connection(rho.clone()).await; - - rho_manager::add_rho(rho).await; - - let get_iota_msg = CommunicationValue::new(CommunicationType::GetIotaData) - .add_typed_default(DataType::IotaId, DataValue::SignedNumber(id.into())); - - if let Ok(iota_data_cv) = get_omega_connection() - .await_response(&get_iota_msg, Some(Duration::from_secs(20))) - .await - { - if let DataValue::Array(users) = iota_data_cv.get_data(DataType::UserIds) { - let mut user_ids: Vec = Vec::new(); - for value in users { - if let DataValue::SignedNumber(user_id) = value { - user_ids.push(*user_id as u64); - } - } - iota.set_user_ids(user_ids).await; - } - } - - // mtp's handshake already confirmed this login/registration to the Iota. - iota.clone().start(); - } - ConnectionKind::AnonymousClient => { - let client = AnonymousClientConnection::from_general(self.clone(), id).await; - client.start(); - } - ConnectionKind::Phi => { - let user_id = id as i64; - let mut rho = rho_manager::get_rho_con_for_user(user_id).await; - - if rho.is_none() { - let get_user_msg = CommunicationValue::new(CommunicationType::GetUserData) - .add_typed_default( - DataType::UserId, - DataValue::SignedNumber(user_id.into()), - ); - - if let Ok(user_data_cv) = get_omega_connection() - .await_response(&get_user_msg, Some(Duration::from_secs(20))) - .await - { - if let DataValue::SignedNumber(iota_id) = - user_data_cv.get_data(DataType::IotaId) - { - if let Some(bound_rho) = - rho_manager::bind_user_to_iota(user_id, *iota_id as i64).await - { - bound_rho.bind_user_id(user_id).await; - rho = Some(bound_rho); - } - } - } - } - - *self.rho_connection.write().await = rho.clone(); - - let app_conn = AppConnection::from_general(self.clone(), id).await; - - if let Some(rho_conn) = rho { - rho_conn.bind_user_id(user_id).await; - rho_conn.add_app_connection(app_conn.clone()).await; - } - - app_conn.start(); - } + ConnectionKind::Client => self.migrate_client().await, + ConnectionKind::Iota => self.migrate_iota().await, + ConnectionKind::AnonymousClient => self.migrate_anonymous_client().await, + ConnectionKind::Phi => self.migrate_phi().await, } true } + + async fn migrate_client(self: &Arc) { + let id = self.id; + let user_id = id as i64; + self.notify_user_connected(user_id).await; + + let client = ClientConnection::from_general(self.clone(), id).await; + let rho = self.find_user_rho(user_id).await; + *self.rho_connection.write().await = rho.clone(); + + if let Some(rho_conn) = rho { + if let Some(response) = self.request_initial_client_state(&rho_conn, user_id).await { + let response = self.add_call_state(response, user_id).await; + log_cv_out!(response); + let _ = self.sender.send(&response).await; + rho_conn.bind_user_id(user_id).await; + rho_conn.add_client_connection(client.clone()).await; + } + } else { + log_err!( + user_id, + PrintType::Client, + "No RhoConnection found for user {}, client not attached to iota", + id + ); + } + + client.start(); + } + + async fn notify_user_connected(&self, user_id: i64) { + let notify = CommunicationValue::new(CommunicationType::UserConnected) + .add_typed_default(DataType::UserId, DataValue::SignedNumber(user_id.into())); + self.state.omega.clone().send_message(¬ify).await; + } + + async fn find_user_rho(&self, user_id: i64) -> Option> { + if let Some(rho) = self.state.rho.get_for_user(user_id).await { + return Some(rho); + } + + let get_user_msg = CommunicationValue::new(CommunicationType::GetUserData) + .add_typed_default(DataType::UserId, DataValue::SignedNumber(user_id.into())); + let Ok(user_data_cv) = self + .state + .omega + .clone() + .await_response(&get_user_msg, Some(Duration::from_secs(20))) + .await + else { + return None; + }; + + let DataValue::SignedNumber(iota_id) = user_data_cv.get_data(DataType::IotaId) else { + return None; + }; + let rho = self + .state + .rho + .bind_user_to_iota(user_id, *iota_id as i64) + .await?; + rho.bind_user_id(user_id).await; + Some(rho) + } + + async fn request_initial_client_state( + &self, + rho: &Arc, + user_id: i64, + ) -> Option { + let session_id = *self.session_id.read().await as i64; + let request = CommunicationValue::new(CommunicationType::ClientConnected) + .add_typed_default(DataType::UserId, DataValue::SignedNumber(user_id.into())) + .add_typed_default( + DataType::SessionId, + DataValue::SignedNumber(session_id.into()), + ); + rho.get_iota_connection() + .clone() + .await_response(&request, Some(Duration::from_secs(20))) + .await + .ok() + } + + async fn add_call_state( + &self, + response: CommunicationValue, + user_id: i64, + ) -> CommunicationValue { + let mut output = CommunicationValue::new(CommunicationType::IdentificationResponse); + + for (key, value) in response.iter_typed_data() { + if key == Some(DataType::Contacts) { + if let Some(contacts) = value.as_array() { + let (contacts, global_calls) = self + .state + .call_state_aggregator + .augment_contacts(user_id as u64, contacts.clone()) + .await; + output = + output.add_typed_default(DataType::Contacts, DataValue::Array(contacts)); + output = + output.add_typed_default(DataType::Calls, DataValue::Array(global_calls)); + continue; + } + } + if let Some(data_type) = key { + output = output.add_typed_default(data_type, value.clone()); + } + } + output + } + + async fn migrate_iota(self: &Arc) { + let id = self.id; + log_in!(id as i64, PrintType::Iota, "Iota {} connected", id); + let notify = CommunicationValue::new(CommunicationType::IotaConnected) + .add_typed_default(DataType::IotaId, DataValue::SignedNumber(id.into())); + self.state.omega.clone().send_message(¬ify).await; + + let iota = IotaConnection::from_general(self.clone(), id).await; + let rho = Arc::new(RhoConnection::new(iota.clone(), Vec::new()).await); + iota.set_rho_connection(rho.clone()).await; + self.state.rho.add(rho).await; + self.load_iota_users(&iota, id).await; + iota.start(); + } + + async fn load_iota_users(&self, iota: &Arc, iota_id: u64) { + let request = CommunicationValue::new(CommunicationType::GetIotaData) + .add_typed_default(DataType::IotaId, DataValue::SignedNumber(iota_id.into())); + let Ok(response) = self + .state + .omega + .clone() + .await_response(&request, Some(Duration::from_secs(20))) + .await + else { + return; + }; + let DataValue::Array(users) = response.get_data(DataType::UserIds) else { + return; + }; + let user_ids = users + .iter() + .filter_map(|value| match value { + DataValue::SignedNumber(id) => Some(*id as u64), + _ => None, + }) + .collect(); + iota.set_user_ids(user_ids).await; + } + + async fn migrate_anonymous_client(self: &Arc) { + let client = AnonymousClientConnection::from_general(self.clone(), self.id).await; + client.start(); + } + + async fn migrate_phi(self: &Arc) { + let user_id = self.id as i64; + let rho = self.find_user_rho(user_id).await; + *self.rho_connection.write().await = rho.clone(); + let app_conn = AppConnection::from_general(self.clone(), self.id).await; + if let Some(rho_conn) = rho { + rho_conn.bind_user_id(user_id).await; + rho_conn.add_app_connection(app_conn.clone()).await; + } + app_conn.start(); + } } diff --git a/src/rho/iota_connection.rs b/src/rho/iota_connection.rs index 2297e8a..d154bf7 100755 --- a/src/rho/iota_connection.rs +++ b/src/rho/iota_connection.rs @@ -1,12 +1,12 @@ +use crate::app_state::AppState; use crate::calls::call_group::CallGroup; -use crate::calls::call_manager; use crate::log_cv_in; use crate::log_cv_out; use crate::log_err; use crate::log_in; use crate::log_out; -use crate::omega::omega_connection::get_omega_connection; use crate::rho::connection::{GeneralConnection, MtpReceiver, MtpSender}; +use crate::util::data_type_id; use crate::util::logger::PrintType; use dashmap::DashMap; use mtp::codec::CommunicationType; @@ -17,18 +17,15 @@ use mtp::codec::DataValue; use mtp::codec::TypeMap; use mtp::crypto::KemPublicKey; use std::collections::BTreeMap; -use std::{collections::HashMap, sync::Arc, sync::LazyLock, time::Duration}; +use std::{collections::HashMap, sync::Arc, time::Duration}; use tokio::sync::RwLock; use tokio::sync::mpsc; -use super::{rho_connection::RhoConnection, rho_manager}; -use crate::omega::omega_connection::OmegaConnection; - -static PENDING_CHAT_SECRETS: LazyLock>> = - LazyLock::new(DashMap::new); +use super::rho_connection::RhoConnection; #[allow(dead_code)] pub struct IotaConnection { + pub state: Arc, pub iota_id: u64, pub client_version: String, pub sender: Arc, @@ -39,11 +36,14 @@ pub struct IotaConnection { pub waiting_tasks: DashMap, CommunicationValue) -> bool + Send + Sync>>, pub rho_connection: Arc>>>, + pending_chat_secrets: DashMap>, + message_slots: Arc, } impl IotaConnection { pub async fn from_general(general: Arc, iota_id: u64) -> Arc { Arc::new(Self { + state: general.state.clone(), ping: Arc::new(RwLock::new(0)), pub_key: Arc::new(RwLock::new(None)), rho_connection: general.rho_connection.clone(), @@ -53,6 +53,8 @@ impl IotaConnection { iota_id: iota_id, client_version: general.client_version.read().await.clone(), waiting_tasks: DashMap::new(), + pending_chat_secrets: DashMap::new(), + message_slots: Arc::new(tokio::sync::Semaphore::new(32)), }) } pub fn start(self: Arc) { @@ -128,7 +130,7 @@ impl IotaConnection { } async fn flush_pending_chat_secrets(&self, user_id: u64) { - let Some((_, messages)) = PENDING_CHAT_SECRETS.remove(&user_id) else { + let Some((_, messages)) = self.pending_chat_secrets.remove(&user_id) else { return; }; @@ -137,13 +139,13 @@ impl IotaConnection { } } - fn store_pending_chat_secret(cv: CommunicationValue) { + fn store_pending_chat_secret(&self, cv: CommunicationValue) { let receiver_id = cv.get_receiver(); if receiver_id == 0 || !cv.is_type(CommunicationType::SetChatSecret) { return; } - PENDING_CHAT_SECRETS + self.pending_chat_secrets .entry(receiver_id) .or_default() .push(cv); @@ -187,6 +189,10 @@ impl IotaConnection { /// Handle incoming message from Iota pub async fn handle_message(self: Arc, cv: CommunicationValue) { + let Ok(permit) = self.message_slots.clone().acquire_owned().await else { + return; + }; + let _permit = permit; let msg_id = cv.get_id(); if let Some((_, task)) = self.waiting_tasks.remove(&msg_id) { if (task)(self.clone(), cv.clone()) { @@ -219,7 +225,10 @@ impl IotaConnection { } if cv.is_type(CommunicationType::CompleteRegisterUser) { - let response_cv = get_omega_connection() + let response_cv = self + .state + .omega + .clone() .await_response( &cv.clone().with_sender(self.iota_id), Some(Duration::from_secs(20)), @@ -265,7 +274,10 @@ impl IotaConnection { async fn handle_omega_forward(self: Arc, cv: CommunicationValue) { let iota_for_closure = self.clone(); - let response_cv = get_omega_connection() + let response_cv = self + .state + .omega + .clone() .await_response(&cv.with_sender(self.iota_id), Some(Duration::from_secs(20))) .await; if let Ok(response_cv) = response_cv { @@ -293,10 +305,13 @@ impl IotaConnection { .map(|(k, v)| { let mut map = BTreeMap::new(); if let Ok(uid) = k.parse::() { - map.insert(DataType::UserId.to_id(&tm), DataValue::SignedNumber(uid)); + map.insert( + data_type_id(DataType::UserId, &tm), + DataValue::SignedNumber(uid), + ); } map.insert( - DataType::LastPing.to_id(&tm), + data_type_id(DataType::LastPing, &tm), DataValue::SignedNumber(v.into()), ); DataValue::container_from_map(&map) @@ -327,11 +342,11 @@ impl IotaConnection { ); if my_user_ids.contains(&(sender_id as u64)) { - if let Some(target_rho) = rho_manager::get_rho_con_for_user(receiver_id as i64).await { + if let Some(target_rho) = self.state.rho.get_for_user(receiver_id as i64).await { target_rho.message_to_iota(cv).await; } else { if cv.is_type(CommunicationType::SetChatSecret) { - Self::store_pending_chat_secret(cv.clone()); + self.store_pending_chat_secret(cv.clone()); let success = CommunicationValue::new(CommunicationType::Success) .with_id(cv.get_id()) .with_sender(cv.get_sender()) @@ -387,7 +402,7 @@ impl IotaConnection { // ============================ // Load Calls // ============================ - let calls: Vec> = call_manager::get_call_groups(user_id).await; + let calls: Vec> = self.state.call_manager.get_call_groups(user_id).await; let mut invites: HashMap> = HashMap::new(); let mut global_calls: Vec = Vec::new(); @@ -412,27 +427,31 @@ impl IotaConnection { // Build base call container let mut base_call_map: BTreeMap = BTreeMap::new(); base_call_map.insert( - DataType::CallId.to_id(&tm), + data_type_id(DataType::CallId, &tm), DataValue::Str(call.call_id.to_string()), ); base_call_map.insert( - DataType::CallMembers.to_id(&tm), + data_type_id(DataType::CallMembers, &tm), DataValue::Array(member_ids), ); if timeout > 0 { base_call_map.insert( - DataType::Timeout.to_id(&tm), + data_type_id(DataType::Timeout, &tm), DataValue::SignedNumber(timeout.into()), ); } if admin { - base_call_map.insert(DataType::HasAdmin.to_id(&tm), DataValue::Bool(true)); + base_call_map + .insert(data_type_id(DataType::HasAdmin, &tm), DataValue::Bool(true)); } if let Some(secret) = call.get_secret_for_user(user_id).await { - base_call_map.insert(DataType::CallSecret.to_id(&tm), secret.to_data_value()); + base_call_map.insert( + data_type_id(DataType::CallSecret, &tm), + secret.to_data_value(), + ); } // Add to global calls with only this user's recipient-specific secret. @@ -473,7 +492,7 @@ impl IotaConnection { entries.iter().cloned().collect(); if let Some(DataValue::SignedNumber(id)) = - user_map.get(&DataType::UserId.to_id(&tm)) + user_map.get(&data_type_id(DataType::UserId, &tm)) { interested_ids.push(*id as i64); @@ -481,7 +500,7 @@ impl IotaConnection { && !call_list.is_empty() { user_map.insert( - DataType::Calls.to_id(&tm), + data_type_id(DataType::Calls, &tm), DataValue::Array(call_list.clone()), ); } @@ -498,7 +517,10 @@ impl IotaConnection { // ============================ // Notify Omega // ============================ - OmegaConnection::user_states(user_id as i64, interested_ids.clone()).await; + self.state + .omega + .user_states(user_id as i64, interested_ids.clone()) + .await; // ============================ // Notify Rho diff --git a/src/rho/rho_connection.rs b/src/rho/rho_connection.rs index e0023a8..c298b7a 100644 --- a/src/rho/rho_connection.rs +++ b/src/rho/rho_connection.rs @@ -1,7 +1,7 @@ -use super::{client_connection::ClientConnection, iota_connection::IotaConnection, rho_manager}; +use super::{client_connection::ClientConnection, iota_connection::IotaConnection}; -use crate::omega::omega_connection::OmegaConnection; use crate::{data::user::UserStatus, rho::app_connection::AppConnection}; +use dashmap::DashMap; use mtp::codec::{CommunicationType, CommunicationValue, DataType, DataValue}; use std::collections::HashMap; use std::sync::Arc; @@ -11,8 +11,8 @@ use uuid::Uuid; pub struct RhoConnection { iota_connection: Arc, user_ids: Arc>>, - client_connections: Arc>>>, - app_connections: Arc>>>, + client_connections: DashMap<(u64, u64), Arc>, + app_connections: DashMap<(u64, String, Uuid), Arc>, } impl RhoConnection { @@ -21,8 +21,8 @@ impl RhoConnection { let rho_connection = Self { iota_connection, user_ids: Arc::new(RwLock::new(user_ids.clone())), - client_connections: Arc::new(RwLock::new(Vec::new())), - app_connections: Arc::new(RwLock::new(Vec::new())), + client_connections: DashMap::new(), + app_connections: DashMap::new(), }; rho_connection @@ -58,8 +58,10 @@ impl RhoConnection { } pub async fn get_client_connections(&self) -> Vec> { - let connections = self.client_connections.read().await; - connections.clone() + self.client_connections + .iter() + .map(|entry| entry.value().clone()) + .collect() } /// Get client connections for a specific user @@ -67,14 +69,11 @@ impl RhoConnection { &self, user_id: i64, ) -> Vec> { - let connections = self.client_connections.read().await; - let mut collections = Vec::new(); - for con in connections.iter() { - if con.get_user_id().await == user_id as u64 { - collections.push(con.clone()); - } - } - collections + self.client_connections + .iter() + .filter(|entry| entry.key().0 == user_id as u64) + .map(|entry| entry.value().clone()) + .collect() } #[allow(unused)] @@ -84,10 +83,10 @@ impl RhoConnection { app_identifier: Option, app_session: Option, ) -> Vec> { - let connections = self.app_connections.read().await; - connections + self.app_connections .iter() - .filter(|conn| { + .filter(|entry| { + let conn = entry.value(); if let Some(uid) = userid { if conn.user_id != uid as u64 { return false; @@ -105,18 +104,22 @@ impl RhoConnection { } true }) - .cloned() + .map(|entry| entry.value().clone()) .collect() } pub async fn add_app_connection(&self, connection: Arc) { - let mut connections = self.app_connections.write().await; - connections.push(connection); + let key = ( + connection.user_id, + connection.app_identifier.clone(), + connection.app_session, + ); + self.app_connections.insert(key, connection); } pub async fn close_app_connection(&self, connection: Arc) { - let mut connections = self.app_connections.write().await; - connections.retain(|c| c.app_session != connection.app_session); + self.app_connections + .retain(|_, c| c.app_session != connection.app_session); } /// Add a client connection @@ -130,37 +133,25 @@ impl RhoConnection { self.iota_connection.send_message(¬ification).await; - let mut should_notify_online = false; - { - let mut connections = self.client_connections.write().await; - let mut keep = Vec::new(); - let mut had_user = false; - - for con in connections.drain(..) { - if con.user_id as i64 == user_id { - had_user = true; - if con.session_id as i64 == session_id { - continue; - } - } - keep.push(con); - } - - if !had_user { - should_notify_online = true; - } - - keep.push(Arc::clone(&connection)); - *connections = keep; - } + let should_notify_online = !self + .client_connections + .iter() + .any(|entry| entry.key().0 == user_id as u64); + self.client_connections + .remove(&(user_id as u64, session_id as u64)); + self.client_connections + .insert((user_id as u64, session_id as u64), Arc::clone(&connection)); if should_notify_online { - OmegaConnection::client_changed( - self.get_iota_id().await as i64, - user_id, - UserStatus::user_online, - ) - .await; + self.iota_connection + .state + .omega + .client_changed( + self.get_iota_id().await as i64, + user_id, + UserStatus::user_online, + ) + .await; } } @@ -168,32 +159,28 @@ impl RhoConnection { pub async fn close_client_connection(&self, connection: Arc) { let target_user_id = connection.user_id as i64; let target_session_id = connection.session_id as i64; - let mut remaining_for_user = false; - - { - let mut connections = self.client_connections.write().await; - let mut keep = Vec::new(); - for con in connections.drain(..) { - if con.user_id as i64 == target_user_id - && con.session_id as i64 == target_session_id - { - continue; - } - if con.user_id as i64 == target_user_id { - remaining_for_user = true; - } - keep.push(con); - } - *connections = keep; - } + self.client_connections + .remove(&(target_user_id as u64, target_session_id as u64)); + let remaining_for_user = self + .client_connections + .iter() + .any(|entry| entry.key().0 == target_user_id as u64); if !remaining_for_user { - OmegaConnection::client_changed( - self.get_iota_id().await as i64, - target_user_id, - UserStatus::user_offline, - ) - .await; + self.iota_connection + .state + .rho + .remove_user_binding(target_user_id, self.get_iota_id().await as i64) + .await; + self.iota_connection + .state + .omega + .client_changed( + self.get_iota_id().await as i64, + target_user_id, + UserStatus::user_offline, + ) + .await; } } @@ -206,15 +193,23 @@ impl RhoConnection { } // Remove from manager - rho_manager::remove_rho(self.get_iota_id().await as i64).await; + self.iota_connection + .state + .rho + .remove(self.get_iota_id().await as i64) + .await; // Notify OmegaConnection - OmegaConnection::close_iota(self.get_iota_id().await as i64).await; + self.iota_connection + .state + .omega + .close_iota(self.get_iota_id().await as i64) + .await; } /// Send message from Iota to specific client pub async fn message_to_client(&self, cv: CommunicationValue) { - let connections = self.client_connections.read().await; + let connections = self.get_client_connections().await; let receiver_id = cv.get_receiver(); let session_id = cv.get_data(DataType::SessionId).as_number(); @@ -238,7 +233,7 @@ impl RhoConnection { /// Set interested users for a specific client pub async fn set_interested(&self, user_id: i64, interested_ids: Vec) { - let connections = self.client_connections.read().await; + let connections = self.get_client_connections().await; for connection in connections.iter() { let conn_user_id = connection.get_user_id().await; if conn_user_id == user_id as u64 { @@ -254,7 +249,7 @@ impl RhoConnection { /// Check if clients are interested in a user #[allow(dead_code)] pub async fn are_they_interested(&self, user_id: i64, user_status: &str) { - let connections = self.client_connections.read().await; + let connections = self.get_client_connections().await; for connection in connections.iter() { connection .clone() @@ -265,7 +260,7 @@ impl RhoConnection { /// Get ping information for all clients pub async fn get_client_pings(&self) -> HashMap { - let connections = self.client_connections.read().await; + let connections = self.get_client_connections().await; let mut pings = HashMap::new(); for connection in connections.iter() { @@ -285,7 +280,6 @@ impl RhoConnection { /// Get count of active client connections #[allow(dead_code)] pub async fn client_count(&self) -> usize { - let connections = self.client_connections.read().await; - connections.len() + self.client_connections.len() } } diff --git a/src/rho/rho_manager.rs b/src/rho/rho_manager.rs index e823a04..da9acd8 100644 --- a/src/rho/rho_manager.rs +++ b/src/rho/rho_manager.rs @@ -1,83 +1,129 @@ +use std::sync::Arc; + +use dashmap::DashMap; + use super::rho_connection::RhoConnection; -use crate::log_in; -use crate::util::logger::PrintType; -use std::{ - collections::HashMap, - sync::{Arc, LazyLock}, -}; -use tokio::sync::RwLock; -pub static RHO_CONNECTIONS: LazyLock>>>> = - LazyLock::new(|| Arc::new(RwLock::new(HashMap::new()))); +/* + * Owns both indices for active Iota connections. Keeping the user index next + * to the Iota index makes binding updates atomic from the manager's caller's + * perspective and lets separate Omikron instances keep separate routing state. + */ +#[derive(Default)] +pub struct RhoManager { + connections: DashMap>, + users: DashMap>, +} -pub async fn get_rho_con_for_user(user_id: i64) -> Option> { - let connections = RHO_CONNECTIONS.read().await; - for rho_connection in connections.values() { - let rho_user_ids = rho_connection.get_user_ids().await; - log_in!( - user_id, - PrintType::Client, - "Comparing user IDs: {:?}", - rho_user_ids - ); - if rho_user_ids.contains(&user_id) { - return Some(Arc::clone(rho_connection)); +impl RhoManager { + pub fn new() -> Self { + Self::default() + } + + pub async fn get_for_user(&self, user_id: i64) -> Option> { + self.users.get(&user_id).map(|entry| entry.value().clone()) + } + + pub async fn contains_iota(&self, iota_id: i64) -> bool { + self.connections.contains_key(&iota_id) + } + + pub async fn bind_user_to_iota( + &self, + user_id: i64, + iota_id: i64, + ) -> Option> { + let rho = self + .connections + .get(&iota_id) + .map(|entry| entry.value().clone()); + if let Some(rho) = rho { + rho.add_user_id(user_id).await; + self.users.insert(user_id, rho.clone()); + Some(rho) + } else { + None } } - None -} -#[allow(dead_code)] -pub async fn contains_iota(iota_id: i64) -> bool { - let connections = RHO_CONNECTIONS.read().await; - connections.contains_key(&iota_id) -} + pub async fn replace_users_for_iota(&self, iota_id: i64, user_ids: Vec) { + let rho = self + .connections + .get(&iota_id) + .map(|entry| entry.value().clone()); + let Some(rho) = rho else { + return; + }; -/// Bind a user ID to an already tracked iota/rho connection. -pub async fn bind_user_to_iota(user_id: i64, iota_id: i64) -> Option> { - let connections = RHO_CONNECTIONS.read().await; - if let Some(rho_connection) = connections.get(&iota_id) { - let rho = Arc::clone(rho_connection); - drop(connections); + let previous_users = rho.get_user_ids().await; + rho.set_user_ids(user_ids.clone()).await; + rho.get_iota_connection() + .set_user_ids( + user_ids + .iter() + .filter_map(|user_id| u64::try_from(*user_id).ok()) + .collect(), + ) + .await; - rho.add_user_id(user_id).await; + for user_id in previous_users { + if self + .users + .get(&user_id) + .is_some_and(|entry| Arc::ptr_eq(entry.value(), &rho)) + { + self.users.remove(&user_id); + } + } + for user_id in user_ids { + self.users.insert(user_id, rho.clone()); + } + } - log_in!( - user_id, - PrintType::Client, - "Bound user {} to iota {}", - user_id, - iota_id - ); + pub async fn remove_user_binding(&self, user_id: i64, iota_id: i64) { + let rho = self.users.get(&user_id).map(|entry| entry.value().clone()); + if let Some(rho) = rho { + if rho.get_iota_id().await as i64 == iota_id { + self.users.remove(&user_id); + } + } + } - Some(rho) - } else { - None + pub async fn remove(&self, iota_id: i64) -> Option> { + let rho = self.connections.remove(&iota_id).map(|(_, rho)| rho); + if let Some(rho) = rho.as_ref() { + self.users.retain(|_, mapped| !Arc::ptr_eq(mapped, rho)); + } + rho + } + + pub async fn add(&self, rho_connection: Arc) { + let iota_id = rho_connection.get_iota_id().await as i64; + let user_ids = rho_connection.get_user_ids().await; + self.connections.insert(iota_id, rho_connection.clone()); + for user_id in user_ids { + self.users.insert(user_id, rho_connection.clone()); + } + } + + pub async fn get_by_iota(&self, iota_id: i64) -> Option> { + self.connections + .get(&iota_id) + .map(|entry| entry.value().clone()) + } + + pub async fn connection_count(&self) -> usize { + self.connections.len() + } + + pub async fn iota_ids(&self) -> Vec { + self.connections.iter().map(|entry| *entry.key()).collect() + } + + pub async fn connections(&self) -> Vec> { + self.connections + .iter() + .map(|entry| entry.value().clone()) + .collect() } } - -/// Remove a RhoConnection by Iota ID -pub async fn remove_rho(iota_id: i64) -> Option> { - let mut connections = RHO_CONNECTIONS.write().await; - connections.remove(&iota_id) -} - -/// Add a RhoConnection to the manager -pub async fn add_rho(rho_connection: Arc) { - let mut connections = RHO_CONNECTIONS.write().await; - let iota_id = rho_connection.get_iota_id().await; - connections.insert(iota_id as i64, rho_connection); -} - -/// Get a RhoConnection by Iota ID directly -#[allow(dead_code)] -pub async fn get_rho_by_iota(iota_id: i64) -> Option> { - let connections = RHO_CONNECTIONS.read().await; - connections.get(&iota_id).map(Arc::clone) -} - -/// Get the count of active connections -pub async fn connection_count() -> usize { - let connections = RHO_CONNECTIONS.read().await; - connections.len() -} diff --git a/src/rho/server.rs b/src/rho/server.rs index 53f8506..78c75a4 100644 --- a/src/rho/server.rs +++ b/src/rho/server.rs @@ -1,13 +1,14 @@ use std::net::{IpAddr, Ipv4Addr}; +use std::sync::Arc; use std::time::Duration; use base64::Engine as _; use base64::engine::general_purpose::STANDARD as BASE64_STD; -use crate::load_keyring; use crate::{ + app_state::AppState, log, log_err, - omega::omega_connection::get_omega_connection, + omega::omega_connection::OmegaConnection, rho::connection::GeneralConnection, util::{file_util::load_file_vec, logger::PrintType}, }; @@ -24,6 +25,7 @@ use mtp::webserver::{MTPWebServer, WebServerConfig}; * through mtp's unauthenticated fallback instead of rejecting it outright. */ pub async fn get_by_connector_id( + omega: Arc, client_id: u64, description: Option, ) -> Option { @@ -38,7 +40,7 @@ pub async fn get_by_connector_id( _ => return None, }; - let response = match get_omega_connection() + let response = match omega .await_response(&request, Some(Duration::from_secs(20))) .await { @@ -62,7 +64,11 @@ pub async fn get_by_connector_id( } /* Only Iota registration goes through mtp's Register flow; users are registered out of band. */ -pub async fn complete_register(pub_key: PublicKeyBundle, description: Option) -> u64 { +pub async fn complete_register( + omega: Arc, + pub_key: PublicKeyBundle, + description: Option, +) -> u64 { println!("Iota register start"); if description.as_deref() != Some("iota") { return 0; @@ -75,7 +81,7 @@ pub async fn complete_register(pub_key: PublicKeyBundle, description: Option Result<(), Box> { +pub async fn start(state: Arc) -> Result<(), Box> { let cert_pem = load_file_vec("certs", "cert.pem").expect("Error loading Pemfile"); let key_pem = load_file_vec("certs", "key.pem").expect("Error loading Keyfile"); let host_config = HostConfig::new( IpAddr::from(Ipv4Addr::new(0, 0, 0, 0)), - port, + state.config.rho_port, cert_pem, key_pem, ) @@ -124,16 +130,33 @@ pub async fn start(port: u16) -> Result<(), Box> { .with_persistent_stream_retries(5, Duration::from_secs(5)), ) .with_authentication( - load_keyring(), - Box::new(|user_id, description| Box::pin(get_by_connector_id(user_id, description))), - Box::new(|pub_key, description| Box::pin(complete_register(pub_key, description))), + state + .keyring_for_host() + .map_err(|error| format!("Unable to copy keyring for host: {error}"))?, + Box::new({ + let omega = state.omega.clone(); + move |user_id, description| { + Box::pin(get_by_connector_id(omega.clone(), user_id, description)) + } + }), + Box::new({ + let omega = state.omega.clone(); + move |pub_key, description| { + Box::pin(complete_register(omega.clone(), pub_key, description)) + } + }), ) .with_authentication_policy(AuthenticationPolicy::AllowAuthentication); let web_config = WebServerConfig::new() .route("/", |_request, response| async move { response.body("OK") })?; let mut host = MTPWebServer::new(host_config, web_config).await?; - log!(0, PrintType::General, "Server listening on port {}.", port); + log!( + 0, + PrintType::General, + "Server listening on port {}.", + state.config.rho_port + ); loop { let conn = match host.accept().await { @@ -153,8 +176,9 @@ pub async fn start(port: u16) -> Result<(), Box> { } }; + let state = state.clone(); tokio::spawn(async move { - let Some(conn) = GeneralConnection::new(conn) else { + let Some(conn) = GeneralConnection::new(conn, state) else { log_err!( 0, PrintType::General, diff --git a/src/services/call_service.rs b/src/services/call_service.rs new file mode 100644 index 0000000..1713815 --- /dev/null +++ b/src/services/call_service.rs @@ -0,0 +1,29 @@ +use std::sync::Arc; + +use uuid::Uuid; + +use crate::{ + app_state::AppState, + calls::{call_group::CallGroup, error::CallError}, +}; + +pub struct CallService { + state: Arc, +} + +impl CallService { + pub fn new(state: Arc) -> Self { + Self { state } + } + + pub async fn get_call(&self, call_id: Uuid) -> Option> { + self.state.call_manager.get_call(call_id).await + } + + pub async fn token_for_user(&self, user_id: u64, call_id: Uuid) -> Result { + self.state + .call_manager + .get_call_token(user_id, call_id) + .await + } +} diff --git a/src/services/call_state.rs b/src/services/call_state.rs new file mode 100644 index 0000000..65ee97f --- /dev/null +++ b/src/services/call_state.rs @@ -0,0 +1,148 @@ +use std::collections::BTreeMap; +use std::sync::Arc; + +use mtp::codec::{DataType, DataValue, TypeMap}; + +use crate::{calls::call_manager::CallManager, util::data_type_id}; + +pub struct CallStateAggregator { + call_manager: Arc, +} + +impl CallStateAggregator { + pub fn new(call_manager: Arc) -> Self { + Self { call_manager } + } + + pub async fn augment_contacts( + &self, + user_id: u64, + contacts: Vec, + ) -> (Vec, Vec) { + let tm = TypeMap::latest(); + let mut invites = std::collections::HashMap::new(); + let mut global_calls = Vec::new(); + + for call in self.call_manager.get_call_groups(user_id).await { + let members = call.members.read().await; + let Some(call_self) = members.iter().find(|member| member.user_id == user_id) else { + continue; + }; + + let mut call_map = BTreeMap::new(); + call_map.insert( + data_type_id(DataType::CallId, &tm), + DataValue::Str(call.call_id.to_string()), + ); + call_map.insert( + data_type_id(DataType::CallMembers, &tm), + DataValue::Array( + members + .iter() + .map(|member| DataValue::SignedNumber(member.user_id.into())) + .collect(), + ), + ); + + let timeout = *call_self.timeout.read().await; + if timeout > 0 { + call_map.insert( + data_type_id(DataType::Timeout, &tm), + DataValue::SignedNumber((timeout as i64).into()), + ); + } + if call_self.has_admin() { + call_map.insert(data_type_id(DataType::HasAdmin, &tm), DataValue::Bool(true)); + } + if let Some(secret) = call.get_secret_for_user(user_id).await { + call_map.insert( + data_type_id(DataType::CallSecret, &tm), + secret.to_data_value(), + ); + } + + let call_value = DataValue::container_from_map(&call_map); + global_calls.push(call_value.clone()); + for member in members.iter().filter(|member| member.user_id != user_id) { + invites + .entry(member.user_id as i64) + .or_insert_with(Vec::new) + .push(call_value.clone()); + } + } + + let user_id_type = data_type_id(DataType::UserId, &tm); + let calls_type = data_type_id(DataType::Calls, &tm); + let augmented_contacts = contacts + .into_iter() + .map(|contact| { + let Some(mut contact_map) = contact.as_map() else { + return contact; + }; + let Some(DataValue::SignedNumber(contact_id)) = contact_map.get(&user_id_type) + else { + return DataValue::container_from_map(&contact_map); + }; + if let Some(calls) = invites.get(&(*contact_id as i64)) { + contact_map.insert(calls_type, DataValue::Array(calls.clone())); + } + DataValue::container_from_map(&contact_map) + }) + .collect(); + + (augmented_contacts, global_calls) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::calls::{call_group::CallGroup, caller::Caller}; + use uuid::Uuid; + + #[tokio::test] + async fn augment_contacts_preserves_contacts_without_call_invites() { + let aggregator = CallStateAggregator::new(Arc::new(CallManager::default())); + let contacts = vec![DataValue::Str("unchanged".to_string())]; + + let (augmented, global_calls) = aggregator.augment_contacts(7, contacts.clone()).await; + + assert_eq!(augmented, contacts); + assert!(global_calls.is_empty()); + } + + #[tokio::test] + async fn augment_contacts_adds_call_to_member_contact() { + let call_id = Uuid::new_v4(); + let manager = Arc::new(CallManager::default()); + let group = Arc::new(CallGroup::new( + call_id, + Arc::new(Caller::new(1, call_id, true)), + )); + group + .members + .write() + .await + .push(Arc::new(Caller::new(2, call_id, false))); + manager.groups.insert(call_id, group); + + let tm = TypeMap::latest(); + let user_id_type = data_type_id(DataType::UserId, &tm); + let calls_type = data_type_id(DataType::Calls, &tm); + let mut contact = BTreeMap::new(); + contact.insert(user_id_type, DataValue::SignedNumber(2.into())); + + let (contacts, global_calls) = CallStateAggregator::new(manager) + .augment_contacts(1, vec![DataValue::container_from_map(&contact)]) + .await; + + assert_eq!(global_calls.len(), 1); + let Some(contact) = contacts[0].as_map() else { + panic!("expected contact map"); + }; + assert!(matches!( + contact.get(&calls_type), + Some(DataValue::Array(_)) + )); + } +} diff --git a/src/services/mod.rs b/src/services/mod.rs new file mode 100644 index 0000000..19cfa24 --- /dev/null +++ b/src/services/mod.rs @@ -0,0 +1,4 @@ +pub mod call_service; +pub mod call_state; +pub mod routing_service; +pub mod user_service; diff --git a/src/services/routing_service.rs b/src/services/routing_service.rs new file mode 100644 index 0000000..86f034f --- /dev/null +++ b/src/services/routing_service.rs @@ -0,0 +1,25 @@ +use std::sync::Arc; + +use mtp::codec::CommunicationValue; + +use crate::{app_state::AppState, rho::rho_connection::RhoConnection}; + +pub struct RoutingService { + state: Arc, +} + +impl RoutingService { + pub fn new(state: Arc) -> Self { + Self { state } + } + + pub async fn route_to_user(&self, user_id: i64, message: CommunicationValue) { + if let Some(connection) = self.state.rho.get_for_user(user_id).await { + connection.message_to_iota(message).await; + } + } + + pub async fn connection_for_user(&self, user_id: i64) -> Option> { + self.state.rho.get_for_user(user_id).await + } +} diff --git a/src/services/user_service.rs b/src/services/user_service.rs new file mode 100644 index 0000000..efc72a5 --- /dev/null +++ b/src/services/user_service.rs @@ -0,0 +1,17 @@ +use std::sync::Arc; + +use crate::{app_state::AppState, rho::rho_connection::RhoConnection}; + +pub struct UserService { + state: Arc, +} + +impl UserService { + pub fn new(state: Arc) -> Self { + Self { state } + } + + pub async fn connection_for_user(&self, user_id: i64) -> Option> { + self.state.rho.get_for_user(user_id).await + } +} diff --git a/src/util/mod.rs b/src/util/mod.rs index d0cc1dd..1a3a05e 100644 --- a/src/util/mod.rs +++ b/src/util/mod.rs @@ -1,2 +1,16 @@ pub mod file_util; pub mod logger; + +use mtp::codec::{DataType, DataTypeId, TypeMap}; + +/* + * Omikron serializes fields using the latest TypeMap, whose built-in fields + * are required by the protocol. Keep the invariant in one place while using + * the non-deprecated conversion API. + */ +pub fn data_type_id(data_type: DataType, type_map: &TypeMap) -> DataTypeId { + let Some(id) = data_type.try_to_id(type_map) else { + panic!("required data type is missing from the TypeMap: {data_type:?}"); + }; + id +}