Skip to content
Merged
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
51 changes: 34 additions & 17 deletions crates/openshell-cli/src/ssh.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ use std::os::unix::process::CommandExt;
use std::path::{Path, PathBuf};
use std::process::{Command, ExitStatus, Stdio};
use std::time::{Duration, Instant};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::io::AsyncWriteExt;
use tokio::net::TcpStream;
use tokio::process::{Child, Command as TokioCommand};
use tokio_stream::wrappers::ReceiverStream;
Expand Down Expand Up @@ -1926,20 +1926,37 @@ pub async fn sandbox_ssh_proxy(
.into_diagnostic()?
.into_inner();

let stdin = tokio::io::stdin();
let stdout = tokio::io::stdout();

let to_remote = tokio::spawn(async move {
let mut stdin = stdin;
let mut buf = vec![0u8; 64 * 1024];
while let Ok(n) = stdin.read(&mut buf).await {
if n == 0 {
break;
// Tokio stdin uses an uncancellable read on the runtime's blocking pool.
// If the relay closes while SSH still holds the pipe open, runtime shutdown
// would wait forever for that read. A dedicated thread can be left behind
// when this ProxyCommand process exits without holding the runtime alive.
let (stdin_tx, mut stdin_rx) = tokio::sync::mpsc::channel(8);
std::thread::Builder::new()
.name("ssh-proxy-stdin".into())
.spawn(move || {
use std::io::Read as _;
let mut stdin = std::io::stdin().lock();
let mut buf = vec![0u8; 64 * 1024];
loop {
match stdin.read(&mut buf) {
Ok(0) | Err(_) => break,
Ok(n) => {
if stdin_tx.blocking_send(buf[..n].to_vec()).is_err() {
break;
}
}
}
}
})
.into_diagnostic()?;
let to_remote = tokio::spawn(async move {
while let Some(data) = stdin_rx.recv().await {
if tx
.send(TcpForwardFrame {
payload: Some(openshell_core::proto::tcp_forward_frame::Payload::Data(
buf[..n].to_vec(),
data,
)),
})
.await
Expand All @@ -1952,8 +1969,10 @@ pub async fn sandbox_ssh_proxy(
let from_remote = tokio::spawn(async move {
let mut stdout = stdout;
loop {
let Ok(Some(frame)) = response.message().await else {
break;
let frame = match response.message().await {
Ok(Some(frame)) => frame,
Ok(None) => return Ok::<_, Report>(()),
Err(error) => return Err(error).into_diagnostic(),
};
let Some(openshell_core::proto::tcp_forward_frame::Payload::Data(data)) = frame.payload
else {
Expand All @@ -1962,16 +1981,14 @@ pub async fn sandbox_ssh_proxy(
if data.is_empty() {
continue;
}
if stdout.write_all(&data).await.is_err() {
break;
}
let _ = stdout.flush().await;
stdout.write_all(&data).await.into_diagnostic()?;
stdout.flush().await.into_diagnostic()?;
}
});
let _ = from_remote.await;
let result = from_remote.await;
to_remote.abort();

Ok(())
result.into_diagnostic()?
}

fn grpc_server_from_ssh_gateway_url(gateway_url: &str) -> Result<String> {
Expand Down
117 changes: 117 additions & 0 deletions crates/openshell-cli/tests/ssh_proxy_shutdown_integration.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,117 @@
// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
// SPDX-License-Identifier: Apache-2.0

use bytes::Bytes;
use http_body_util::{BodyExt, StreamBody};
use hyper::body::Frame;
use hyper::service::service_fn;
use hyper_util::rt::{TokioExecutor, TokioIo};
use openshell_core::proto::{TcpForwardFrame, tcp_forward_frame};
use prost::Message;
use std::convert::Infallible;
use std::process::Stdio;
use std::time::Duration;
use tokio::io::AsyncWriteExt;

async fn proxy_exits_with_stdin_open(relay_status: tonic::Status) {
let expected_code = relay_status.code();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (socket, _) = listener.accept().await.unwrap();
let service = service_fn(move |request: hyper::Request<hyper::body::Incoming>| {
let status = relay_status.clone();
async move {
assert_eq!(request.uri().path(), "/openshell.v1.OpenShell/ForwardTcp");
let body = StreamBody::new(futures::stream::once(async move {
// Wait until stdin actually traverses the proxy. After this
// chunk its reader blocks again, with the parent pipe open.
let mut inbound = request.into_body();
let mut pending = Vec::new();
'input: while let Some(frame) = inbound.frame().await {
if let Ok(data) = frame.unwrap().into_data() {
pending.extend_from_slice(&data);
while pending.len() >= 5 {
let length =
u32::from_be_bytes(pending[1..5].try_into().unwrap()) as usize;
if pending.len() < 5 + length {
break;
}
let frame =
TcpForwardFrame::decode(&pending[5..5 + length]).unwrap();
pending.drain(..5 + length);
if matches!(frame.payload, Some(tcp_forward_frame::Payload::Data(data)) if data == b"probe")
{
break 'input;
}
}
}
}
tokio::time::sleep(Duration::from_millis(100)).await;
let trailers = status.into_http::<()>().into_parts().0.headers;
Ok::<Frame<Bytes>, Infallible>(Frame::trailers(trailers))
}));
let mut response = hyper::Response::new(body);
response
.headers_mut()
.insert("content-type", "application/grpc".parse().unwrap());
Ok::<_, Infallible>(response)
}
});
let _ = hyper::server::conn::http2::Builder::new(TokioExecutor::new())
.serve_connection(TokioIo::new(socket), service)
.await;
});
let config = tempfile::tempdir().unwrap();
let mut child = tokio::process::Command::new(env!("CARGO_BIN_EXE_openshell"))
.args([
"ssh-proxy",
"--gateway",
&format!("http://{address}/proxy/connect"),
"--sandbox",
"repro",
"--token",
"test-token",
])
.env("XDG_CONFIG_HOME", config.path())
.env("OPENSHELL_TELEMETRY_ENABLED", "false")
.stdin(Stdio::piped())
.stdout(Stdio::null())
.stderr(Stdio::piped())
.kill_on_drop(true)
.spawn()
.unwrap();
child
.stdin
.as_mut()
.unwrap()
.write_all(b"probe")
.await
.unwrap();
// Keep child.stdin alive throughout wait: closing it would conceal the bug.
let status = tokio::time::timeout(Duration::from_secs(5), child.wait())
.await
.expect("SSH proxy must exit even while the parent holds stdin open")
.unwrap();
let mut stderr = String::new();
tokio::io::AsyncReadExt::read_to_string(child.stderr.as_mut().unwrap(), &mut stderr)
.await
.unwrap();
if expected_code == tonic::Code::Ok {
assert!(status.success(), "{stderr}");
} else {
assert!(!status.success(), "relay failure must propagate to SSH");
assert!(stderr.contains("Deadline expired"), "{stderr}");
}
server.abort();
}

#[tokio::test]
async fn proxy_exits_after_relay_error_with_stdin_open() {
proxy_exits_with_stdin_open(tonic::Status::deadline_exceeded("relay open timed out")).await;
}

#[tokio::test]
async fn proxy_exits_after_clean_relay_close_with_stdin_open() {
proxy_exits_with_stdin_open(tonic::Status::ok("")).await;
}
11 changes: 11 additions & 0 deletions crates/openshell-driver-docker/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5376,6 +5376,7 @@ async fn wait_for_docker_supervisor_ready(
_ if state.running == Some(false) => {
let log_tail = docker_container_log_tail(docker, supervisor_id).await;
let sandbox_log_tail = docker_container_log_tail(docker, sandbox_id).await;
warn!(sandbox_id, supervisor_id, supervisor_logs = %log_tail, sandbox_logs = %sandbox_log_tail, "Docker supervisor exited before becoming ready");
return Err(Status::unavailable(format!(
"Docker supervisor exited before becoming ready{}{}",
format_log_tail(&log_tail),
Expand All @@ -5392,8 +5393,18 @@ fn format_log_tail(log_tail: &str) -> String {
}

fn format_named_log_tail(label: &str, log_tail: &str) -> String {
// gRPC status messages travel in HTTP/2 headers. Two 16 KiB container
// tails exceed the client's 16 KiB header budget and hide the real error
// behind PROTOCOL_ERROR. Allow for up to 3x percent-encoding expansion.
const MAX_STATUS_LOG_TAIL_BYTES: usize = 1024;
if log_tail.is_empty() {
String::new()
} else if log_tail.len() > MAX_STATUS_LOG_TAIL_BYTES {
let mut start = log_tail.len() - MAX_STATUS_LOG_TAIL_BYTES;
while !log_tail.is_char_boundary(start) {
start += 1;
}
format!("; {label}: [truncated] {}", &log_tail[start..])
} else {
format!("; {label}: {log_tail}")
}
Expand Down
26 changes: 26 additions & 0 deletions crates/openshell-driver-docker/src/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,32 @@ use std::io::Read as _;
use std::sync::Arc;
use tempfile::TempDir;

#[test]
fn startup_error_log_tails_fit_grpc_header_budget() {
// Multibyte text exercises both the UTF-8 cut and worst-case gRPC message
// percent encoding. Preserve the final diagnostic from each container.
let logs = format!("{}\nstartup timed out", "🦀".repeat(8192));
let message = format!(
"Docker supervisor exited before becoming ready{}{}",
format_log_tail(&logs),
format_named_log_tail("sandbox log tail", &logs),
);
assert_eq!(message.matches("[truncated]").count(), 2);
assert_eq!(message.matches("startup timed out").count(), 2);
let response = Status::unavailable(message).into_http::<()>();
let header_bytes: usize = response
.headers()
.iter()
.map(|(name, value)| name.as_str().len() + value.as_bytes().len() + 32)
.sum();
assert!(
header_bytes < 16 * 1024,
"status headers: {header_bytes} bytes"
);
assert_eq!(format_log_tail("small error"), "; log tail: small error");
assert!(format_log_tail("").is_empty());
}

fn test_launch_authentication() -> Vec<u8> {
serde_json::to_vec(&SandboxLaunchAuthentication {
supervisor: SupervisorAuthBundle {
Expand Down
1 change: 1 addition & 0 deletions crates/openshell-server/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,7 @@ tonic-prost-build = { workspace = true }
protoc-bin-vendored = { workspace = true }

[dev-dependencies]
tokio = { workspace = true, features = ["test-util"] }
# Tests import the example profiles from providers/ the way an operator
# would; the feature is test-only and never reaches a release binary.
openshell-providers = { path = "../openshell-providers", features = ["example-profiles"] }
Expand Down
Loading
Loading