Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
266 changes: 83 additions & 183 deletions crates/openshell-supervisor-process/src/ssh.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand Down Expand Up @@ -332,6 +337,7 @@ async fn handle_connection(
#[derive(Default)]
struct ChannelState {
input_sender: Option<InputSender>,
input_task: Option<input::InputTask>,
process: Option<Arc<dyn openshell_isolation_interface::contract::BoundaryProcess>>,
terminal: Option<Arc<dyn openshell_isolation_interface::contract::BoundaryTerminal>>,
pty_request: Option<PtyRequest>,
Expand All @@ -344,14 +350,17 @@ struct ChannelState {
}

enum InputSender {
Process(mpsc::Sender<Vec<u8>>),
Process(input::InputSender),
Main(tokio::sync::mpsc::Sender<Vec<u8>>),
}

impl InputSender {
fn send(&self, data: Vec<u8>) -> 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(_) => {
Expand Down Expand Up @@ -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()
{
Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand All @@ -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::<Vec<u8>>();
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();
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -1281,6 +1326,19 @@ mod tests {

async fn main_test_client(
main_session: Option<Arc<MainSession>>,
) -> russh::client::Handle<AcceptAnyServerKey> {
test_client(
main_session,
Arc::new(RejectingExec),
russh::client::Config::default(),
)
.await
}

pub(super) async fn test_client(
main_session: Option<Arc<MainSession>>,
boundary_exec: Arc<dyn openshell_isolation_interface::contract::BoundaryExec>,
client_config: russh::client::Config,
) -> russh::client::Handle<AcceptAnyServerKey> {
let host_key = {
let mut rng = rand::rng();
Expand All @@ -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) =
Expand All @@ -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,
)
Expand Down Expand Up @@ -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::<Vec<u8>>();

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::<Vec<u8>>();

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
// -----------------------------------------------------------------------
Expand Down Expand Up @@ -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::<Vec<u8>>();
let (tx_b, rx_b) = mpsc::channel::<Vec<u8>>();

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");
}
}
Loading
Loading