Skip to content

Commit 8d418f1

Browse files
authored
fix(supervisor): bound pending exec stdin and cancel stalled writers (#3846)
* fix(supervisor): bound pending exec stdin and cancel stalled writers Signed-off-by: Shiju <shiju@nvidia.com> * docs(supervisor): separate pending stdin guidance from CLI modes Keep the pending-input limit beside the RPC lifecycle contract so the streaming CLI documentation can merge independently. Signed-off-by: Shiju <shiju@nvidia.com> --------- Signed-off-by: Shiju <shiju@nvidia.com>
1 parent 6e369f2 commit 8d418f1

4 files changed

Lines changed: 756 additions & 182 deletions

File tree

‎crates/openshell-supervisor-process/src/ssh.rs‎

Lines changed: 82 additions & 182 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,11 @@
33

44
//! Embedded SSH server for sandbox access.
55
6+
mod input;
7+
8+
#[cfg(test)]
9+
mod exec_input_tests;
10+
611
use crate::main_session::{MainOutput, MainSession};
712
#[cfg(unix)]
813
use libc;
@@ -17,7 +22,7 @@ use russh::{ChannelId, ChannelOpenFailure, Sig};
1722
use std::borrow::Cow;
1823
use std::collections::HashMap;
1924
use std::path::{Path, PathBuf};
20-
use std::sync::{Arc, mpsc};
25+
use std::sync::Arc;
2126
use std::time::Duration;
2227
use tokio::net::UnixListener;
2328
use tracing::warn;
@@ -353,6 +358,7 @@ async fn handle_connection(
353358
#[derive(Default)]
354359
struct ChannelState {
355360
input_sender: Option<InputSender>,
361+
input_task: Option<input::InputTask>,
356362
process: Option<Arc<dyn openshell_isolation_interface::contract::BoundaryProcess>>,
357363
terminal: Option<Arc<dyn openshell_isolation_interface::contract::BoundaryTerminal>>,
358364
pty_request: Option<PtyRequest>,
@@ -368,14 +374,17 @@ struct ChannelState {
368374
}
369375

370376
enum InputSender {
371-
Process(mpsc::Sender<Vec<u8>>),
377+
Process(input::InputSender),
372378
Main(tokio::sync::mpsc::Sender<Vec<u8>>),
373379
}
374380

375381
impl InputSender {
376382
fn send(&self, data: Vec<u8>) -> Result<(), &'static str> {
377383
match self {
378-
Self::Process(sender) => sender.send(data).map_err(|_| "process stdin closed"),
384+
Self::Process(sender) => sender.send(&data).map_err(|error| match error {
385+
input::SendError::Full => "process stdin buffer is full",
386+
input::SendError::Closed => "process stdin closed",
387+
}),
379388
Self::Main(sender) => sender.try_send(data).map_err(|error| match error {
380389
tokio::sync::mpsc::error::TrySendError::Full(_) => "canonical stdin buffer is full",
381390
tokio::sync::mpsc::error::TrySendError::Closed(_) => {
@@ -458,15 +467,15 @@ impl russh::server::Handler for SshHandler {
458467

459468
/// Clean up per-channel state when the channel is closed.
460469
///
461-
/// This is the final cleanup and subsumes `channel_eof` — if `channel_close`
462-
/// fires without a preceding `channel_eof`, all resources (`pty_master` File,
463-
/// `input_sender`) are dropped here.
470+
/// Unlike EOF, close cancels pending stdin writes before terminating the
471+
/// process. A child that does not read stdin must not retain queued input.
464472
async fn channel_close(
465473
&mut self,
466474
channel: ChannelId,
467475
_session: &mut Session,
468476
) -> Result<(), Self::Error> {
469477
if let Some(mut state) = self.channels.remove(&channel) {
478+
state.input_task.take();
470479
if state.main_attached
471480
&& let Some(main_session) = self.main_session.as_ref()
472481
{
@@ -849,6 +858,17 @@ impl russh::server::Handler for SshHandler {
849858
warn!("data on unknown channel {channel:?}");
850859
return Ok(());
851860
};
861+
if !state.main_attached {
862+
// Russh replenishes receive credit before this callback. Waiting
863+
// for space here would block signals, close, and output credit on
864+
// the same SSH connection. Reject overflow before copying input.
865+
if let Some(InputSender::Process(sender)) = &state.input_sender
866+
&& sender.send(data) == Err(input::SendError::Full)
867+
{
868+
self.reject_exec_input(channel, session)?;
869+
}
870+
return Ok(());
871+
}
852872
// A viewer has no process stdin to interrupt. Ctrl-C closes only its
853873
// attachment; the input owner's Ctrl-C still reaches the process.
854874
if state.main_attached && state.main_input_owner.is_none() && data.contains(&0x03) {
@@ -900,10 +920,9 @@ impl russh::server::Handler for SshHandler {
900920
channel: ChannelId,
901921
_session: &mut Session,
902922
) -> Result<(), Self::Error> {
903-
// Drop the input sender so the stdin writer thread sees a
904-
// disconnected channel and closes the child's stdin pipe. This
905-
// is essential for commands like `cat | tar xf -` which need
906-
// stdin EOF to know the input stream is complete.
923+
// Drop only the sender: the writer drains accepted bytes and then
924+
// closes stdin. Commands such as `cat | tar xf -` need this half-close
925+
// to finish while their output remains available to the SSH client.
907926
if let Some(state) = self.channels.get_mut(&channel) {
908927
if state.main_attached
909928
&& let Some(owner) = state.main_input_owner.take()
@@ -973,6 +992,33 @@ impl russh::server::Handler for SshHandler {
973992
}
974993

975994
impl SshHandler {
995+
fn reject_exec_input(
996+
&mut self,
997+
channel: ChannelId,
998+
session: &mut Session,
999+
) -> anyhow::Result<()> {
1000+
if let Some(mut state) = self.channels.remove(&channel) {
1001+
// A locally initiated close may never call channel_close. Release
1002+
// input here and keep backend termination off the session loop.
1003+
state.input_task.take();
1004+
if let Some(process) = state.process.take() {
1005+
tokio::spawn(async move {
1006+
if let Err(error) = process.terminate().await {
1007+
warn!(%error, "failed to terminate exec after stdin overflow");
1008+
}
1009+
});
1010+
}
1011+
}
1012+
// Session methods enqueue directly. Do not add stderr data here: with
1013+
// no output credit, russh would retain it and stop draining Handle
1014+
// messages for other channels. Exit status needs no output credit.
1015+
warn!(?channel, "process stdin buffer is full; terminating exec");
1016+
session.exit_status_request(channel, 74)?;
1017+
session.eof(channel)?;
1018+
session.close(channel)?;
1019+
Ok(())
1020+
}
1021+
9761022
async fn start_shell(
9771023
&mut self,
9781024
channel: ChannelId,
@@ -1027,7 +1073,7 @@ impl SshHandler {
10271073
handle: Handle,
10281074
spec: openshell_isolation_interface::contract::ExecSpec,
10291075
) -> anyhow::Result<()> {
1030-
use tokio::io::{AsyncReadExt, AsyncWriteExt};
1076+
use tokio::io::AsyncReadExt;
10311077

10321078
let mut exec = self
10331079
.boundary_exec
@@ -1042,18 +1088,15 @@ impl SshHandler {
10421088
state.terminal = exec.terminal.take();
10431089
let output_status = exec.output_status.take();
10441090

1045-
if let Some(mut stdin) = exec.stdin.take() {
1046-
let (sender, receiver) = mpsc::channel::<Vec<u8>>();
1047-
let runtime = tokio::runtime::Handle::current();
1048-
std::thread::spawn(move || {
1049-
while let Ok(bytes) = receiver.recv() {
1050-
if runtime.block_on(stdin.write_all(&bytes)).is_err() {
1051-
break;
1052-
}
1053-
}
1054-
});
1091+
if let Some(stdin) = exec.stdin.take() {
1092+
let (sender, task) = input::InputTask::spawn(stdin);
10551093
state.input_sender = Some(InputSender::Process(sender));
1094+
state.input_task = Some(task);
10561095
}
1096+
let input_abort = state
1097+
.input_task
1098+
.as_ref()
1099+
.map(input::InputTask::abort_handle);
10571100

10581101
let mut stdout = exec.stdout;
10591102
let stdout_handle = handle.clone();
@@ -1101,6 +1144,9 @@ impl SshHandler {
11011144
}
11021145
status
11031146
};
1147+
if let Some(input_abort) = input_abort {
1148+
input_abort.abort();
1149+
}
11041150
let code = match status {
11051151
Some(openshell_isolation_interface::contract::BoundaryExitStatus::Exited(code)) => {
11061152
code.max(0).cast_unsigned()
@@ -1270,7 +1316,6 @@ fn direct_tcpip_target(
12701316
mod tests {
12711317
use super::*;
12721318
use std::io::Write as _;
1273-
use std::process::{Command, Stdio};
12741319

12751320
pub(super) struct AcceptAnyServerKey;
12761321

@@ -1332,6 +1377,19 @@ mod tests {
13321377

13331378
async fn main_test_client(
13341379
main_session: Option<Arc<MainSession>>,
1380+
) -> russh::client::Handle<AcceptAnyServerKey> {
1381+
test_client(
1382+
main_session,
1383+
Arc::new(RejectingExec),
1384+
russh::client::Config::default(),
1385+
)
1386+
.await
1387+
}
1388+
1389+
pub(super) async fn test_client(
1390+
main_session: Option<Arc<MainSession>>,
1391+
boundary_exec: Arc<dyn openshell_isolation_interface::contract::BoundaryExec>,
1392+
client_config: russh::client::Config,
13351393
) -> russh::client::Handle<AcceptAnyServerKey> {
13361394
let host_key = {
13371395
let mut rng = rand::rng();
@@ -1343,11 +1401,7 @@ mod tests {
13431401
};
13441402
server_config.keys.push(host_key);
13451403

1346-
let handler = SshHandler::new(
1347-
Arc::new(TestLoopbackConnector),
1348-
Arc::new(RejectingExec),
1349-
main_session,
1350-
);
1404+
let handler = SshHandler::new(Arc::new(TestLoopbackConnector), boundary_exec, main_session);
13511405
let (server_stream, client_stream) = tokio::io::duplex(64 * 1024);
13521406
tokio::spawn(async move {
13531407
if let Ok(session) =
@@ -1358,7 +1412,7 @@ mod tests {
13581412
});
13591413

13601414
let mut client = russh::client::connect_stream(
1361-
Arc::new(russh::client::Config::default()),
1415+
Arc::new(client_config),
13621416
client_stream,
13631417
AcceptAnyServerKey,
13641418
)
@@ -1855,109 +1909,6 @@ mod tests {
18551909
drop(listener);
18561910
}
18571911

1858-
/// Verify that dropping the input sender (the operation `channel_eof`
1859-
/// performs) causes the stdin writer loop to exit and close the child's
1860-
/// stdin pipe. Without this, commands like `cat | tar xf -` used by
1861-
/// `sync --up` hang forever waiting for EOF on stdin.
1862-
#[test]
1863-
fn dropping_input_sender_closes_child_stdin() {
1864-
let (sender, receiver) = mpsc::channel::<Vec<u8>>();
1865-
1866-
let mut child = Command::new("cat")
1867-
.stdin(Stdio::piped())
1868-
.stdout(Stdio::piped())
1869-
.spawn()
1870-
.expect("failed to spawn cat");
1871-
1872-
let child_stdin = child.stdin.take().expect("stdin must be piped");
1873-
1874-
// Replicate the stdin writer loop from spawn_pipe_exec.
1875-
std::thread::spawn(move || {
1876-
let mut stdin = child_stdin;
1877-
while let Ok(bytes) = receiver.recv() {
1878-
if stdin.write_all(&bytes).is_err() {
1879-
break;
1880-
}
1881-
let _ = stdin.flush();
1882-
}
1883-
});
1884-
1885-
sender.send(b"hello".to_vec()).unwrap();
1886-
1887-
// Simulate what channel_eof does: drop the sender.
1888-
drop(sender);
1889-
1890-
// cat should see EOF on stdin and exit. Use a timeout so the test
1891-
// fails fast instead of hanging if the mechanism is broken.
1892-
let (done_tx, done_rx) = mpsc::channel();
1893-
std::thread::spawn(move || {
1894-
let _ = done_tx.send(child.wait_with_output());
1895-
});
1896-
let output = done_rx
1897-
.recv_timeout(Duration::from_secs(5))
1898-
.expect("cat hung for 5s — stdin was not closed (channel_eof bug)")
1899-
.expect("failed to wait for cat");
1900-
1901-
assert!(
1902-
output.status.success(),
1903-
"cat exited with {:?}",
1904-
output.status
1905-
);
1906-
assert_eq!(output.stdout, b"hello");
1907-
}
1908-
1909-
/// Verify that the stdin writer delivers all buffered data before exiting
1910-
/// when the sender is dropped. This ensures channel_eof doesn't cause
1911-
/// data loss — only signals "no more data after this".
1912-
#[test]
1913-
fn stdin_writer_delivers_buffered_data_before_eof() {
1914-
let (sender, receiver) = mpsc::channel::<Vec<u8>>();
1915-
1916-
let mut child = Command::new("wc")
1917-
.arg("-c")
1918-
.stdin(Stdio::piped())
1919-
.stdout(Stdio::piped())
1920-
.spawn()
1921-
.expect("failed to spawn wc");
1922-
1923-
let child_stdin = child.stdin.take().expect("stdin must be piped");
1924-
1925-
std::thread::spawn(move || {
1926-
let mut stdin = child_stdin;
1927-
while let Ok(bytes) = receiver.recv() {
1928-
if stdin.write_all(&bytes).is_err() {
1929-
break;
1930-
}
1931-
let _ = stdin.flush();
1932-
}
1933-
});
1934-
1935-
// Send multiple chunks, then drop the sender.
1936-
for _ in 0..100 {
1937-
sender.send(vec![0u8; 1024]).unwrap();
1938-
}
1939-
drop(sender);
1940-
1941-
let (done_tx, done_rx) = mpsc::channel();
1942-
std::thread::spawn(move || {
1943-
let _ = done_tx.send(child.wait_with_output());
1944-
});
1945-
let output = done_rx
1946-
.recv_timeout(Duration::from_secs(5))
1947-
.expect("wc hung for 5s — stdin was not closed")
1948-
.expect("failed to wait for wc");
1949-
1950-
let count: usize = String::from_utf8_lossy(&output.stdout)
1951-
.trim()
1952-
.parse()
1953-
.expect("wc output was not a number");
1954-
assert_eq!(
1955-
count,
1956-
100 * 1024,
1957-
"expected all 100 KiB delivered before EOF"
1958-
);
1959-
}
1960-
19611912
// -----------------------------------------------------------------------
19621913
// SEC-007: is_loopback_host tests
19631914
// -----------------------------------------------------------------------
@@ -2007,55 +1958,4 @@ mod tests {
20071958
assert!(!is_loopback_host("not-an-ip"));
20081959
assert!(!is_loopback_host("[]"));
20091960
}
2010-
2011-
#[test]
2012-
fn channel_state_independent_input_senders() {
2013-
// Verify that each channel gets its own input sender so that
2014-
// data() and channel_eof() affect only the targeted channel.
2015-
let (tx_a, rx_a) = mpsc::channel::<Vec<u8>>();
2016-
let (tx_b, rx_b) = mpsc::channel::<Vec<u8>>();
2017-
2018-
let mut state_a = ChannelState {
2019-
input_sender: Some(InputSender::Process(tx_a)),
2020-
..Default::default()
2021-
};
2022-
let state_b = ChannelState {
2023-
input_sender: Some(InputSender::Process(tx_b)),
2024-
..Default::default()
2025-
};
2026-
2027-
// Send data to channel A only.
2028-
state_a
2029-
.input_sender
2030-
.as_ref()
2031-
.unwrap()
2032-
.send(b"hello-a".to_vec())
2033-
.unwrap();
2034-
// Send data to channel B only.
2035-
state_b
2036-
.input_sender
2037-
.as_ref()
2038-
.unwrap()
2039-
.send(b"hello-b".to_vec())
2040-
.unwrap();
2041-
2042-
assert_eq!(rx_a.recv().unwrap(), b"hello-a");
2043-
assert_eq!(rx_b.recv().unwrap(), b"hello-b");
2044-
2045-
// EOF on channel A (drop sender) should not affect channel B.
2046-
state_a.input_sender.take();
2047-
assert!(
2048-
rx_a.recv().is_err(),
2049-
"channel A sender dropped, recv should fail"
2050-
);
2051-
2052-
// Channel B should still be functional.
2053-
state_b
2054-
.input_sender
2055-
.as_ref()
2056-
.unwrap()
2057-
.send(b"still-alive".to_vec())
2058-
.unwrap();
2059-
assert_eq!(rx_b.recv().unwrap(), b"still-alive");
2060-
}
20611961
}

0 commit comments

Comments
 (0)