diff --git a/Cargo.lock b/Cargo.lock index 8ceda8d5a..2778738a1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2833,6 +2833,7 @@ dependencies = [ "serde_qs", "tokio", "tokio-tungstenite", + "tokio-util", "tracing", "wildcard", ] diff --git a/bin/cli/Cargo.toml b/bin/cli/Cargo.toml index 35028b168..8d4fa34e0 100644 --- a/bin/cli/Cargo.toml +++ b/bin/cli/Cargo.toml @@ -24,6 +24,7 @@ noise.workspace = true tokio-tungstenite.workspace = true futures-util.workspace = true comfy-table.workspace = true +tokio-util.workspace = true serde_json.workspace = true crossterm.workspace = true serde_qs.workspace = true diff --git a/bin/cli/src/command/ssh.rs b/bin/cli/src/command/ssh.rs index d5342c15b..ace225271 100644 --- a/bin/cli/src/command/ssh.rs +++ b/bin/cli/src/command/ssh.rs @@ -8,6 +8,7 @@ use komodo_client::{ }; use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _}; use tokio_tungstenite::tungstenite; +use tokio_util::sync::CancellationToken; pub async fn handle( Ssh { @@ -34,36 +35,49 @@ pub async fn handle( // Send first resize messsage, bailing if it fails to get the size. write_tx.send(resize_message()?).await?; + let cancel = CancellationToken::new(); + let forward_resize = async { - while sigwinch.recv().await.is_some() { + while future_or_cancel(sigwinch.recv(), &cancel) + .await + .flatten() + .is_some() + { if let Ok(resize_message) = resize_message() && write_tx.send(resize_message).await.is_err() { break; } } + cancel.cancel(); }; let forward_stdin = async { let mut stdin = tokio::io::stdin(); let mut buf = [0u8; 8192]; - loop { + while let Some(Ok(n)) = future_or_cancel( // Read into buffer starting from index 1, // leaving first byte to represent 'data' message. - let n = match stdin.read(&mut buf[1..]).await { - Ok(0) => break, // EOF - Ok(n) => n, - Err(_) => break, - }; + stdin.read(&mut buf[1..]), + &cancel, + ) + .await + { + // EOF + if n == 0 { + break; + } // Check for disconnect sequence (alt + q) if buf[1..(n + 1)] == [197, 147] { break; } + // Forward bytes let bytes = Bytes::copy_from_slice(&buf[..(n + 1)]); if write_tx.send(bytes).await.is_err() { break; }; } + cancel.cancel(); }; // ===================== @@ -92,71 +106,69 @@ pub async fn handle( .split(); let forward_write = async { - while let Some(bytes) = write_rx.recv().await { + while let Some(bytes) = + future_or_cancel(write_rx.recv(), &cancel).await.flatten() + { if let Err(e) = ws_write.send(tungstenite::Message::Binary(bytes)).await { + cancel.cancel(); return Some(e); }; } + cancel.cancel(); None }; let forward_read = async { let mut stdout = tokio::io::stdout(); - loop { - match ws_read.next().await { - Some(Ok(tungstenite::Message::Binary(bytes))) => { - if let Err(e) = - tokio::io::copy(&mut bytes.as_ref(), &mut stdout) - .await - .context("Failed to copy bytes to stdout") - { - return Some(e); - } - let _ = stdout.flush().await; - } - Some(Ok(tungstenite::Message::Text(text))) => { - if let Err(e) = - tokio::io::copy(&mut text.as_ref(), &mut stdout) - .await - .context("Failed to copy text to stdout") - { - return Some(e); - } - let _ = stdout.flush().await; - } - Some(Ok(tungstenite::Message::Close(_))) => break, - Some(Err(e)) => { + while let Some(msg) = + future_or_cancel(ws_read.next(), &cancel).await.flatten() + { + let bytes = match msg { + Ok(tungstenite::Message::Binary(bytes)) => bytes, + Ok(tungstenite::Message::Text(text)) => text.into(), + Err(e) => { + cancel.cancel(); return Some( anyhow::Error::from(e).context("Websocket read error"), ); } - None => break, - _ => {} + Ok(tungstenite::Message::Close(_)) => break, + _ => continue, + }; + if let Err(e) = stdout + .write_all(&bytes) + .await + .context("Failed to write text to stdout") + { + cancel.cancel(); + return Some(e); } + let _ = stdout.flush().await; } + cancel.cancel(); None }; let guard = RawModeGuard::enable_raw_mode()?; - tokio::select! { - _ = forward_resize => drop(guard), - _ = forward_stdin => drop(guard), - e = forward_write => { - drop(guard); - if let Some(e) = e { - eprintln!("\nFailed to forward stdin | {e:#}"); - } - }, - e = forward_read => { - drop(guard); - if let Some(e) = e { - eprintln!("\nFailed to forward stdout | {e:#}"); - } - }, - }; + let (_, _, write_error, read_error) = tokio::join!( + forward_resize, + forward_stdin, + forward_write, + forward_read + ); + + drop(guard); + + if let Some(e) = write_error { + eprintln!("\nFailed to forward stdin | {e:#}"); + } + + if let Some(e) = read_error { + eprintln!("\nFailed to forward stdout | {e:#}"); + } println!("\n\n{} {}", "connection".bold(), "closed".red().bold()); @@ -191,3 +203,13 @@ impl Drop for RawModeGuard { } } } + +async fn future_or_cancel>( + fut: F, + cancel: &CancellationToken, +) -> Option { + tokio::select! { + res = fut => Some(res), + _ = cancel.cancelled() => None + } +}