From e5822cefb84df83622a657beccb3b1889fd6e56a Mon Sep 17 00:00:00 2001 From: mbecker20 Date: Sat, 27 Sep 2025 12:59:16 -0700 Subject: [PATCH] clean up socket handling --- bin/core/src/connection/client.rs | 91 +++++----- bin/core/src/connection/mod.rs | 221 ++++++++++++------------- bin/core/src/connection/server.rs | 19 ++- bin/periphery/src/connection/client.rs | 10 +- bin/periphery/src/connection/mod.rs | 4 +- bin/periphery/src/connection/server.rs | 14 +- 6 files changed, 184 insertions(+), 175 deletions(-) diff --git a/bin/core/src/connection/client.rs b/bin/core/src/connection/client.rs index a96023ddb..c98eb2566 100644 --- a/bin/core/src/connection/client.rs +++ b/bin/core/src/connection/client.rs @@ -8,14 +8,17 @@ use serror::{deserialize_error_bytes, serialize_error_bytes}; use tokio_tungstenite::Connector; use transport::{ MessageState, - auth::{AddressConnectionIdentifiers, ClientLoginFlow}, + auth::{ + AddressConnectionIdentifiers, ClientLoginFlow, + ConnectionIdentifiers, + }, fix_ws_address, websocket::{Websocket, tungstenite::TungsteniteWebsocket}, }; use crate::{ config::{core_config, core_connection_query}, - connection::PeripheryConnectionArgs, + connection::{PeripheryConnection, PeripheryConnectionArgs}, periphery::ConnectionChannels, state::periphery_connections, }; @@ -37,7 +40,7 @@ impl PeripheryConnectionArgs<'_> { AddressConnectionIdentifiers::extract(&address)?; let endpoint = format!("{address}/?{}", core_connection_query()); - let (connection, mut write_receiver) = + let (connection, mut receiver) = periphery_connections().insert(server_id, self).await; let channels = connection.channels.clone(); @@ -51,7 +54,7 @@ impl PeripheryConnectionArgs<'_> { } }; - let (socket, accept) = match ws { + let (mut socket, accept) = match ws { Ok(res) => res, Err(e) => { connection.set_error(e).await; @@ -63,17 +66,15 @@ impl PeripheryConnectionArgs<'_> { } }; - let mut handler = super::WebsocketHandler { - socket, - connection_identifiers: identifiers.build( - accept.as_bytes(), - core_connection_query().as_bytes(), - ), - write_receiver: &mut write_receiver, - connection: &connection, - }; + let identifiers = identifiers.build( + accept.as_bytes(), + core_connection_query().as_bytes(), + ); - if let Err(e) = handle_login(&mut handler, &passkey).await { + if let Err(e) = connection + .client_login(&mut socket, identifiers, &passkey) + .await + { if connection.cancel.is_cancelled() { break; } @@ -85,7 +86,7 @@ impl PeripheryConnectionArgs<'_> { continue; }; - handler.handle().await + connection.handle_socket(socket, &mut receiver).await } }); @@ -93,35 +94,40 @@ impl PeripheryConnectionArgs<'_> { } } -/// Custom Core -> Periphery side only login wrapper -/// to implement passkey support for backward compatibility -async fn handle_login( - handler: &mut super::WebsocketHandler<'_, TungsteniteWebsocket>, - // for legacy auth - passkey: &str, -) -> anyhow::Result<()> { - // Get the required auth type - let bytes = handler - .socket - .recv_bytes() - .await - .context("Failed to receive login type indicator")?; +impl PeripheryConnection { + /// Custom Core -> Periphery side only login wrapper + /// to implement passkey support for backward compatibility + async fn client_login( + &self, + socket: &mut TungsteniteWebsocket, + identifiers: ConnectionIdentifiers<'_>, + // for legacy auth + passkey: &str, + ) -> anyhow::Result<()> { + // Get the required auth type + let bytes = socket + .recv_bytes() + .await + .context("Failed to receive login type indicator")?; - match bytes.iter().as_slice() { - // Noise auth - &[0] => handler.login::().await, - // Passkey auth - &[1] => handle_passkey_login(handler, passkey).await, - other => { - Err(anyhow!( + match bytes.iter().as_slice() { + // Noise auth + &[0] => { + self + .handle_login::<_, ClientLoginFlow>(socket, identifiers) + .await + } + // Passkey auth + &[1] => handle_passkey_login(socket, passkey).await, + other => Err(anyhow!( "Receieved invalid login type pattern: {other:?}" - )) + )), } } } async fn handle_passkey_login( - handler: &mut super::WebsocketHandler<'_, TungsteniteWebsocket>, + socket: &mut TungsteniteWebsocket, // for legacy auth passkey: &str, ) -> anyhow::Result<()> { @@ -138,15 +144,13 @@ async fn handle_passkey_login( }; passkey.push(MessageState::Successful.as_byte()); - handler - .socket + socket .send(passkey.into()) .await .context("Failed to send passkey")?; // Receive login state message and return based on value - let state_msg = handler - .socket + let state_msg = socket .recv_bytes() .await .context("Failed to receive authentication state message")?; @@ -164,8 +168,7 @@ async fn handle_passkey_login( if let Err(e) = res { let mut bytes = serialize_error_bytes(&e); bytes.push(MessageState::Failed.as_byte()); - if let Err(e) = handler - .socket + if let Err(e) = socket .send(bytes.into()) .await .context("Failed to send login failed to client") @@ -174,7 +177,7 @@ async fn handle_passkey_login( warn!("{e:#}"); } // Close socket - let _ = handler.socket.close(None).await; + let _ = socket.close(None).await; // Return the original error Err(e) } else { diff --git a/bin/core/src/connection/mod.rs b/bin/core/src/connection/mod.rs index 7aa387d70..4da1c83a9 100644 --- a/bin/core/src/connection/mod.rs +++ b/bin/core/src/connection/mod.rs @@ -31,125 +31,11 @@ use crate::{config::core_config, periphery::ConnectionChannels}; pub mod client; pub mod server; -pub struct WebsocketHandler<'a, W> { - pub socket: W, - pub connection_identifiers: ConnectionIdentifiers<'a>, - pub write_receiver: &'a mut BufferedReceiver, - pub connection: &'a PeripheryConnection, -} - -impl WebsocketHandler<'_, W> { - async fn login(&mut self) -> anyhow::Result<()> { - let core_private_key = if let Some(private_key) = - optional_str(&self.connection.core_private_key) - { - private_key - } else { - &core_config().private_key - }; - - let periphery_public_key = - optional_str(&self.connection.periphery_public_key) - .or(core_config().periphery_public_key.as_deref()); - - // Periphery -> Core connection requires a public key pinned - if self.connection.address.is_empty() - && periphery_public_key.is_none() - { - return Err(anyhow!( - "Must either configure Server 'Periphery Public Key' or set KOMODO_PERIPHERY_PUBLIC_KEY for Periphery -> Core connection." - )); - } - - L::login( - &mut self.socket, - self.connection_identifiers, - core_private_key, - &PeripheryPublicKeyValidator { - expected: periphery_public_key, - }, - ) - .await - } - - async fn handle(self) { - let WebsocketHandler { - socket, - write_receiver, - connection, - .. - } = self; - - let handler_cancel = CancellationToken::new(); - - connection.set_connected(true); - connection.clear_error().await; - - let (mut ws_write, mut ws_read) = socket.split(); - - let forward_writes = async { - loop { - let next = tokio::select! { - next = write_receiver.recv() => next, - _ = connection.cancel.cancelled() => break, - _ = handler_cancel.cancelled() => break, - }; - - let message = match next { - Some(request) => Bytes::copy_from_slice(request), - // Sender Dropped (shouldn't happen, a reference is held on 'connection'). - None => break, - }; - - match ws_write.send(message).await { - Ok(_) => write_receiver.clear_buffer(), - Err(e) => { - connection.set_error(e.into()).await; - break; - } - } - } - // Cancel again if not already - let _ = ws_write.close(None).await; - handler_cancel.cancel(); - }; - - let handle_reads = async { - loop { - let next = tokio::select! { - next = ws_read.recv() => next, - _ = connection.cancel.cancelled() => break, - _ = handler_cancel.cancelled() => break, - }; - - match next { - Ok(WebsocketMessage::Binary(bytes)) => { - connection.handle_incoming_bytes(bytes).await - } - Ok(WebsocketMessage::Close(_)) - | Ok(WebsocketMessage::Closed) => { - connection.set_error(anyhow!("Connection closed")).await; - break; - } - Err(e) => { - connection.set_error(e.into()).await; - } - }; - } - // Cancel again if not already - handler_cancel.cancel(); - }; - - tokio::join!(forward_writes, handle_reads); - - connection.set_connected(false); - } -} - pub struct PeripheryPublicKeyValidator<'a> { /// If None, ignore public key. pub expected: Option<&'a str>, } + impl PublicKeyValidator for PeripheryPublicKeyValidator<'_> { fn validate(&self, public_key: String) -> anyhow::Result<()> { if let Some(expected) = self.expected @@ -278,6 +164,111 @@ impl PeripheryConnection { ) } + pub async fn handle_login( + &self, + socket: &mut W, + identifiers: ConnectionIdentifiers<'_>, + ) -> anyhow::Result<()> { + let core_private_key = if let Some(private_key) = + optional_str(&self.core_private_key) + { + private_key + } else { + &core_config().private_key + }; + + let periphery_public_key = + optional_str(&self.periphery_public_key) + .or(core_config().periphery_public_key.as_deref()); + + // Periphery -> Core connection requires a public key pinned + if self.address.is_empty() && periphery_public_key.is_none() { + return Err(anyhow!( + "Must either configure Server 'Periphery Public Key' or set KOMODO_PERIPHERY_PUBLIC_KEY for Periphery -> Core connection." + )); + } + + L::login( + socket, + identifiers, + core_private_key, + &PeripheryPublicKeyValidator { + expected: periphery_public_key, + }, + ) + .await + } + + pub async fn handle_socket( + &self, + socket: W, + receiver: &mut BufferedReceiver, + ) { + let handler_cancel = CancellationToken::new(); + + self.set_connected(true); + self.clear_error().await; + + let (mut ws_write, mut ws_read) = socket.split(); + + let forward_writes = async { + loop { + let next = tokio::select! { + next = receiver.recv() => next, + _ = self.cancel.cancelled() => break, + _ = handler_cancel.cancelled() => break, + }; + + let message = match next { + Some(request) => Bytes::copy_from_slice(request), + // Sender Dropped (shouldn't happen, a reference is held on 'connection'). + None => break, + }; + + match ws_write.send(message).await { + Ok(_) => receiver.clear_buffer(), + Err(e) => { + self.set_error(e.into()).await; + break; + } + } + } + // Cancel again if not already + let _ = ws_write.close(None).await; + handler_cancel.cancel(); + }; + + let handle_reads = async { + loop { + let next = tokio::select! { + next = ws_read.recv() => next, + _ = self.cancel.cancelled() => break, + _ = handler_cancel.cancelled() => break, + }; + + match next { + Ok(WebsocketMessage::Binary(bytes)) => { + self.handle_incoming_bytes(bytes).await + } + Ok(WebsocketMessage::Close(_)) + | Ok(WebsocketMessage::Closed) => { + self.set_error(anyhow!("Connection closed")).await; + break; + } + Err(e) => { + self.set_error(e.into()).await; + } + }; + } + // Cancel again if not already + handler_cancel.cancel(); + }; + + tokio::join!(forward_writes, handle_reads); + + self.set_connected(false); + } + pub async fn handle_incoming_bytes(&self, bytes: Bytes) { let id = match id_from_transport_bytes(&bytes) { Ok(res) => res, diff --git a/bin/core/src/connection/server.rs b/bin/core/src/connection/server.rs index 35b3b2b13..a2a77670f 100644 --- a/bin/core/src/connection/server.rs +++ b/bin/core/src/connection/server.rs @@ -56,7 +56,7 @@ pub async fn handler( ); } - let (connection, mut write_receiver) = periphery_connections() + let (connection, mut receiver) = periphery_connections() .insert( server.id.clone(), PeripheryConnectionArgs { @@ -69,17 +69,18 @@ pub async fn handler( Ok(ws.on_upgrade(|socket| async move { let query = format!("server={}", urlencoding::encode(&_server)); - let mut handler = super::WebsocketHandler { - socket: AxumWebsocket(socket), - connection_identifiers: identifiers.build(query.as_bytes()), - write_receiver: &mut write_receiver, - connection: &connection, - }; + let mut socket = AxumWebsocket(socket); - if let Err(e) = handler.login::().await { + if let Err(e) = connection + .handle_login::<_, ServerLoginFlow>( + &mut socket, + identifiers.build(query.as_bytes()), + ) + .await + { connection.set_error(e).await; } - handler.handle().await + connection.handle_socket(socket, &mut receiver).await })) } diff --git a/bin/periphery/src/connection/client.rs b/bin/periphery/src/connection/client.rs index 8707e4eae..d3268e168 100644 --- a/bin/periphery/src/connection/client.rs +++ b/bin/periphery/src/connection/client.rs @@ -60,7 +60,7 @@ pub async fn handler( info!("Connected to core connection websocket"); } - if let Err(e) = super::login::<_, ClientLoginFlow>( + if let Err(e) = super::handle_login::<_, ClientLoginFlow>( &mut socket, identifiers.build(accept.as_bytes(), query.as_bytes()), ) @@ -79,7 +79,13 @@ pub async fn handler( already_logged_login_error = false; - super::handle(socket, &args, &channel.sender, &mut receiver).await + super::handle_socket( + socket, + &args, + &channel.sender, + &mut receiver, + ) + .await } } diff --git a/bin/periphery/src/connection/mod.rs b/bin/periphery/src/connection/mod.rs index eab68ea9b..83c9906f6 100644 --- a/bin/periphery/src/connection/mod.rs +++ b/bin/periphery/src/connection/mod.rs @@ -86,7 +86,7 @@ impl PublicKeyValidator for CorePublicKeyValidator { } } -async fn login( +async fn handle_login( socket: &mut W, identifiers: ConnectionIdentifiers<'_>, ) -> anyhow::Result<()> { @@ -99,7 +99,7 @@ async fn login( .await } -async fn handle( +async fn handle_socket( socket: W, args: &Arc, sender: &Sender, diff --git a/bin/periphery/src/connection/server.rs b/bin/periphery/src/connection/server.rs index 25d51bcd7..40a6abd3c 100644 --- a/bin/periphery/src/connection/server.rs +++ b/bin/periphery/src/connection/server.rs @@ -87,7 +87,8 @@ async fn handler( let args = Arc::new(Args { core }); - let channel = core_channels().get_or_insert_default(&args.core).await; + let channel = + core_channels().get_or_insert_default(&args.core).await; // Ensure the receiver is free before upgrading connection. // Due to ownership, it needs to be re-locked inside the ws handler, @@ -146,7 +147,13 @@ async fn handler( } }; - super::handle(socket, &args, &channel.sender, &mut receiver).await + super::handle_socket( + socket, + &args, + &channel.sender, + &mut receiver, + ) + .await })) } @@ -165,7 +172,8 @@ async fn handle_login( .send(Bytes::from_owner([0])) .await .context("Failed to send login type indicator")?; - super::login::<_, ServerLoginFlow>(socket, identifiers).await + super::handle_login::<_, ServerLoginFlow>(socket, identifiers) + .await } (None, Some(passkeys)) => { handle_passkey_login(socket, passkeys).await