mirror of
https://github.com/moghtech/komodo.git
synced 2026-09-06 16:00:44 +00:00
clean up socket handling
This commit is contained in:
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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>,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user