Better websocket

This commit is contained in:
2026-08-27 16:18:19 +02:00
parent 02a444d6ed
commit a4d675bda9
10 changed files with 141 additions and 142 deletions
Generated
+13
View File
@@ -256,6 +256,7 @@ dependencies = [
"axum", "axum",
"bs58", "bs58",
"ed25519-dalek", "ed25519-dalek",
"futures-util",
"rand 0.8.7", "rand 0.8.7",
"rusqlite", "rusqlite",
"serde", "serde",
@@ -313,6 +314,17 @@ version = "0.3.34"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e" checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e"
[[package]]
name = "futures-macro"
version = "0.3.34"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9fb9654ba8355388abeb8dcb4fc62f511300867002afc858860463bdd9fe0c44"
dependencies = [
"proc-macro2",
"quote",
"syn 3.0.3",
]
[[package]] [[package]]
name = "futures-sink" name = "futures-sink"
version = "0.3.34" version = "0.3.34"
@@ -332,6 +344,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc" checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc"
dependencies = [ dependencies = [
"futures-core", "futures-core",
"futures-macro",
"futures-sink", "futures-sink",
"futures-task", "futures-task",
"pin-project-lite", "pin-project-lite",
+1
View File
@@ -15,3 +15,4 @@ bs58 = "0.5.1"
tower-http = { version = "0.7.0", features = ["fs", "cors"] } tower-http = { version = "0.7.0", features = ["fs", "cors"] }
rusqlite = { version = "0.31", features = ["bundled"] } rusqlite = { version = "0.31", features = ["bundled"] }
uuid = { version = "1.24.1", features = ["v4"] } uuid = { version = "1.24.1", features = ["v4"] }
futures-util = "0.3.34"
+1
View File
@@ -4,6 +4,7 @@ pub mod protocol;
pub mod server; pub mod server;
pub mod types; pub mod types;
pub mod vc_server; pub mod vc_server;
pub mod ws;
use std::{ use std::{
net::{IpAddr, Ipv4Addr, SocketAddr}, net::{IpAddr, Ipv4Addr, SocketAddr},
+24 -38
View File
@@ -3,10 +3,9 @@ use std::{
time::{SystemTime, UNIX_EPOCH}, time::{SystemTime, UNIX_EPOCH},
}; };
use axum::extract::ws::WebSocket;
use ed25519_dalek::{Signer, VerifyingKey}; use ed25519_dalek::{Signer, VerifyingKey};
use crate::server::Server; use crate::{server::Server, ws::EnclaveWebSocket};
use super::*; use super::*;
use crate::server::UserConnections; use crate::server::UserConnections;
@@ -14,22 +13,20 @@ use crate::server::UserConnections;
impl UserConnections { impl UserConnections {
pub async fn initialize( pub async fn initialize(
server: &Arc<Server>, server: &Arc<Server>,
mut socket: WebSocket, socket: Arc<EnclaveWebSocket>,
) -> anyhow::Result<(WebSocket, VerifyingKey, ClientMeta)> { ) -> anyhow::Result<(Arc<EnclaveWebSocket>, VerifyingKey, ClientMeta)> {
let Some(ServerMethod::Initialize { let Some(ServerMethod::Initialize {
public_key: public_key_string, public_key: public_key_string,
signature, signature,
timestamp, timestamp,
hostname, hostname,
}) = read_socket(&mut socket).await? }) = socket.read().await?
else { else {
send_socket( socket
&mut socket, .send(&ClientMethod::Error {
&ClientMethod::Error {
error: Cow::Borrowed("Initialization required"), error: Cow::Borrowed("Initialization required"),
}, })
)
.await?; .await?;
return Err(anyhow::anyhow!( return Err(anyhow::anyhow!(
@@ -40,14 +37,12 @@ impl UserConnections {
let server_timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_millis() as u64; let server_timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_millis() as u64;
if server_timestamp.saturating_sub(timestamp) > 2000 { if server_timestamp.saturating_sub(timestamp) > 2000 {
send_socket( socket
&mut socket, .send(&ClientMethod::Error {
&ClientMethod::Error {
error: Cow::Borrowed( error: Cow::Borrowed(
"Timestamp doesn't match, make sure it's in secs and is (<= 2secs)", "Timestamp doesn't match, make sure it's in secs and is (<= 2secs)",
), ),
}, })
)
.await?; .await?;
return Err(anyhow::anyhow!("Client tampstamp wasn't correct")); return Err(anyhow::anyhow!("Client tampstamp wasn't correct"));
@@ -55,8 +50,7 @@ impl UserConnections {
if hostname != server.config.public_hostname || !server.config.hostnames.contains(&hostname) if hostname != server.config.public_hostname || !server.config.hostnames.contains(&hostname)
{ {
send_socket( socket.send(
&mut socket,
&ClientMethod::Error { &ClientMethod::Error {
error: Cow::Owned(format!("Invalid Hostname, to avoid man-in-the-middle attacks, please use the correct hostname: {}", server.config.public_hostname)), error: Cow::Owned(format!("Invalid Hostname, to avoid man-in-the-middle attacks, please use the correct hostname: {}", server.config.public_hostname)),
}, },
@@ -67,12 +61,10 @@ impl UserConnections {
} }
let Ok(public_key) = crate::crypto::from_string(&public_key_string) else { let Ok(public_key) = crate::crypto::from_string(&public_key_string) else {
send_socket( socket
&mut socket, .send(&ClientMethod::Error {
&ClientMethod::Error {
error: Cow::Borrowed("Invalid public key"), error: Cow::Borrowed("Invalid public key"),
}, })
)
.await?; .await?;
return Err(anyhow::anyhow!("Invalid public key")); return Err(anyhow::anyhow!("Invalid public key"));
@@ -85,21 +77,18 @@ impl UserConnections {
) )
.is_err() .is_err()
{ {
send_socket( socket
&mut socket, .send(&ClientMethod::Error {
&ClientMethod::Error {
error: Cow::Borrowed("Invalid signature"), error: Cow::Borrowed("Invalid signature"),
}, })
)
.await?; .await?;
return Err(anyhow::anyhow!("Invalid signature")); return Err(anyhow::anyhow!("Invalid signature"));
} }
{ {
send_socket( socket
&mut socket, .send(&ClientMethod::Initialized {
&ClientMethod::Initialized {
public_key: crate::crypto::to_string(&server.key.verifying_key()), public_key: crate::crypto::to_string(&server.key.verifying_key()),
signature: crate::crypto::to_string_sig(&server.key.sign( signature: crate::crypto::to_string_sig(&server.key.sign(
format!("{server_timestamp}@{hostname}@{public_key_string}").as_bytes(), format!("{server_timestamp}@{hostname}@{public_key_string}").as_bytes(),
@@ -107,18 +96,15 @@ impl UserConnections {
timestamp: server_timestamp, timestamp: server_timestamp,
hostname, hostname,
}, })
)
.await?; .await?;
} }
let Some(ServerMethod::Meta(meta)) = read_socket(&mut socket).await? else { let Some(ServerMethod::Meta(meta)) = socket.read().await? else {
send_socket( socket
&mut socket, .send(&ClientMethod::Error {
&ClientMethod::Error {
error: Cow::Borrowed("Expected meta"), error: Cow::Borrowed("Expected meta"),
}, })
)
.await?; .await?;
return Err(anyhow::anyhow!( return Err(anyhow::anyhow!(
+6 -10
View File
@@ -1,17 +1,15 @@
use crate::data::messages::{MessageData, StoredMessage}; use crate::data::messages::{MessageData, StoredMessage};
use crate::protocol::{ClientMethod, send_socket}; use crate::protocol::ClientMethod;
use crate::server::Server; use crate::server::Server;
use axum::extract::ws::WebSocket;
use ed25519_dalek::{Verifier, VerifyingKey}; use ed25519_dalek::{Verifier, VerifyingKey};
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
use tokio::sync::Mutex;
pub async fn send_message( pub async fn send_message(
server: &Arc<Server>, server: &Arc<Server>,
verifying_key: VerifyingKey, verifying_key: VerifyingKey,
_socket: &Arc<Mutex<WebSocket>>, _socket: &Arc<crate::ws::EnclaveWebSocket>,
message: MessageData, message: MessageData,
channel_id: String, channel_id: String,
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
@@ -58,7 +56,7 @@ pub async fn send_message(
pub async fn get_messages( pub async fn get_messages(
server: &Arc<Server>, server: &Arc<Server>,
_verifying_key: VerifyingKey, _verifying_key: VerifyingKey,
socket: &Arc<Mutex<WebSocket>>, socket: &Arc<crate::ws::EnclaveWebSocket>,
channel_id: String, channel_id: String,
chunk: u32, chunk: u32,
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
@@ -68,12 +66,10 @@ pub async fn get_messages(
.message_store .message_store
.get_recent_messages(&channel_id, CHUNK_SIZE, chunk)?; .get_recent_messages(&channel_id, CHUNK_SIZE, chunk)?;
send_socket( socket
&mut *socket.lock().await, .send(&ClientMethod::Messages {
&ClientMethod::Messages {
messages: HashMap::from([(channel_id, messages)]), messages: HashMap::from([(channel_id, messages)]),
}, })
)
.await?; .await?;
Ok(()) Ok(())
+8 -60
View File
@@ -1,9 +1,7 @@
use std::{borrow::Cow, collections::HashMap, sync::Arc}; use std::{borrow::Cow, collections::HashMap, sync::Arc};
use axum::extract::ws::{Message, Utf8Bytes, WebSocket};
use ed25519_dalek::VerifyingKey; use ed25519_dalek::VerifyingKey;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tokio::sync::Mutex;
use crate::{data::messages::StoredMessage, server::Server, types::ClientMeta}; use crate::{data::messages::StoredMessage, server::Server, types::ClientMeta};
@@ -110,21 +108,15 @@ pub enum ServerMethod {
pub async fn read_loop( pub async fn read_loop(
server: &Arc<Server>, server: &Arc<Server>,
verifying_key: VerifyingKey, verifying_key: VerifyingKey,
socket: &Arc<Mutex<WebSocket>>, socket: &Arc<crate::ws::EnclaveWebSocket>,
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
let mut socket_lock = socket.lock().await; while let Some(message) = socket.read().await? {
while let Some(message) = read_socket(&mut *socket_lock).await? {
drop(socket_lock);
match message { match message {
ServerMethod::Initialize { .. } => { ServerMethod::Initialize { .. } => {
send_socket( socket
&mut *socket.lock().await, .send(&ClientMethod::Error {
&ClientMethod::Error {
error: Cow::Borrowed("Already initialized"), error: Cow::Borrowed("Already initialized"),
}, })
)
.await?; .await?;
} }
@@ -181,13 +173,11 @@ pub async fn read_loop(
.await .await
.insert(pin, (verifying_key, channel_id.clone())); .insert(pin, (verifying_key, channel_id.clone()));
send_socket( socket
&mut *socket.lock().await, .send(&ClientMethod::JoinVoice {
&ClientMethod::JoinVoice {
channel_id: channel_id.clone(), channel_id: channel_id.clone(),
pin, pin,
}, })
)
.await?; .await?;
} }
@@ -199,49 +189,7 @@ pub async fn read_loop(
.await?; .await?;
} }
} }
socket_lock = socket.lock().await;
} }
Ok(()) Ok(())
} }
pub async fn read_socket(socket: &mut WebSocket) -> anyhow::Result<Option<ServerMethod>> {
match socket.recv().await.transpose()? {
Some(Message::Text(text)) => match serde_json::from_str(&text.to_string()) {
Ok(msg) => Ok(Some(msg)),
Err(e) => {
send_socket(
socket,
&ClientMethod::Error {
error: Cow::Owned(format!("Unable to parse message: {e}")),
},
)
.await?;
Ok(None)
}
},
Some(Message::Ping(v)) => {
socket.send(Message::Pong(v)).await?;
Ok(None)
}
Some(_) => Ok(None),
None => Ok(None),
}
}
pub async fn send_socket(socket: &mut WebSocket, message: &ClientMethod) -> anyhow::Result<()> {
socket
.send(Message::Text(Utf8Bytes::from(serde_json::to_string(
message,
)?)))
.await?;
Ok(())
}
+3 -8
View File
@@ -1,23 +1,18 @@
use std::sync::Arc; use std::sync::Arc;
use axum::extract::ws::WebSocket;
use ed25519_dalek::VerifyingKey; use ed25519_dalek::VerifyingKey;
use tokio::sync::Mutex;
use crate::{ use crate::{protocol::ClientMethod, server::Server};
protocol::{ClientMethod, send_socket},
server::Server,
};
pub async fn get_users( pub async fn get_users(
server: &Arc<Server>, server: &Arc<Server>,
_verifying_key: VerifyingKey, _verifying_key: VerifyingKey,
socket: &Arc<Mutex<WebSocket>>, socket: &Arc<crate::ws::EnclaveWebSocket>,
pubkeys: Vec<String>, pubkeys: Vec<String>,
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
let users = server.user_store.get_users(&pubkeys).await?; let users = server.user_store.get_users(&pubkeys).await?;
send_socket(&mut *socket.lock().await, &ClientMethod::Users { users }).await?; socket.send(&ClientMethod::Users { users }).await?;
Ok(()) Ok(())
} }
+6 -7
View File
@@ -19,8 +19,9 @@ use tokio::{
use crate::{ use crate::{
data::{config::Config, messages::MessageStore, users::UserMetaStore}, data::{config::Config, messages::MessageStore, users::UserMetaStore},
protocol::{ClientMethod, read_loop, send_socket}, protocol::{ClientMethod, read_loop},
types::ClientMeta, types::ClientMeta,
ws::EnclaveWebSocket,
}; };
pub struct VoiceConnection { pub struct VoiceConnection {
@@ -33,7 +34,7 @@ pub struct UserConnections {
pub meta: ClientMeta, pub meta: ClientMeta,
pub counter: AtomicU16, pub counter: AtomicU16,
pub public_key: VerifyingKey, pub public_key: VerifyingKey,
pub connections: Mutex<HashMap<u16, Arc<Mutex<WebSocket>>>>, pub connections: Mutex<HashMap<u16, Arc<crate::ws::EnclaveWebSocket>>>,
pub voice: Mutex<Option<VoiceConnection>>, pub voice: Mutex<Option<VoiceConnection>>,
} }
@@ -66,7 +67,7 @@ impl Server {
let s = self.clone(); let s = self.clone();
ws.on_upgrade(move |socket: WebSocket| async move { ws.on_upgrade(move |socket: WebSocket| async move {
match UserConnections::initialize(&s, socket).await { match UserConnections::initialize(&s, Arc::new(EnclaveWebSocket::new(socket))).await {
Ok((client, public_key, meta)) => { Ok((client, public_key, meta)) => {
if let Err(e) = s if let Err(e) = s
.user_store .user_store
@@ -76,8 +77,6 @@ impl Server {
eprintln!("Failed to upsert client: {e}"); eprintln!("Failed to upsert client: {e}");
} }
let client = Arc::new(Mutex::new(client));
let mut clients_meta = s.clients.lock().await; let mut clients_meta = s.clients.lock().await;
let clients = clients_meta let clients = clients_meta
@@ -154,7 +153,7 @@ impl Server {
impl UserConnections { impl UserConnections {
pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> { pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> {
for (_, conn) in self.connections.lock().await.iter() { for (_, conn) in self.connections.lock().await.iter() {
send_socket(&mut *conn.lock().await, message).await?; conn.send(message).await?;
} }
Ok(()) Ok(())
@@ -162,7 +161,7 @@ impl UserConnections {
pub async fn send_to(&self, id: u16, message: &ClientMethod) -> anyhow::Result<bool> { pub async fn send_to(&self, id: u16, message: &ClientMethod) -> anyhow::Result<bool> {
if let Some(conn) = self.connections.lock().await.get(&id) { if let Some(conn) = self.connections.lock().await.get(&id) {
send_socket(&mut *conn.lock().await, message).await?; conn.send(message).await?;
Ok(true) Ok(true)
} else { } else {
+4 -9
View File
@@ -4,10 +4,7 @@ use anyhow::Context;
use ed25519_dalek::VerifyingKey; use ed25519_dalek::VerifyingKey;
use tokio::net::UdpSocket; use tokio::net::UdpSocket;
use crate::{ use crate::{protocol::ClientMethod, server::Server};
protocol::{ClientMethod, send_socket},
server::Server,
};
use tokio::time::Instant; use tokio::time::Instant;
@@ -107,12 +104,10 @@ impl Server {
if now.duration_since(voice.last_speaking_sent).as_millis() >= 600 { if now.duration_since(voice.last_speaking_sent).as_millis() >= 600 {
for conn in user.connections.lock().await.values() { for conn in user.connections.lock().await.values() {
let _ = send_socket( let _ = conn
&mut *conn.lock().await, .send(&ClientMethod::Speaking {
&ClientMethod::Speaking {
pubkey: crate::crypto::to_string(sender), pubkey: crate::crypto::to_string(sender),
}, })
)
.await; .await;
} }
+65
View File
@@ -0,0 +1,65 @@
use std::borrow::Cow;
use axum::extract::ws::{Message, Utf8Bytes, WebSocket};
use futures_util::{
SinkExt, StreamExt,
stream::{SplitSink, SplitStream},
};
use tokio::sync::Mutex;
use crate::protocol::{ClientMethod, ServerMethod};
pub struct EnclaveWebSocket {
tx: Mutex<SplitSink<WebSocket, Message>>,
rx: Mutex<SplitStream<WebSocket>>,
}
impl EnclaveWebSocket {
pub fn new(ws: WebSocket) -> Self {
let (tx, rx) = ws.split();
Self {
tx: Mutex::new(tx),
rx: Mutex::new(rx),
}
}
pub async fn read(&self) -> anyhow::Result<Option<ServerMethod>> {
match self.rx.lock().await.next().await.transpose()? {
Some(Message::Text(text)) => match serde_json::from_str(&text.to_string()) {
Ok(msg) => Ok(Some(msg)),
Err(e) => {
self.send(&ClientMethod::Error {
error: Cow::Owned(format!("Unable to parse message: {e}")),
})
.await?;
Ok(None)
}
},
Some(Message::Ping(v)) => {
self.tx.lock().await.send(Message::Pong(v)).await?;
Ok(None)
}
Some(_) => Ok(None),
None => Ok(None),
}
}
pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> {
self.tx
.lock()
.await
.send(Message::Text(Utf8Bytes::from(serde_json::to_string(
message,
)?)))
.await?;
Ok(())
}
}