Merge pull request #2 from recurse-chat/multi-client-stuff

Multi client stuff
This commit is contained in:
2026-08-14 17:10:28 +02:00
committed by GitHub
3 changed files with 146 additions and 83 deletions
+36 -18
View File
@@ -4,25 +4,29 @@ use std::{
}; };
use axum::extract::ws::WebSocket; use axum::extract::ws::WebSocket;
use ed25519_dalek::Signer; use ed25519_dalek::{Signer, VerifyingKey};
use crate::server::Server; use crate::server::Server;
use super::*; use super::*;
use crate::server::UserConnections;
impl super::Client { impl UserConnections {
pub async fn initialize(server: &Arc<Server>, mut socket: WebSocket) -> anyhow::Result<Self> { pub async fn initialize(
server: &Arc<Server>,
mut socket: WebSocket,
) -> anyhow::Result<(WebSocket, 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,
}) = Client::read_socket(&mut socket).await? }) = read_socket(&mut socket).await?
else { else {
Client::send_socket( send_socket(
&mut socket, &mut socket,
ClientMethod::Error { &ClientMethod::Error {
error: Cow::Borrowed("Initialization required"), error: Cow::Borrowed("Initialization required"),
}, },
) )
@@ -33,24 +37,42 @@ impl super::Client {
)); ));
}; };
let public_key = crate::signature::from_string(&public_key_string)?; let Ok(public_key) = crate::signature::from_string(&public_key_string) else {
send_socket(
&mut socket,
&ClientMethod::Error {
error: Cow::Borrowed("Invalid public key"),
},
)
.await?;
return Err(anyhow::anyhow!("Invalid public key"));
};
if public_key if public_key
.verify_strict( .verify_strict(
format!("{timestamp}@{hostname}").as_bytes(), format!("{timestamp}@{hostname}").as_bytes(),
&crate::signature::from_string_sig(&signature)?, &crate::signature::from_string_sig(&signature)?,
) )
.is_ok() .is_err()
{ {
send_socket(
&mut socket,
&ClientMethod::Error {
error: Cow::Borrowed("Invalid signature"),
},
)
.await?;
return Err(anyhow::anyhow!("Invalid signature")); return Err(anyhow::anyhow!("Invalid signature"));
} }
{ {
let timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs(); let timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs();
Client::send_socket( send_socket(
&mut socket, &mut socket,
ClientMethod::Initialized { &ClientMethod::Initialized {
public_key: crate::signature::to_string(&server.key.verifying_key()), public_key: crate::signature::to_string(&server.key.verifying_key()),
signature: server signature: server
.key .key
@@ -64,10 +86,10 @@ impl super::Client {
.await?; .await?;
} }
let Some(ServerMethod::Meta(meta)) = Client::read_socket(&mut socket).await? else { let Some(ServerMethod::Meta(meta)) = read_socket(&mut socket).await? else {
Client::send_socket( send_socket(
&mut socket, &mut socket,
ClientMethod::Error { &ClientMethod::Error {
error: Cow::Borrowed("Expected meta"), error: Cow::Borrowed("Expected meta"),
}, },
) )
@@ -78,10 +100,6 @@ impl super::Client {
)); ));
}; };
Ok(Self { Ok((socket, public_key, meta))
socket,
meta,
public_key,
})
} }
} }
+23 -35
View File
@@ -1,17 +1,11 @@
use std::borrow::Cow; use std::{borrow::Cow, sync::Arc};
use axum::extract::ws::{Message, Utf8Bytes, WebSocket}; use axum::extract::ws::{Message, Utf8Bytes, WebSocket};
use ed25519_dalek::VerifyingKey;
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use tokio::sync::Mutex;
pub mod initialize; pub mod initialize;
pub struct Client {
pub socket: WebSocket,
pub meta: ClientMeta,
pub public_key: VerifyingKey,
}
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClientMeta {} pub struct ClientMeta {}
@@ -55,17 +49,20 @@ pub enum ServerMethod {
}, },
} }
impl Client { pub async fn read_loop(socket: &Arc<Mutex<WebSocket>>) -> anyhow::Result<()> {
pub async fn read_loop(&mut self) -> anyhow::Result<()> { while let Some(message) = read_socket(&mut *socket.lock().await).await? {
while let Some(message) = self.read().await? {
match message { match message {
ServerMethod::Initialize { .. } => { ServerMethod::Initialize { .. } => {
self.send(ClientMethod::Error { send_socket(
&mut *socket.lock().await,
&ClientMethod::Error {
error: Cow::Borrowed("Already initialized"), error: Cow::Borrowed("Already initialized"),
}) },
)
.await?; .await?;
} }
#[allow(unused_variables)]
ServerMethod::Meta(meta) => {} ServerMethod::Meta(meta) => {}
ServerMethod::Error { error } => { ServerMethod::Error { error } => {
@@ -75,25 +72,25 @@ impl Client {
} }
Ok(()) Ok(())
} }
pub async fn read_socket(socket: &mut WebSocket) -> anyhow::Result<Option<ServerMethod>> { pub async fn read_socket(socket: &mut WebSocket) -> anyhow::Result<Option<ServerMethod>> {
match socket.recv().await.transpose()? { match socket.recv().await.transpose()? {
Some(Message::Text(text)) => { Some(Message::Text(text)) => match serde_json::from_str(&text.to_string()) {
if let Ok(msg) = serde_json::from_str(&text.to_string()) { Ok(msg) => Ok(Some(msg)),
Ok(Some(msg))
} else { Err(e) => {
Client::send_socket( send_socket(
socket, socket,
ClientMethod::Error { &ClientMethod::Error {
error: Cow::Borrowed("Unable to parse message: {text}"), error: Cow::Owned(format!("Unable to parse message: {e}")),
}, },
) )
.await?; .await?;
Ok(None) Ok(None)
} }
} },
Some(Message::Ping(v)) => { Some(Message::Ping(v)) => {
socket.send(Message::Pong(v)).await?; socket.send(Message::Pong(v)).await?;
@@ -105,23 +102,14 @@ impl Client {
None => Ok(None), None => Ok(None),
} }
} }
pub async fn send_socket(socket: &mut WebSocket, message: ClientMethod) -> anyhow::Result<()> { pub async fn send_socket(socket: &mut WebSocket, message: &ClientMethod) -> anyhow::Result<()> {
socket socket
.send(Message::Text(Utf8Bytes::from(serde_json::to_string( .send(Message::Text(Utf8Bytes::from(serde_json::to_string(
&message, message,
)?))) )?)))
.await?; .await?;
Ok(()) Ok(())
}
pub async fn read(&mut self) -> anyhow::Result<Option<ServerMethod>> {
Self::read_socket(&mut self.socket).await
}
pub async fn send(&mut self, message: ClientMethod) -> anyhow::Result<()> {
Self::send_socket(&mut self.socket, message).await
}
} }
+63 -6
View File
@@ -1,16 +1,31 @@
use std::sync::Arc; use std::{
collections::HashMap,
sync::{Arc, atomic::AtomicU16},
};
use axum::{ use axum::{
extract::{WebSocketUpgrade, ws::WebSocket}, extract::{WebSocketUpgrade, ws::WebSocket},
response::Response, response::Response,
}; };
use ed25519_dalek::SigningKey; use ed25519_dalek::{SigningKey, VerifyingKey};
use tokio::sync::Mutex;
use crate::{config::Config, protocol::Client}; use crate::{
config::Config,
protocol::{ClientMeta, ClientMethod, read_loop, send_socket},
};
pub struct UserConnections {
pub meta: ClientMeta,
pub counter: AtomicU16,
pub public_key: VerifyingKey,
pub connections: HashMap<u16, Arc<Mutex<WebSocket>>>,
}
pub struct Server { pub struct Server {
pub key: SigningKey, pub key: SigningKey,
pub config: Config, pub config: Config,
pub clients: Mutex<HashMap<VerifyingKey, UserConnections>>,
} }
impl Server { impl Server {
@@ -18,6 +33,7 @@ impl Server {
Ok(Arc::new(Self { Ok(Arc::new(Self {
key: crate::signature::get().await?, key: crate::signature::get().await?,
config: Config::get().await?, config: Config::get().await?,
clients: Mutex::new(HashMap::new()),
})) }))
} }
} }
@@ -27,9 +43,30 @@ 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 Client::initialize(&s, socket).await { match UserConnections::initialize(&s, socket).await {
Ok(mut client) => { Ok((client, public_key, meta)) => {
if let Err(e) = client.read_loop().await { let client = Arc::new(Mutex::new(client));
let mut clients_meta = s.clients.lock().await;
let client_meta =
clients_meta
.entry(public_key)
.or_insert_with(|| UserConnections {
meta,
public_key: public_key,
counter: AtomicU16::new(0),
connections: HashMap::new(),
});
client_meta.connections.insert(
client_meta
.counter
.fetch_add(1, std::sync::atomic::Ordering::Relaxed),
client.clone(),
);
if let Err(e) = read_loop(&client).await {
eprintln!("Failed to handle client: {e}"); eprintln!("Failed to handle client: {e}");
} else { } else {
println!("Client connection closed") println!("Client connection closed")
@@ -43,3 +80,23 @@ impl Server {
}) })
} }
} }
impl UserConnections {
pub async fn send(&self, message: &ClientMethod) -> anyhow::Result<()> {
for (_, conn) in &self.connections {
send_socket(&mut *conn.lock().await, message).await?;
}
Ok(())
}
pub async fn send_to(&self, id: u16, message: &ClientMethod) -> anyhow::Result<bool> {
if let Some(conn) = self.connections.get(&id) {
send_socket(&mut *conn.lock().await, message).await?;
Ok(true)
} else {
Ok(false)
}
}
}