clean up socket handling

This commit is contained in:
mbecker20
2025-09-27 14:23:49 -07:00
parent 4baab194cf
commit e5822cefb8
6 changed files with 184 additions and 175 deletions
+47 -44
View File
@@ -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::<ClientLoginFlow>().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 {
+106 -115
View File
@@ -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<Bytes>,
pub connection: &'a PeripheryConnection,
}
impl<W: Websocket> WebsocketHandler<'_, W> {
async fn login<L: LoginFlow>(&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<W: Websocket, L: LoginFlow>(
&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<W: Websocket>(
&self,
socket: W,
receiver: &mut BufferedReceiver<Bytes>,
) {
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,
+10 -9
View File
@@ -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::<ServerLoginFlow>().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
}))
}
+8 -2
View File
@@ -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
}
}
+2 -2
View File
@@ -86,7 +86,7 @@ impl PublicKeyValidator for CorePublicKeyValidator {
}
}
async fn login<W: Websocket, L: LoginFlow>(
async fn handle_login<W: Websocket, L: LoginFlow>(
socket: &mut W,
identifiers: ConnectionIdentifiers<'_>,
) -> anyhow::Result<()> {
@@ -99,7 +99,7 @@ async fn login<W: Websocket, L: LoginFlow>(
.await
}
async fn handle<W: Websocket>(
async fn handle_socket<W: Websocket>(
socket: W,
args: &Arc<Args>,
sender: &Sender<Bytes>,
+11 -3
View File
@@ -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