diff --git a/crates/openshell-supervisor-process/src/ssh.rs b/crates/openshell-supervisor-process/src/ssh.rs index 25c5cfdbbd..19cb557355 100644 --- a/crates/openshell-supervisor-process/src/ssh.rs +++ b/crates/openshell-supervisor-process/src/ssh.rs @@ -3,6 +3,11 @@ //! Embedded SSH server for sandbox access. +mod input; + +#[cfg(test)] +mod exec_input_tests; + use crate::main_session::{MainOutput, MainSession}; #[cfg(unix)] use libc; @@ -17,7 +22,7 @@ use russh::{ChannelId, ChannelOpenFailure, Sig}; use std::borrow::Cow; use std::collections::HashMap; use std::path::{Path, PathBuf}; -use std::sync::{Arc, mpsc}; +use std::sync::Arc; use std::time::Duration; use tokio::net::UnixListener; use tracing::warn; @@ -332,6 +337,7 @@ async fn handle_connection( #[derive(Default)] struct ChannelState { input_sender: Option, + input_task: Option, process: Option>, terminal: Option>, pty_request: Option, @@ -344,14 +350,17 @@ struct ChannelState { } enum InputSender { - Process(mpsc::Sender>), + Process(input::InputSender), Main(tokio::sync::mpsc::Sender>), } impl InputSender { fn send(&self, data: Vec) -> Result<(), &'static str> { match self { - Self::Process(sender) => sender.send(data).map_err(|_| "process stdin closed"), + Self::Process(sender) => sender.send(&data).map_err(|error| match error { + input::SendError::Full => "process stdin buffer is full", + input::SendError::Closed => "process stdin closed", + }), Self::Main(sender) => sender.try_send(data).map_err(|error| match error { tokio::sync::mpsc::error::TrySendError::Full(_) => "canonical stdin buffer is full", tokio::sync::mpsc::error::TrySendError::Closed(_) => { @@ -434,15 +443,15 @@ impl russh::server::Handler for SshHandler { /// Clean up per-channel state when the channel is closed. /// - /// This is the final cleanup and subsumes `channel_eof` — if `channel_close` - /// fires without a preceding `channel_eof`, all resources (`pty_master` File, - /// `input_sender`) are dropped here. + /// Unlike EOF, close cancels pending stdin writes before terminating the + /// process. A child that does not read stdin must not retain queued input. async fn channel_close( &mut self, channel: ChannelId, _session: &mut Session, ) -> Result<(), Self::Error> { if let Some(mut state) = self.channels.remove(&channel) { + state.input_task.take(); if state.main_attached && let Some(main_session) = self.main_session.as_ref() { @@ -824,6 +833,17 @@ impl russh::server::Handler for SshHandler { warn!("data on unknown channel {channel:?}"); return Ok(()); }; + if !state.main_attached { + // Russh replenishes receive credit before this callback. Waiting + // for space here would block signals, close, and output credit on + // the same SSH connection. Reject overflow before copying input. + if let Some(InputSender::Process(sender)) = &state.input_sender + && sender.send(data) == Err(input::SendError::Full) + { + self.reject_exec_input(channel, session)?; + } + return Ok(()); + } // A viewer has no process stdin to interrupt. Ctrl-C closes only its // attachment; the input owner's Ctrl-C still reaches the process. if state.main_attached && state.main_input_owner.is_none() && data.contains(&0x03) { @@ -851,10 +871,9 @@ impl russh::server::Handler for SshHandler { channel: ChannelId, _session: &mut Session, ) -> Result<(), Self::Error> { - // Drop the input sender so the stdin writer thread sees a - // disconnected channel and closes the child's stdin pipe. This - // is essential for commands like `cat | tar xf -` which need - // stdin EOF to know the input stream is complete. + // Drop only the sender: the writer drains accepted bytes and then + // closes stdin. Commands such as `cat | tar xf -` need this half-close + // to finish while their output remains available to the SSH client. if let Some(state) = self.channels.get_mut(&channel) { if state.main_attached && let Some(owner) = state.main_input_owner.take() @@ -923,6 +942,33 @@ impl russh::server::Handler for SshHandler { } impl SshHandler { + fn reject_exec_input( + &mut self, + channel: ChannelId, + session: &mut Session, + ) -> anyhow::Result<()> { + if let Some(mut state) = self.channels.remove(&channel) { + // A locally initiated close may never call channel_close. Release + // input here and keep backend termination off the session loop. + state.input_task.take(); + if let Some(process) = state.process.take() { + tokio::spawn(async move { + if let Err(error) = process.terminate().await { + warn!(%error, "failed to terminate exec after stdin overflow"); + } + }); + } + } + // Session methods enqueue directly. Do not add stderr data here: with + // no output credit, russh would retain it and stop draining Handle + // messages for other channels. Exit status needs no output credit. + warn!(?channel, "process stdin buffer is full; terminating exec"); + session.exit_status_request(channel, 74)?; + session.eof(channel)?; + session.close(channel)?; + Ok(()) + } + async fn start_shell( &mut self, channel: ChannelId, @@ -977,7 +1023,7 @@ impl SshHandler { handle: Handle, spec: openshell_isolation_interface::contract::ExecSpec, ) -> anyhow::Result<()> { - use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::io::AsyncReadExt; let mut exec = self .boundary_exec @@ -992,18 +1038,15 @@ impl SshHandler { state.terminal = exec.terminal.take(); let output_status = exec.output_status.take(); - if let Some(mut stdin) = exec.stdin.take() { - let (sender, receiver) = mpsc::channel::>(); - let runtime = tokio::runtime::Handle::current(); - std::thread::spawn(move || { - while let Ok(bytes) = receiver.recv() { - if runtime.block_on(stdin.write_all(&bytes)).is_err() { - break; - } - } - }); + if let Some(stdin) = exec.stdin.take() { + let (sender, task) = input::InputTask::spawn(stdin); state.input_sender = Some(InputSender::Process(sender)); + state.input_task = Some(task); } + let input_abort = state + .input_task + .as_ref() + .map(input::InputTask::abort_handle); let mut stdout = exec.stdout; let stdout_handle = handle.clone(); @@ -1051,6 +1094,9 @@ impl SshHandler { } status }; + if let Some(input_abort) = input_abort { + input_abort.abort(); + } let code = match status { Some(openshell_isolation_interface::contract::BoundaryExitStatus::Exited(code)) => { code.max(0).cast_unsigned() @@ -1219,9 +1265,8 @@ fn direct_tcpip_target( mod tests { use super::*; use std::io::Write as _; - use std::process::{Command, Stdio}; - struct AcceptAnyServerKey; + pub(super) struct AcceptAnyServerKey; impl russh::client::Handler for AcceptAnyServerKey { type Error = russh::Error; @@ -1281,6 +1326,19 @@ mod tests { async fn main_test_client( main_session: Option>, + ) -> russh::client::Handle { + test_client( + main_session, + Arc::new(RejectingExec), + russh::client::Config::default(), + ) + .await + } + + pub(super) async fn test_client( + main_session: Option>, + boundary_exec: Arc, + client_config: russh::client::Config, ) -> russh::client::Handle { let host_key = { let mut rng = rand::rng(); @@ -1292,11 +1350,7 @@ mod tests { }; server_config.keys.push(host_key); - let handler = SshHandler::new( - Arc::new(TestLoopbackConnector), - Arc::new(RejectingExec), - main_session, - ); + let handler = SshHandler::new(Arc::new(TestLoopbackConnector), boundary_exec, main_session); let (server_stream, client_stream) = tokio::io::duplex(64 * 1024); tokio::spawn(async move { if let Ok(session) = @@ -1307,7 +1361,7 @@ mod tests { }); let mut client = russh::client::connect_stream( - Arc::new(russh::client::Config::default()), + Arc::new(client_config), client_stream, AcceptAnyServerKey, ) @@ -1804,109 +1858,6 @@ mod tests { drop(listener); } - /// Verify that dropping the input sender (the operation `channel_eof` - /// performs) causes the stdin writer loop to exit and close the child's - /// stdin pipe. Without this, commands like `cat | tar xf -` used by - /// `sync --up` hang forever waiting for EOF on stdin. - #[test] - fn dropping_input_sender_closes_child_stdin() { - let (sender, receiver) = mpsc::channel::>(); - - let mut child = Command::new("cat") - .stdin(Stdio::piped()) - .stdout(Stdio::piped()) - .spawn() - .expect("failed to spawn cat"); - - let child_stdin = child.stdin.take().expect("stdin must be piped"); - - // Replicate the stdin writer loop from spawn_pipe_exec. - std::thread::spawn(move || { - let mut stdin = child_stdin; - while let Ok(bytes) = receiver.recv() { - if stdin.write_all(&bytes).is_err() { - break; - } - let _ = stdin.flush(); - } - }); - - sender.send(b"hello".to_vec()).unwrap(); - - // Simulate what channel_eof does: drop the sender. - drop(sender); - - // cat should see EOF on stdin and exit. Use a timeout so the test - // fails fast instead of hanging if the mechanism is broken. - let (done_tx, done_rx) = mpsc::channel(); - std::thread::spawn(move || { - let _ = done_tx.send(child.wait_with_output()); - }); - let output = done_rx - .recv_timeout(Duration::from_secs(5)) - .expect("cat hung for 5s — stdin was not closed (channel_eof bug)") - .expect("failed to wait for cat"); - - assert!( - output.status.success(), - "cat exited with {:?}", - output.status - ); - assert_eq!(output.stdout, b"hello"); - } - - /// Verify that the stdin writer delivers all buffered data before exiting - /// when the sender is dropped. This ensures channel_eof doesn't cause - /// data loss — only signals "no more data after this". - #[test] - fn stdin_writer_delivers_buffered_data_before_eof() { - let (sender, receiver) = mpsc::channel::>(); - - let mut child = Command::new("wc") - .arg("-c") - .stdin(Stdio::piped()) - .stdout(Stdio::piped()) - .spawn() - .expect("failed to spawn wc"); - - let child_stdin = child.stdin.take().expect("stdin must be piped"); - - std::thread::spawn(move || { - let mut stdin = child_stdin; - while let Ok(bytes) = receiver.recv() { - if stdin.write_all(&bytes).is_err() { - break; - } - let _ = stdin.flush(); - } - }); - - // Send multiple chunks, then drop the sender. - for _ in 0..100 { - sender.send(vec![0u8; 1024]).unwrap(); - } - drop(sender); - - let (done_tx, done_rx) = mpsc::channel(); - std::thread::spawn(move || { - let _ = done_tx.send(child.wait_with_output()); - }); - let output = done_rx - .recv_timeout(Duration::from_secs(5)) - .expect("wc hung for 5s — stdin was not closed") - .expect("failed to wait for wc"); - - let count: usize = String::from_utf8_lossy(&output.stdout) - .trim() - .parse() - .expect("wc output was not a number"); - assert_eq!( - count, - 100 * 1024, - "expected all 100 KiB delivered before EOF" - ); - } - // ----------------------------------------------------------------------- // SEC-007: is_loopback_host tests // ----------------------------------------------------------------------- @@ -1956,55 +1907,4 @@ mod tests { assert!(!is_loopback_host("not-an-ip")); assert!(!is_loopback_host("[]")); } - - #[test] - fn channel_state_independent_input_senders() { - // Verify that each channel gets its own input sender so that - // data() and channel_eof() affect only the targeted channel. - let (tx_a, rx_a) = mpsc::channel::>(); - let (tx_b, rx_b) = mpsc::channel::>(); - - let mut state_a = ChannelState { - input_sender: Some(InputSender::Process(tx_a)), - ..Default::default() - }; - let state_b = ChannelState { - input_sender: Some(InputSender::Process(tx_b)), - ..Default::default() - }; - - // Send data to channel A only. - state_a - .input_sender - .as_ref() - .unwrap() - .send(b"hello-a".to_vec()) - .unwrap(); - // Send data to channel B only. - state_b - .input_sender - .as_ref() - .unwrap() - .send(b"hello-b".to_vec()) - .unwrap(); - - assert_eq!(rx_a.recv().unwrap(), b"hello-a"); - assert_eq!(rx_b.recv().unwrap(), b"hello-b"); - - // EOF on channel A (drop sender) should not affect channel B. - state_a.input_sender.take(); - assert!( - rx_a.recv().is_err(), - "channel A sender dropped, recv should fail" - ); - - // Channel B should still be functional. - state_b - .input_sender - .as_ref() - .unwrap() - .send(b"still-alive".to_vec()) - .unwrap(); - assert_eq!(rx_b.recv().unwrap(), b"still-alive"); - } } diff --git a/crates/openshell-supervisor-process/src/ssh/exec_input_tests.rs b/crates/openshell-supervisor-process/src/ssh/exec_input_tests.rs new file mode 100644 index 0000000000..4aea1c4cde --- /dev/null +++ b/crates/openshell-supervisor-process/src/ssh/exec_input_tests.rs @@ -0,0 +1,487 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Real SSH sessions with controlled boundary I/O, including stalled stdin. + +use super::input::MAX_PENDING_INPUT; +use super::tests::test_client; +use openshell_isolation_interface::contract::{ + BackendError, BoundaryExec, BoundaryExitStatus, BoundaryProcess, BoundarySignal, ExecSession, + ExecSpec, +}; +use russh::ChannelMsg; +use std::collections::VecDeque; +use std::pin::Pin; +use std::sync::{Arc, Mutex}; +use std::task::{Context, Poll}; +use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWrite, AsyncWriteExt, DuplexStream}; +use tokio::sync::{mpsc, oneshot, watch}; + +const DEADLINE: Duration = Duration::from_secs(5); + +struct TrackedInput { + stream: DuplexStream, + dropped: Option>, +} + +impl AsyncWrite for TrackedInput { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + bytes: &[u8], + ) -> Poll> { + Pin::new(&mut self.stream).poll_write(cx, bytes) + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.stream).poll_shutdown(cx) + } +} + +impl Drop for TrackedInput { + fn drop(&mut self) { + if let Some(dropped) = self.dropped.take() { + let _ = dropped.send(()); + } + } +} + +struct TestProcess { + status: watch::Sender>, + signals: mpsc::UnboundedSender, + terminated: watch::Sender, +} + +#[async_trait::async_trait] +impl BoundaryProcess for TestProcess { + async fn wait(&self) -> Result { + let mut status = self.status.subscribe(); + Ok(*status + .wait_for(Option::is_some) + .await + .unwrap() + .as_ref() + .unwrap()) + } + + async fn signal(&self, signal: BoundarySignal) -> Result<(), BackendError> { + self.signals.send(signal).unwrap(); + Ok(()) + } + + async fn terminate(&self) -> Result<(), BackendError> { + self.terminated.send_replace(true); + self.status + .send_replace(Some(BoundaryExitStatus::Exited(0))); + Ok(()) + } +} + +struct TestExec(Mutex>); + +#[async_trait::async_trait] +impl BoundaryExec for TestExec { + async fn exec(&self, _spec: ExecSpec) -> Result { + Ok(self + .0 + .lock() + .unwrap() + .pop_front() + .expect("test exec session")) + } +} + +struct Control { + stdin: DuplexStream, + stdout: DuplexStream, + stderr: DuplexStream, + dropped: oneshot::Receiver<()>, + process: Arc, + signals: mpsc::UnboundedReceiver, +} + +fn exec_session() -> (ExecSession, Control) { + exec_session_with_capacity(1) +} + +fn exec_session_with_capacity(capacity: usize) -> (ExecSession, Control) { + // A one-byte sink makes every ordinary SSH packet stall in write_all until + // the test starts reading. Retain the read end to distinguish cancellation + // from a broken pipe caused by the fixture closing its own input. + let (stdin, stdin_reader) = tokio::io::duplex(capacity); + let (stdout, stdout_writer) = tokio::io::duplex(64 * 1024); + let (stderr, stderr_writer) = tokio::io::duplex(64 * 1024); + let (dropped_tx, dropped) = oneshot::channel(); + let (signals_tx, signals) = mpsc::unbounded_channel(); + let process = Arc::new(TestProcess { + status: watch::channel(None).0, + signals: signals_tx, + terminated: watch::channel(false).0, + }); + ( + ExecSession { + process: process.clone(), + stdin: Some(Box::new(TrackedInput { + stream: stdin, + dropped: Some(dropped_tx), + })), + stdout: Box::new(stdout), + stderr: Some(Box::new(stderr)), + terminal: None, + output_status: None, + }, + Control { + stdin: stdin_reader, + stdout: stdout_writer, + stderr: stderr_writer, + dropped, + process, + signals, + }, + ) +} + +async fn start_exec( + client: &russh::client::Handle, +) -> russh::Channel { + let mut channel = client.channel_open_session().await.unwrap(); + channel.exec(true, "cat").await.unwrap(); + assert!(matches!(channel.wait().await, Some(ChannelMsg::Success))); + channel +} + +#[tokio::test] +async fn full_input_allows_output_signals_and_channel_close() { + tokio::time::timeout(DEADLINE, async { + let (session, mut control) = exec_session(); + let client = test_client( + None, + Arc::new(TestExec(Mutex::new([session].into()))), + russh::client::Config::default(), + ) + .await; + let mut channel = start_exec(&client).await; + channel + .data(vec![7; MAX_PENDING_INPUT].as_slice()) + .await + .unwrap(); + // The signal follows all input packets on the wire and acts as a + // barrier proving that the handler accepted the full pending budget. + channel.signal(russh::Sig::INT).await.unwrap(); + assert_eq!(control.signals.recv().await, Some(BoundarySignal::Int)); + // Exceed the default SSH output window so this also requires incoming + // window adjustments while stdin remains full. + let output = tokio::spawn(async move { + control + .stdout + .write_all(&vec![9; 3 * 1024 * 1024]) + .await + .unwrap(); + }); + control + .stderr + .write_all(b"stderr while stdin is full") + .await + .unwrap(); + let mut stdout = Vec::new(); + let mut stderr = Vec::new(); + while stdout.len() < 3 * 1024 * 1024 || stderr.len() < 26 { + match channel.wait().await.unwrap() { + ChannelMsg::Data { data } => stdout.extend_from_slice(&data), + ChannelMsg::ExtendedData { data, ext: 1 } => stderr.extend_from_slice(&data), + ChannelMsg::WindowAdjusted { .. } => {} + message => panic!("unexpected message: {message:?}"), + } + } + assert_eq!(stdout, vec![9; 3 * 1024 * 1024]); + assert_eq!(stderr, b"stderr while stdin is full"); + output.await.unwrap(); + channel.close().await.unwrap(); + control.dropped.await.expect("close releases blocked stdin"); + control + .process + .terminated + .subscribe() + .wait_for(|done| *done) + .await + .unwrap(); + }) + .await + .expect("full stdin blocked output, signal, or close"); +} + +#[tokio::test] +async fn disconnect_cancels_stalled_input_without_main_session() { + tokio::time::timeout(DEADLINE, async { + let (session, mut control) = exec_session(); + let client = test_client( + None, + Arc::new(TestExec(Mutex::new([session].into()))), + russh::client::Config::default(), + ) + .await; + let channel = start_exec(&client).await; + channel + .data(vec![7; MAX_PENDING_INPUT].as_slice()) + .await + .unwrap(); + channel.signal(russh::Sig::INT).await.unwrap(); + control.signals.recv().await.unwrap(); + client + .disconnect(russh::Disconnect::ByApplication, "test disconnect", "") + .await + .unwrap(); + control + .dropped + .await + .expect("disconnect releases blocked stdin"); + }) + .await + .expect("disconnect retained a blocked stdin writer"); +} + +#[tokio::test] +async fn process_exit_cancels_stalled_input() { + tokio::time::timeout(DEADLINE, async { + let (session, mut control) = exec_session(); + let client = test_client( + None, + Arc::new(TestExec(Mutex::new([session].into()))), + russh::client::Config::default(), + ) + .await; + let mut channel = start_exec(&client).await; + channel + .data(vec![7; MAX_PENDING_INPUT].as_slice()) + .await + .unwrap(); + channel.signal(russh::Sig::INT).await.unwrap(); + control.signals.recv().await.unwrap(); + control + .process + .status + .send_replace(Some(BoundaryExitStatus::Exited(0))); + drop(control.stdout); + drop(control.stderr); + control + .dropped + .await + .expect("process exit releases blocked stdin"); + while let Some(message) = channel.wait().await { + if let ChannelMsg::ExitStatus { exit_status } = message { + assert_eq!(exit_status, 0); + return; + } + } + panic!("missing process exit status"); + }) + .await + .expect("process exit retained a blocked stdin writer"); +} + +#[tokio::test] +async fn overflow_reports_failure_and_preserves_sibling_channel() { + tokio::time::timeout(DEADLINE, async { + let (first, mut a) = exec_session(); + let (second, mut b) = exec_session(); + let client = test_client( + None, + Arc::new(TestExec(Mutex::new([first, second].into()))), + russh::client::Config::default(), + ) + .await; + let mut channel = start_exec(&client).await; + let sibling = start_exec(&client).await; + channel + .data(vec![7; MAX_PENDING_INPUT].as_slice()) + .await + .unwrap(); + channel.signal(russh::Sig::INT).await.unwrap(); + a.signals.recv().await.unwrap(); + channel.data(&b"overflow"[..]).await.unwrap(); + a.dropped.await.expect("overflow releases blocked stdin"); + a.process + .terminated + .subscribe() + .wait_for(|done| *done) + .await + .unwrap(); + // Termination reports success in this fixture; the SSH result must + // still report rejected input, with no later successful exit status. + drop(a.stdout); + drop(a.stderr); + let mut statuses = Vec::new(); + let mut stderr = Vec::new(); + while let Some(message) = channel.wait().await { + match message { + ChannelMsg::ExitStatus { exit_status } => statuses.push(exit_status), + ChannelMsg::ExtendedData { data, ext: 1 } => stderr.extend_from_slice(&data), + ChannelMsg::Close => break, + _ => {} + } + } + assert_eq!(statuses, [74]); + assert!(stderr.is_empty()); + sibling.data(&b"still usable"[..]).await.unwrap(); + sibling.eof().await.unwrap(); + let mut received = Vec::new(); + b.stdin.read_to_end(&mut received).await.unwrap(); + assert_eq!(received, b"still usable"); + sibling.close().await.unwrap(); + }) + .await + .expect("overflow blocked cleanup or a sibling channel"); +} + +#[tokio::test] +async fn eof_drains_accepted_input_and_allows_larger_total_transfers() { + tokio::time::timeout(Duration::from_secs(20), async { + let (session, mut control) = exec_session_with_capacity(64 * 1024); + let client = test_client( + None, + Arc::new(TestExec(Mutex::new([session].into()))), + russh::client::Config::default(), + ) + .await; + let mut channel = start_exec(&client).await; + // Consume each batch before sending the next, allowing a transfer + // larger than the pending-byte budget without relying on scheduling. + let payload = vec![42; 64 * 1024]; + let (read_tx, mut read_rx) = mpsc::channel(1); + let reader = tokio::spawn(async move { + let mut total = 0; + let mut received = vec![0; 64 * 1024]; + for _ in 0..80 { + control.stdin.read_exact(&mut received).await.unwrap(); + assert!(received.iter().all(|byte| *byte == 42)); + total += received.len(); + read_tx.send(()).await.unwrap(); + } + let mut tail = Vec::new(); + control.stdin.read_to_end(&mut tail).await.unwrap(); + assert_eq!(tail, b"accepted before EOF"); + total + }); + for _ in 0..80 { + channel.data(payload.as_slice()).await.unwrap(); + read_rx.recv().await.unwrap(); + } + channel.data(&b"accepted before EOF"[..]).await.unwrap(); + channel.eof().await.unwrap(); + assert_eq!(reader.await.unwrap(), 5 * 1024 * 1024); + control.stdout.write_all(b"complete").await.unwrap(); + drop(control.stdout); + drop(control.stderr); + control + .process + .status + .send_replace(Some(BoundaryExitStatus::Exited(0))); + let mut output = Vec::new(); + let mut statuses = Vec::new(); + while let Some(message) = channel.wait().await { + match message { + ChannelMsg::Data { data } => output.extend_from_slice(&data), + ChannelMsg::ExitStatus { exit_status } => statuses.push(exit_status), + ChannelMsg::Close => break, + _ => {} + } + } + assert_eq!(output, b"complete"); + assert_eq!(statuses, [0]); + }) + .await + .expect("EOF did not drain accepted input"); +} + +#[tokio::test] +async fn eof_preserves_queued_input_until_the_child_reads() { + tokio::time::timeout(DEADLINE, async { + let (session, mut control) = exec_session(); + let client = test_client( + None, + Arc::new(TestExec(Mutex::new([session].into()))), + russh::client::Config::default(), + ) + .await; + let channel = start_exec(&client).await; + channel.data(vec![42; 64 * 1024].as_slice()).await.unwrap(); + channel.eof().await.unwrap(); + channel.signal(russh::Sig::INT).await.unwrap(); + control.signals.recv().await.unwrap(); + // EOF has reached the handler while write_all is still blocked. Only + // now let the child consume the accepted input and observe stdin EOF. + let mut received = Vec::new(); + control.stdin.read_to_end(&mut received).await.unwrap(); + assert_eq!(received, vec![42; 64 * 1024]); + channel.close().await.unwrap(); + }) + .await + .expect("EOF discarded or retained queued input"); +} + +#[tokio::test] +async fn overflow_reports_failure_without_output_window_credit() { + tokio::time::timeout(DEADLINE, async { + let (session, mut control) = exec_session(); + let (second, sibling_control) = exec_session(); + let client = test_client( + None, + Arc::new(TestExec(Mutex::new([session, second].into()))), + russh::client::Config { + window_size: 0, + ..Default::default() + }, + ) + .await; + let mut channel = start_exec(&client).await; + let mut sibling = start_exec(&client).await; + channel + .data(vec![7; MAX_PENDING_INPUT].as_slice()) + .await + .unwrap(); + channel.signal(russh::Sig::INT).await.unwrap(); + control.signals.recv().await.unwrap(); + // Overflow must not queue a diagnostic behind the zero output window. + // That would block Handle messages, including the sibling's status. + channel.data(&b"overflow"[..]).await.unwrap(); + control.dropped.await.unwrap(); + loop { + match channel.wait().await { + Some(ChannelMsg::WindowAdjusted { .. }) => {} + Some(ChannelMsg::ExitStatus { exit_status: 74 }) => break, + message => panic!("expected failure before output credit: {message:?}"), + } + } + control + .process + .terminated + .subscribe() + .wait_for(|done| *done) + .await + .unwrap(); + sibling_control + .process + .status + .send_replace(Some(BoundaryExitStatus::Exited(0))); + drop(sibling_control.stdout); + drop(sibling_control.stderr); + loop { + match sibling.wait().await { + Some(ChannelMsg::ExitStatus { exit_status: 0 }) => break, + Some(ChannelMsg::Eof | ChannelMsg::WindowAdjusted { .. }) => {} + message => panic!("sibling did not complete: {message:?}"), + } + } + client + .disconnect(russh::Disconnect::ByApplication, "done", "") + .await + .unwrap(); + }) + .await + .expect("overflow waited for output credit before reporting failure"); +} diff --git a/crates/openshell-supervisor-process/src/ssh/input.rs b/crates/openshell-supervisor-process/src/ssh/input.rs new file mode 100644 index 0000000000..391fa62b70 --- /dev/null +++ b/crates/openshell-supervisor-process/src/ssh/input.rs @@ -0,0 +1,185 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Bounded stdin delivery without blocking the shared SSH session callback. + +use std::collections::VecDeque; +use std::sync::{Arc, Mutex}; + +use openshell_isolation_interface::contract::BoundaryInput; +use tokio::io::AsyncWriteExt; +use tokio::sync::Notify; +use tokio::task::{AbortHandle, JoinHandle}; + +/// Includes queued bytes and the chunk currently being written. Consumed bytes +/// release capacity, so this is not a limit on the total command input size. +pub(super) const MAX_PENDING_INPUT: usize = 4 * 1024 * 1024; +const WRITE_CHUNK: usize = 16 * 1024; + +#[derive(Default)] +struct Buffer { + // Coalescing packets avoids unbounded per-message overhead for tiny writes. + bytes: VecDeque, + in_flight: usize, + closed: bool, +} + +#[derive(Default)] +struct Shared { + buffer: Mutex, + ready: Notify, +} + +pub(super) struct InputSender(Arc); + +#[derive(Debug, PartialEq, Eq)] +pub(super) enum SendError { + Full, + Closed, +} + +impl InputSender { + pub(super) fn send(&self, bytes: &[u8]) -> Result<(), SendError> { + let mut buffer = self.0.buffer.lock().map_err(|_| SendError::Closed)?; + if buffer.closed { + return Err(SendError::Closed); + } + if bytes.len() > MAX_PENDING_INPUT - buffer.bytes.len() - buffer.in_flight { + return Err(SendError::Full); + } + buffer.bytes.extend(bytes); + drop(buffer); + self.0.ready.notify_one(); + Ok(()) + } +} + +impl Drop for InputSender { + fn drop(&mut self) { + if let Ok(mut buffer) = self.0.buffer.lock() { + buffer.closed = true; + } + // EOF wakes an idle writer but leaves accepted bytes to drain. + self.0.ready.notify_one(); + } +} + +/// The channel owns cancellation independently of stdin EOF. Dropping a task +/// handle alone would detach it and retain a blocked write and its queued input. +pub(super) struct InputTask { + shared: Arc, + task: JoinHandle<()>, +} + +impl InputTask { + pub(super) fn spawn(stdin: BoundaryInput) -> (InputSender, Self) { + let shared = Arc::new(Shared::default()); + let receiver = Receiver(shared.clone()); + let task = tokio::spawn(receiver.write_to(stdin)); + (InputSender(shared.clone()), Self { shared, task }) + } + + pub(super) fn abort_handle(&self) -> AbortHandle { + self.task.abort_handle() + } +} + +impl Drop for InputTask { + fn drop(&mut self) { + // Release queued allocations even before the runtime polls the aborted + // writer. Its in-flight chunk and stdin are released when it is dropped. + if let Ok(mut buffer) = self.shared.buffer.lock() { + buffer.closed = true; + buffer.bytes = VecDeque::new(); + } + self.task.abort(); + } +} + +struct Receiver(Arc); + +impl Receiver { + async fn write_to(self, mut stdin: BoundaryInput) { + loop { + let ready = self.0.ready.notified(); + let chunk = { + let Ok(mut buffer) = self.0.buffer.lock() else { + break; + }; + if buffer.bytes.is_empty() && buffer.closed { + break; + } + let size = buffer.bytes.len().min(WRITE_CHUNK); + buffer.in_flight = size; + buffer.bytes.drain(..size).collect::>() + }; + if chunk.is_empty() { + ready.await; + continue; + } + if stdin.write_all(&chunk).await.is_err() { + break; + } + let Ok(mut buffer) = self.0.buffer.lock() else { + break; + }; + buffer.in_flight = 0; + } + } +} + +impl Drop for Receiver { + fn drop(&mut self) { + // A failed or cancelled writer must refuse later input and release the + // queue even if the SSH channel has not yet received a close packet. + if let Ok(mut buffer) = self.0.buffer.lock() { + buffer.closed = true; + buffer.in_flight = 0; + buffer.bytes = VecDeque::new(); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tokio::io::AsyncReadExt; + + #[tokio::test] + async fn bound_includes_in_flight_input_and_rejects_before_copying() { + let (stdin, mut reader) = tokio::io::duplex(1); + let (sender, _task) = InputTask::spawn(Box::new(stdin)); + assert_eq!( + sender.send(&vec![0; MAX_PENDING_INPUT + 1]), + Err(SendError::Full) + ); + assert_eq!(sender.0.buffer.lock().unwrap().bytes.capacity(), 0); + sender.send(&vec![42; MAX_PENDING_INPUT]).unwrap(); + // Once the one-byte pipe is full, the rest of this chunk is retained + // by write_all and must still count against the pending-byte budget. + reader.read_u8().await.unwrap(); + assert_eq!(sender.send(&[1]), Err(SendError::Full)); + let buffer = sender.0.buffer.lock().unwrap(); + assert_eq!(buffer.bytes.len() + buffer.in_flight, MAX_PENDING_INPUT); + assert_eq!(buffer.in_flight, WRITE_CHUNK); + } + + #[tokio::test] + async fn tiny_packets_share_one_byte_buffer_and_eof_drains_them() { + let (stdin, mut reader) = tokio::io::duplex(64 * 1024); + let (sender, _task) = InputTask::spawn(Box::new(stdin)); + for _ in 0..8192 { + sender.send(&[42]).unwrap(); + } + drop(sender); + let mut received = Vec::new(); + tokio::time::timeout( + std::time::Duration::from_secs(5), + reader.read_to_end(&mut received), + ) + .await + .unwrap() + .unwrap(); + assert_eq!(received, vec![42; 8192]); + } +} diff --git a/docs/how-it-works/sandboxes/overview.mdx b/docs/how-it-works/sandboxes/overview.mdx index b20756c2e3..1ed1aa31a6 100644 --- a/docs/how-it-works/sandboxes/overview.mdx +++ b/docs/how-it-works/sandboxes/overview.mdx @@ -305,6 +305,8 @@ and receive the exit event and final gRPC status. Input EOF does not immediately close the SSH output channel and is distinct from cancelling the RPC. With a PTY, input closure is not equivalent to sending a terminal Ctrl-D keystroke. +The supervisor limits pending stdin to 4 MiB per exec channel. If input arrives faster than the command consumes it and this buffer fills, further input fails the exec with exit code 74. The supervisor logs `process stdin buffer is full`. This also applies to SSH file transfers that use exec channels. The limit counts input waiting to be written, so clients that support larger transfers can continue as the command consumes data. Input EOF drains bytes already accepted; cancelling the exec discards pending input. + Run a one-shot command inside a running sandbox without opening an interactive shell: ```shell