diff --git a/crates/openshell-cli/src/main.rs b/crates/openshell-cli/src/main.rs index f656fd8833..01f6b8a564 100644 --- a/crates/openshell-cli/src/main.rs +++ b/crates/openshell-cli/src/main.rs @@ -1702,6 +1702,13 @@ enum SandboxCommands { #[arg(long, overrides_with = "tty")] no_tty: bool, + /// Stream stdin as it arrives without allocating a pseudo-terminal. + /// Starts the command before stdin closes. Input is limited to 4 MiB + /// per command. Exceeding the limit cancels execution; input may already + /// have been processed. + #[arg(long, conflicts_with = "tty")] + stream_stdin: bool, + /// Run the command without sourcing shell login/profile startup files. /// /// Default sources them so tool-specific env (`VIRTUAL_ENV`, etc.) is @@ -3574,6 +3581,7 @@ async fn run_async() -> Result<()> { timeout, tty, no_tty, + stream_stdin, envs, command, no_login_shell, @@ -3595,6 +3603,7 @@ async fn run_async() -> Result<()> { workdir.as_deref(), timeout, tty_override, + stream_stdin, &env_map, no_login_shell, &tls, @@ -4205,6 +4214,55 @@ mod tests { use std::ffi::OsString; use std::fs; + #[test] + fn sandbox_exec_stream_stdin_is_explicit() { + for flags in [ + vec![], + vec!["--stream-stdin"], + vec!["--stream-stdin", "--no-tty"], + vec!["--no-tty", "--stream-stdin"], + vec!["--stream-stdin", "--tty", "--no-tty"], + vec!["--tty", "--no-tty", "--stream-stdin"], + ] { + let mut args = vec!["openshell", "sandbox", "exec", "-n", "sandbox-1"]; + args.extend(flags.iter().copied()); + args.extend(["--", "cat"]); + let cli = Cli::try_parse_from(args).expect("exec options should parse"); + let Some(Commands::Sandbox { + command: + Some(SandboxCommands::Exec { + stream_stdin, + tty, + no_tty, + .. + }), + }) = cli.command + else { + panic!("expected sandbox exec"); + }; + assert_eq!(stream_stdin, flags.contains(&"--stream-stdin")); + assert!(!tty); + assert_eq!(no_tty, flags.contains(&"--no-tty")); + } + } + + #[test] + fn sandbox_exec_stream_stdin_conflicts_with_tty() { + for flags in [ + vec!["--stream-stdin", "--tty"], + vec!["--tty", "--stream-stdin"], + vec!["--stream-stdin", "--no-tty", "--tty"], + vec!["--no-tty", "--tty", "--stream-stdin"], + ] { + let mut args = vec!["openshell", "sandbox", "exec", "-n", "sandbox-1"]; + args.extend(flags); + args.extend(["--", "cat"]); + let error = + Cli::try_parse_from(args).expect_err("streaming stdin must not allocate a PTY"); + assert_eq!(error.kind(), clap::error::ErrorKind::ArgumentConflict); + } + } + #[test] fn policy_update_parses_explicit_l7_scope_and_endpoint_path() { let cli = Cli::try_parse_from([ diff --git a/crates/openshell-cli/src/run.rs b/crates/openshell-cli/src/run.rs index ce63bc9f7b..c2db21bdba 100644 --- a/crates/openshell-cli/src/run.rs +++ b/crates/openshell-cli/src/run.rs @@ -1850,6 +1850,9 @@ const MAX_EXEC_STDIN_BYTES: usize = 4 * 1024 * 1024; /// Execute a command in a running sandbox via gRPC, streaming output to the terminal. /// +/// With `stream_stdin`, starts before local EOF and sends up to 4 MiB without a +/// pseudo-terminal. Exceeding that limit cancels execution; the command may have +/// processed partial input. /// Returns the remote command's exit code, or an error if the event stream /// closes before the command reports an exit status. #[allow(clippy::too_many_arguments, clippy::implicit_hasher)] @@ -1860,6 +1863,7 @@ pub async fn sandbox_exec_grpc( workdir: Option<&str>, timeout_seconds: u32, tty_override: Option, + stream_stdin: bool, environment: &HashMap, no_login_shell: bool, tls: &TlsOptions, @@ -1890,15 +1894,16 @@ pub async fn sandbox_exec_grpc( )); } - // Resolve TTY mode: explicit --tty / --no-tty wins, otherwise auto-detect. + // Streaming stdin preserves stdout/stderr separately and never allocates + // a PTY. Other invocations retain explicit overrides and auto-detection. let stdin_is_terminal = std::io::stdin().is_terminal(); - let tty = tty_override.unwrap_or_else(|| stdin_is_terminal && std::io::stdout().is_terminal()); + let tty = !stream_stdin + && tty_override.unwrap_or_else(|| stdin_is_terminal && std::io::stdout().is_terminal()); - // Preserve unary exec for small pipes, including older gateways whose - // interactive RPC closes the SSH channel when stdin reaches EOF. Retain - // the existing 4 MiB input cap because the supervisor's process stdin - // queue is unbounded; larger input should use file upload instead. - let stdin_prefix = if stdin_is_terminal { + // Finite pipes retain atomic oversize rejection and unary exec for small + // input. Streaming starts immediately, enforcing the same total byte cap + // while reading because the supervisor's process stdin queue is unbounded. + let stdin_prefix = if stream_stdin || stdin_is_terminal { Vec::new() } else { tokio::task::spawn_blocking(|| { @@ -1948,7 +1953,8 @@ pub async fn sandbox_exec_grpc( "exec command or environment exceeds the gateway's 1 MiB message limit" )); } - if (tty && stdin_is_terminal) || request.encoded_len() > MAX_EXEC_REQUEST_BYTES { + if stream_stdin || (tty && stdin_is_terminal) || request.encoded_len() > MAX_EXEC_REQUEST_BYTES + { return sandbox_exec_streaming_grpc( client, &sandbox, @@ -1960,6 +1966,7 @@ pub async fn sandbox_exec_grpc( tty, stdin_is_terminal, std::mem::take(&mut request.stdin), + (stream_stdin || !stdin_is_terminal).then_some(MAX_EXEC_STDIN_BYTES), ) .await; } @@ -2334,6 +2341,91 @@ impl Drop for TaskGuard { } } +// Only an explicit local EOF closes this request-body stream. A failed reader +// drops its sender while the RPC remains open until response cancellation. +enum ExecInputMessage { + Frame(Box), + Eof, +} + +fn exec_input_stream( + input_rx: tokio::sync::mpsc::Receiver, +) -> impl futures::Stream + Send { + futures::stream::unfold(input_rx, |mut input_rx| async move { + match input_rx.recv().await { + Some(ExecInputMessage::Frame(frame)) => Some((*frame, input_rx)), + Some(ExecInputMessage::Eof) => None, + None => futures::future::pending().await, + } + }) +} + +// Reading at most the remaining allowance preserves every permitted byte. +// At the limit, one additional byte distinguishes EOF from an oversized input; +// that byte is never forwarded, even if the remote command consumes eagerly. +fn forward_exec_stdin( + mut reader: impl Read, + prefix: &[u8], + limit: Option, + mut send: impl FnMut(&[u8]) -> bool, +) -> std::io::Result<()> { + let limit_error = || { + std::io::Error::other( + "streamed stdin exceeds the 4 MiB limit; the command may have processed partial input; use `sandbox upload` for larger input", + ) + }; + if limit.is_some_and(|limit| prefix.len() > limit) { + return Err(limit_error()); + } + let mut buf = [0u8; 4096]; + for chunk in prefix.chunks(buf.len()) { + if !send(chunk) { + return Ok(()); + } + } + let mut remaining = limit.map(|limit| limit - prefix.len()); + loop { + let read_size = remaining.map_or(buf.len(), |remaining| remaining.clamp(1, buf.len())); + match reader.read(&mut buf[..read_size]) { + Ok(0) => return Ok(()), + Err(error) if error.kind() == ErrorKind::Interrupted => {} + Err(error) => return Err(error), + Ok(n) => { + if let Some(remaining) = &mut remaining { + if n > *remaining { + return Err(limit_error()); + } + *remaining -= n; + } + if !send(&buf[..n]) { + return Ok(()); + } + } + } + } +} + +// Keep the EOF decision beside the reader result so failures cannot queue a +// successful end-of-input marker before the response loop cancels the RPC. +fn write_exec_stdin_frames( + reader: impl Read, + prefix: &[u8], + limit: Option, + sender: &tokio::sync::mpsc::Sender, +) -> std::io::Result<()> { + use openshell_core::proto::{ExecSandboxInput, exec_sandbox_input}; + + forward_exec_stdin(reader, prefix, limit, |chunk| { + sender + .blocking_send(ExecInputMessage::Frame(Box::new(ExecSandboxInput { + payload: Some(exec_sandbox_input::Payload::Stdin(chunk.to_vec())), + }))) + .is_ok() + })?; + let _ = sender.blocking_send(ExecInputMessage::Eof); + Ok(()) +} + #[allow(clippy::too_many_arguments)] async fn sandbox_exec_streaming_grpc( mut client: crate::tls::GrpcClient, @@ -2346,11 +2438,11 @@ async fn sandbox_exec_streaming_grpc( tty: bool, stdin_is_terminal: bool, stdin_prefix: Vec, + stdin_limit: Option, ) -> Result { #[cfg(unix)] use openshell_core::proto::ExecSandboxWindowResize; use openshell_core::proto::{ExecSandboxInput, exec_sandbox_input}; - use tokio_stream::wrappers::ReceiverStream; let (cols, rows) = if tty { local_terminal_size().unwrap_or((80, 24)) @@ -2358,11 +2450,11 @@ async fn sandbox_exec_streaming_grpc( (0, 0) }; - let (input_tx, input_rx) = tokio::sync::mpsc::channel::(64); + let (input_tx, input_rx) = tokio::sync::mpsc::channel::(64); // Send the start message with exec metadata. input_tx - .send(ExecSandboxInput { + .send(ExecInputMessage::Frame(Box::new(ExecSandboxInput { payload: Some(exec_sandbox_input::Payload::Start(ExecSandboxRequest { request_id: String::new(), sandbox: sandbox.object_name().to_string(), @@ -2379,12 +2471,12 @@ async fn sandbox_exec_streaming_grpc( cols, rows, })), - }) + }))) .await .into_diagnostic()?; let mut stream = client - .exec_sandbox_interactive(ReceiverStream::new(input_rx)) + .exec_sandbox_interactive(exec_input_stream(input_rx)) .await .into_diagnostic()? .into_inner(); @@ -2397,46 +2489,18 @@ async fn sandbox_exec_streaming_grpc( None }; - // Stdin reader on a detached OS thread. Using std::thread (not - // spawn_blocking) so the tokio runtime shutdown doesn't wait for a - // thread blocked on stdin.read(). The thread exits when the channel - // closes (blocking_send returns Err) or stdin hits EOF. + // A detached OS thread keeps an idle stdin read from blocking Tokio runtime + // shutdown. It can outlive this operation until input/EOF arrives, but never + // keeps the CLI process alive after the response completes or fails. let stdin_tx = input_tx.clone(); let (stdin_result_tx, mut stdin_result_rx) = tokio::sync::oneshot::channel(); std::thread::spawn(move || { - let mut stdin = std::io::stdin().lock(); - let mut buf = [0u8; 4096]; - let result = (|| { - for chunk in stdin_prefix.chunks(buf.len()) { - if stdin_tx - .blocking_send(ExecSandboxInput { - payload: Some(exec_sandbox_input::Payload::Stdin(chunk.to_vec())), - }) - .is_err() - { - return Ok(()); - } - } - loop { - match stdin.read(&mut buf) { - Ok(0) => return Ok(()), - Err(error) if error.kind() == ErrorKind::Interrupted => {} - Err(error) => return Err(error), - Ok(n) => { - if stdin_tx - .blocking_send(ExecSandboxInput { - payload: Some(exec_sandbox_input::Payload::Stdin( - buf[..n].to_vec(), - )), - }) - .is_err() - { - return Ok(()); - } - } - } - } - })(); + let result = write_exec_stdin_frames( + std::io::stdin().lock(), + &stdin_prefix, + stdin_limit, + &stdin_tx, + ); let _ = stdin_result_tx.send(result); }); @@ -2455,7 +2519,11 @@ async fn sandbox_exec_streaming_grpc( ExecSandboxWindowResize { cols, rows }, )), }; - if resize_tx.send(msg).await.is_err() { + if resize_tx + .send(ExecInputMessage::Frame(Box::new(msg))) + .await + .is_err() + { break; } } @@ -2467,10 +2535,8 @@ async fn sandbox_exec_streaming_grpc( #[cfg(unix)] let _resize_guard = resize_task.map(TaskGuard); - // Keep a sender until the reader confirms clean EOF. On a read error, - // cancel the response stream before the gateway can treat channel EOF as - // successful completion of a partial command. - let mut pipe_input_tx = Some(input_tx); + // Retain a sender to invalidate the request on a read error. The request + // stream sends EOF only after the reader explicitly reports clean EOF. let mut exit_code = 0i32; let mut exit_seen = false; @@ -2483,24 +2549,13 @@ async fn sandbox_exec_streaming_grpc( result = &mut stdin_result_rx, if !stdin_reader_done => { stdin_reader_done = true; match result.into_diagnostic()? { - Ok(()) => { - let sender = pipe_input_tx.take().expect("stdin sender is held until EOF"); - drop(sender); - } + Ok(()) => {} Err(error) => { - let sender = pipe_input_tx.take().expect("stdin sender is held until EOF"); - // A clean request EOF would make the gateway execute - // the truncated input. An invalid frame makes the - // gateway abort the command instead. - let abort = ExecSandboxInput { payload: None }; - if tokio::time::timeout(Duration::from_secs(5), sender.send(abort)) - .await - .is_err() - { - // Keep the request body open if a blocked remote - // stdin prevents delivery of the abort frame. - std::mem::forget(sender); - } + // An invalid frame aborts the command if it reaches the + // gateway. If delivery is blocked, closing the response + // cancels the relay independently of request termination. + let abort = ExecInputMessage::Frame(Box::new(ExecSandboxInput { payload: None })); + let _ = tokio::time::timeout(Duration::from_secs(5), input_tx.send(abort)).await; drop(stream); return Err(error).into_diagnostic(); } @@ -2525,7 +2580,8 @@ async fn sandbox_exec_streaming_grpc( Some(exec_sandbox_event::Payload::Exit(exit)) => { exit_code = exit.exit_code; exit_seen = true; - break; + // A terminal event does not guarantee successful gRPC trailers. + // Keep draining so a relay failure cannot become a successful exit. } None => {} } @@ -6490,6 +6546,155 @@ mod tests { service_url_for_gateway, workspace_member_to_json, }; + #[test] + fn exec_stdin_limit_counts_prefix_and_never_forwards_the_extra_byte() { + let prefix = vec![b'p'; 4095]; + let mut forwarded = Vec::new(); + let error = super::forward_exec_stdin(&b"xy"[..], &prefix, Some(4096), |chunk| { + assert!(chunk.len() <= 4096); + forwarded.extend_from_slice(chunk); + true + }) + .expect_err("one byte above the limit must fail"); + assert_eq!(forwarded.len(), 4096); + assert_eq!(forwarded.last(), Some(&b'x')); + assert!( + error + .to_string() + .contains("may have processed partial input") + ); + } + + #[test] + fn exec_stdin_accepts_exact_limit_and_preserves_chunk_order() { + let prefix = vec![b'p'; 4097]; + let input = vec![b'i'; 4097]; + let mut forwarded = Vec::new(); + super::forward_exec_stdin(input.as_slice(), &prefix, Some(8194), |chunk| { + assert!(chunk.len() <= 4096); + forwarded.extend_from_slice(chunk); + true + }) + .expect("exact limit followed by EOF must succeed"); + assert_eq!(forwarded, [prefix, input].concat()); + } + + #[test] + fn exec_stdin_propagates_read_failure_after_partial_input() { + struct FailedReader(R); + impl std::io::Read for FailedReader { + fn read(&mut self, buffer: &mut [u8]) -> std::io::Result { + match self.0.read(buffer)? { + 0 => Err(std::io::Error::other("synthetic read failure")), + size => Ok(size), + } + } + } + let mut forwarded = Vec::new(); + let error = + super::forward_exec_stdin(FailedReader(&b"request"[..]), &[], Some(4096), |chunk| { + forwarded.extend_from_slice(chunk); + true + }) + .expect_err("reader failures must not become EOF"); + assert_eq!(forwarded, b"request"); + assert_eq!(error.to_string(), "synthetic read failure"); + } + + #[tokio::test] + async fn exec_input_requires_explicit_clean_eof() { + use futures::StreamExt; + let (sender, receiver) = tokio::sync::mpsc::channel(1); + let stream = super::exec_input_stream(receiver); + tokio::pin!(stream); + sender.send(super::ExecInputMessage::Eof).await.unwrap(); + assert!(stream.next().await.is_none()); + + let (sender, receiver) = tokio::sync::mpsc::channel(1); + let stream = super::exec_input_stream(receiver); + tokio::pin!(stream); + drop(sender); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(10), stream.next()) + .await + .is_err() + ); + } + + fn assert_exec_stdin_writer_error(reader: impl std::io::Read, prefix: &[u8]) -> std::io::Error { + use futures::{FutureExt, StreamExt}; + + let (sender, receiver) = tokio::sync::mpsc::channel(2); + let error = super::write_exec_stdin_frames(reader, prefix, Some(2), &sender) + .expect_err("a failed reader must not queue EOF"); + drop(sender); + let mut stream = Box::pin(super::exec_input_stream(receiver)); + let frame = stream + .next() + .now_or_never() + .expect("the permitted bytes are already queued") + .expect("the request body must contain its input frame"); + assert_eq!( + frame.payload, + Some(openshell_core::proto::exec_sandbox_input::Payload::Stdin( + b"ab".to_vec() + )), + ); + assert!( + stream.next().now_or_never().is_none(), + "a failed reader must leave the request pending, not queue an explicit EOF", + ); + error + } + + #[test] + fn exec_stdin_writer_overflow_does_not_queue_eof() { + let error = assert_exec_stdin_writer_error(&b"abc"[..], &[]); + assert!(error.to_string().contains("partial input")); + } + + #[test] + fn exec_stdin_writer_read_error_does_not_queue_eof() { + struct FailedReader; + impl std::io::Read for FailedReader { + fn read(&mut self, _: &mut [u8]) -> std::io::Result { + Err(std::io::Error::other("synthetic read failure")) + } + } + + let error = assert_exec_stdin_writer_error(FailedReader, b"ab"); + assert_eq!(error.to_string(), "synthetic read failure"); + } + + #[test] + fn exec_stdin_writer_exact_limit_queues_eof() { + use futures::{FutureExt, StreamExt}; + + let (sender, receiver) = tokio::sync::mpsc::channel(2); + super::write_exec_stdin_frames(&b"ab"[..], &[], Some(2), &sender) + .expect("the exact limit followed by EOF is valid"); + drop(sender); + let mut stream = Box::pin(super::exec_input_stream(receiver)); + let frame = stream + .next() + .now_or_never() + .expect("the input is already queued") + .expect("the request body must contain its input frame"); + assert_eq!( + frame.payload, + Some(openshell_core::proto::exec_sandbox_input::Payload::Stdin( + b"ab".to_vec() + )), + ); + assert!( + stream + .next() + .now_or_never() + .expect("EOF is already queued") + .is_none(), + ); + } + #[test] fn zero_exec_timeout_is_omitted() { assert!(proto_execution_timeout(0).unwrap().is_none()); diff --git a/crates/openshell-cli/tests/sandbox_exec_streaming_integration.rs b/crates/openshell-cli/tests/sandbox_exec_streaming_integration.rs new file mode 100644 index 0000000000..8d60915723 --- /dev/null +++ b/crates/openshell-cli/tests/sandbox_exec_streaming_integration.rs @@ -0,0 +1,789 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +#![cfg(unix)] + +mod helpers; + +use std::path::Path; +use std::process::{Output, Stdio}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use helpers::{build_ca, build_client_cert, build_server_cert}; +use openshell_core::proto::open_shell_server::{OpenShell, OpenShellServer}; +use openshell_core::proto::{self, exec_sandbox_event, exec_sandbox_input}; +use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; +use tokio::net::TcpListener; +use tokio::process::{Child, Command}; +use tokio::sync::{Notify, mpsc}; +use tokio::task::JoinHandle; +use tokio::time::timeout; +use tokio_stream::wrappers::{ReceiverStream, TcpListenerStream}; +use tonic::transport::{Certificate, Identity, Server, ServerTlsConfig}; +use tonic::{Request, Response, Status}; + +const DEADLINE: Duration = Duration::from_secs(10); +const STDIN_LIMIT: usize = 4 * 1024 * 1024; +type EventStream = ReceiverStream>; + +#[derive(Clone, Copy)] +enum Scenario { + Echo, + Count, + AwaitCancellation, + ReadAfterCancellation, + EarlyExit, + ErrorAfterExit, + MissingExit, + Disconnect, +} + +#[derive(Default)] +struct Calls { + lookups: usize, + unary: Vec, + starts: Vec, + input_bytes: usize, + request_ended: bool, + request_error: bool, + response_cancelled: bool, +} + +#[derive(Clone)] +struct MockGateway { + scenario: Scenario, + calls: Arc>, + finished: Arc, +} + +fn stdout(data: impl Into>) -> proto::ExecSandboxEvent { + proto::ExecSandboxEvent { + payload: Some(exec_sandbox_event::Payload::Stdout( + proto::ExecSandboxStdout { data: data.into() }, + )), + } +} + +fn exit(code: i32) -> proto::ExecSandboxEvent { + proto::ExecSandboxEvent { + payload: Some(exec_sandbox_event::Payload::Exit(proto::ExecSandboxExit { + exit_code: code, + })), + } +} + +impl MockGateway { + async fn exchange( + self, + mut input: tonic::Streaming, + output: mpsc::Sender>, + ) { + let Some(exec_sandbox_input::Payload::Start(start)) = input + .message() + .await + .expect("read start frame") + .expect("start frame") + .payload + else { + panic!("first frame must start the command"); + }; + assert!(!start.tty, "streaming pipes must not allocate a TTY"); + assert!(start.stdin.is_empty(), "stdin must follow the start frame"); + self.calls.lock().unwrap().starts.push(start); + + match self.scenario { + Scenario::EarlyExit => { + let _ = output.send(Ok(exit(0))).await; + return; + } + Scenario::ErrorAfterExit => { + let _ = output.send(Ok(exit(0))).await; + // Deliver Exit before the later failing trailer so a client that + // stops at Exit incorrectly reports success instead of failure. + tokio::time::sleep(Duration::from_millis(100)).await; + let _ = output + .send(Err(Status::internal("failure after exit"))) + .await; + return; + } + Scenario::MissingExit => { + let _ = output.send(Ok(stdout(b"incomplete\n".to_vec()))).await; + return; + } + Scenario::Disconnect => { + let _ = output + .send(Err(Status::unavailable("relay disconnected"))) + .await; + return; + } + Scenario::Echo | Scenario::Count | Scenario::AwaitCancellation => {} + Scenario::ReadAfterCancellation => { + // Delay request polling until the client has cancelled its + // response. Tonic can then expose CANCEL as request EOF, even + // though the CLI never queued its explicit stdin EOF marker. + output.closed().await; + self.calls.lock().unwrap().response_cancelled = true; + } + } + + loop { + let message = tokio::select! { + biased; + () = output.closed(), if !matches!(self.scenario, Scenario::ReadAfterCancellation) => { + self.calls.lock().unwrap().response_cancelled = true; + return; + } + message = input.message() => message, + }; + match message { + Ok(Some(frame)) => match frame.payload { + Some(exec_sandbox_input::Payload::Stdin(bytes)) => { + self.calls.lock().unwrap().input_bytes += bytes.len(); + if matches!(self.scenario, Scenario::Echo) + && output.send(Ok(stdout(bytes))).await.is_err() + { + return; + } + } + // The malformed abort frame can arrive before cancellation + // or be discarded by it. Neither outcome proves clean EOF. + None if matches!(self.scenario, Scenario::ReadAfterCancellation) => {} + None if matches!(self.scenario, Scenario::AwaitCancellation) => break, + None => return, + unexpected => panic!("unexpected input after start: {unexpected:?}"), + }, + Ok(None) => { + // Request completion alone cannot distinguish clean stdin + // EOF from HTTP/2 cancellation. Observe the response too. + self.calls.lock().unwrap().request_ended = true; + break; + } + Err(_) => { + self.calls.lock().unwrap().request_error = true; + if matches!(self.scenario, Scenario::AwaitCancellation) { + break; + } + return; + } + } + } + + match self.scenario { + Scenario::Echo => { + let _ = output.send(Ok(stdout(b"after-eof\n".to_vec()))).await; + let _ = output + .send(Ok(proto::ExecSandboxEvent { + payload: Some(exec_sandbox_event::Payload::Stderr( + proto::ExecSandboxStderr { + data: b"remote-stderr\n".to_vec(), + }, + )), + })) + .await; + let _ = output.send(Ok(exit(7))).await; + } + Scenario::Count => { + let total = self.calls.lock().unwrap().input_bytes; + let _ = output + .send(Ok(stdout(format!("{total}\n").into_bytes()))) + .await; + let _ = output.send(Ok(exit(0))).await; + } + Scenario::AwaitCancellation | Scenario::ReadAfterCancellation => { + // Keep the response open so this witness cannot be caused by + // the mock finishing normally after an ambiguous request end. + output.closed().await; + self.calls.lock().unwrap().response_cancelled = true; + } + _ => unreachable!(), + } + } +} + +// Generate unused trait methods inside the async_trait expansion so this mock +// implements only the RPC behavior under test without hand-written boilerplate. +macro_rules! mock_gateway { + ( + unary { $( $method:ident($request:ty) -> $response:ty; )* } + client_stream { $( $client_method:ident($client_request:ty) -> $client_response:ty; )* } + server_stream { $( $server_method:ident($server_request:ty) -> $stream_type:ident($server_response:ty); )* } + bidi { $( $bidi_method:ident($bidi_request:ty) -> $bidi_type:ident($bidi_response:ty); )* } + ) => { + #[tonic::async_trait] + impl OpenShell for MockGateway { + $(async fn $method(&self, _: Request<$request>) -> Result, Status> { + Err(Status::unimplemented("unused test RPC")) + })* + $(async fn $client_method(&self, _: Request>) -> Result, Status> { + Err(Status::unimplemented("unused test RPC")) + })* + $(type $stream_type = ReceiverStream>; + async fn $server_method(&self, _: Request<$server_request>) -> Result, Status> { + Err(Status::unimplemented("unused test RPC")) + })* + $(type $bidi_type = ReceiverStream>; + async fn $bidi_method(&self, _: Request>) -> Result, Status> { + Err(Status::unimplemented("unused test RPC")) + })* + + async fn get_sandbox(&self, _: Request) -> Result, Status> { + self.calls.lock().unwrap().lookups += 1; + Ok(Response::new(proto::SandboxResponse { + sandbox: Some(proto::Sandbox { + metadata: Some(proto::datamodel::v1::ObjectMeta { + id: "test-id".to_string(), + name: "test-sandbox".to_string(), + workspace: "default".to_string(), + ..Default::default() + }), + status: Some(proto::SandboxStatus { + phase: proto::SandboxPhase::Ready.into(), + ..Default::default() + }), + ..Default::default() + }), + ..Default::default() + })) + } + + type ExecSandboxStream = EventStream; + async fn exec_sandbox(&self, request: Request) -> Result, Status> { + let request = request.into_inner(); + let bytes = request.stdin.clone(); + self.calls.lock().unwrap().unary.push(request); + let (sender, receiver) = mpsc::channel(2); + sender.send(Ok(stdout(bytes))).await.unwrap(); + sender.send(Ok(exit(0))).await.unwrap(); + Ok(Response::new(ReceiverStream::new(receiver))) + } + + type ExecSandboxInteractiveStream = EventStream; + async fn exec_sandbox_interactive(&self, request: Request>) -> Result, Status> { + let (sender, receiver) = mpsc::channel(1); + let service = self.clone(); + tokio::spawn(async move { + service.clone().exchange(request.into_inner(), sender).await; + service.finished.notify_one(); + }); + Ok(Response::new(ReceiverStream::new(receiver))) + } + } + }; +} + +mock_gateway! { + unary { + health(proto::HealthRequest) -> proto::HealthResponse; + get_current_user(proto::GetCurrentUserRequest) -> proto::GetCurrentUserResponse; + get_gateway_info(proto::GetGatewayInfoRequest) -> proto::GetGatewayInfoResponse; + create_sandbox(proto::CreateSandboxRequest) -> proto::SandboxResponse; + begin_rootfs_tar_staging(proto::BeginRootfsTarStagingRequest) -> proto::BeginRootfsTarStagingResponse; + list_sandboxes(proto::ListSandboxesRequest) -> proto::ListSandboxesResponse; + create_sandbox_template(proto::CreateSandboxTemplateRequest) -> proto::SandboxTemplateResponse; + get_sandbox_template(proto::GetSandboxTemplateRequest) -> proto::SandboxTemplateResponse; + list_sandbox_templates(proto::ListSandboxTemplatesRequest) -> proto::ListSandboxTemplatesResponse; + delete_sandbox_template(proto::DeleteSandboxTemplateRequest) -> proto::DeleteSandboxTemplateResponse; + list_sandbox_providers(proto::ListSandboxProvidersRequest) -> proto::ListSandboxProvidersResponse; + attach_sandbox_provider(proto::AttachSandboxProviderRequest) -> proto::AttachSandboxProviderResponse; + detach_sandbox_provider(proto::DetachSandboxProviderRequest) -> proto::DetachSandboxProviderResponse; + get_sandbox_provider_status(proto::GetSandboxProviderStatusRequest) -> proto::GetSandboxProviderStatusResponse; + delete_sandbox(proto::DeleteSandboxRequest) -> proto::DeleteSandboxResponse; + stop_sandbox(proto::StopSandboxRequest) -> proto::SandboxResponse; + start_sandbox(proto::StartSandboxRequest) -> proto::SandboxResponse; + create_ssh_session(proto::CreateSshSessionRequest) -> proto::CreateSshSessionResponse; + expose_service(proto::ExposeServiceRequest) -> proto::ServiceEndpointResponse; + get_service(proto::GetServiceRequest) -> proto::ServiceEndpointResponse; + list_services(proto::ListServicesRequest) -> proto::ListServicesResponse; + delete_service(proto::DeleteServiceRequest) -> proto::DeleteServiceResponse; + revoke_ssh_session(proto::RevokeSshSessionRequest) -> proto::RevokeSshSessionResponse; + create_provider(proto::CreateProviderRequest) -> proto::ProviderResponse; + get_provider(proto::GetProviderRequest) -> proto::ProviderResponse; + list_providers(proto::ListProvidersRequest) -> proto::ListProvidersResponse; + list_provider_profiles(proto::ListProviderProfilesRequest) -> proto::ListProviderProfilesResponse; + get_provider_profile(proto::GetProviderProfileRequest) -> proto::ProviderProfileResponse; + import_provider_profiles(proto::ImportProviderProfilesRequest) -> proto::ImportProviderProfilesResponse; + update_provider_profiles(proto::UpdateProviderProfilesRequest) -> proto::UpdateProviderProfilesResponse; + lint_provider_profiles(proto::LintProviderProfilesRequest) -> proto::LintProviderProfilesResponse; + update_provider(proto::UpdateProviderRequest) -> proto::ProviderResponse; + get_provider_refresh_status(proto::GetProviderRefreshStatusRequest) -> proto::GetProviderRefreshStatusResponse; + configure_provider_refresh(proto::ConfigureProviderRefreshRequest) -> proto::ConfigureProviderRefreshResponse; + rotate_provider_credential(proto::RotateProviderCredentialRequest) -> proto::RotateProviderCredentialResponse; + delete_provider_refresh(proto::DeleteProviderRefreshRequest) -> proto::DeleteProviderRefreshResponse; + delete_provider(proto::DeleteProviderRequest) -> proto::DeleteProviderResponse; + delete_provider_profile(proto::DeleteProviderProfileRequest) -> proto::DeleteProviderProfileResponse; + get_sandbox_config(proto::GetSandboxConfigRequest) -> proto::GetSandboxConfigResponse; + get_gateway_config(proto::GetGatewayConfigRequest) -> proto::GetGatewayConfigResponse; + update_config(proto::UpdateConfigRequest) -> proto::UpdateConfigResponse; + get_sandbox_policy_status(proto::GetSandboxPolicyStatusRequest) -> proto::GetSandboxPolicyStatusResponse; + list_sandbox_policies(proto::ListSandboxPoliciesRequest) -> proto::ListSandboxPoliciesResponse; + report_policy_status(proto::ReportPolicyStatusRequest) -> proto::ReportPolicyStatusResponse; + report_endpoint_status(proto::ReportEndpointStatusRequest) -> proto::ReportEndpointStatusResponse; + report_provider_readiness(proto::ReportProviderReadinessRequest) -> proto::ReportProviderReadinessResponse; + report_sandbox_configuration(proto::ReportSandboxConfigurationRequest) -> proto::ReportSandboxConfigurationResponse; + get_sandbox_provider_environment(proto::GetSandboxProviderEnvironmentRequest) -> proto::GetSandboxProviderEnvironmentResponse; + exchange_provider_subject_token(proto::ExchangeProviderSubjectTokenRequest) -> proto::ExchangeProviderSubjectTokenResponse; + get_sandbox_logs(proto::GetSandboxLogsRequest) -> proto::GetSandboxLogsResponse; + report_main_process_exit(proto::ReportMainProcessExitRequest) -> proto::ReportMainProcessExitResponse; + finalize_main_process_exit(proto::FinalizeMainProcessExitRequest) -> proto::FinalizeMainProcessExitResponse; + peer_report_provider_readiness(proto::ReportProviderReadinessRequest) -> proto::ReportProviderReadinessResponse; + peer_report_endpoint_status(proto::ReportEndpointStatusRequest) -> proto::ReportEndpointStatusResponse; + peer_get_sandbox_provider_status(proto::GetSandboxProviderStatusRequest) -> proto::GetSandboxProviderStatusResponse; + submit_policy_analysis(proto::SubmitPolicyAnalysisRequest) -> proto::SubmitPolicyAnalysisResponse; + get_draft_policy(proto::GetDraftPolicyRequest) -> proto::GetDraftPolicyResponse; + approve_draft_chunk(proto::ApproveDraftChunkRequest) -> proto::ApproveDraftChunkResponse; + reject_draft_chunk(proto::RejectDraftChunkRequest) -> proto::RejectDraftChunkResponse; + approve_all_draft_chunks(proto::ApproveAllDraftChunksRequest) -> proto::ApproveAllDraftChunksResponse; + edit_draft_chunk(proto::EditDraftChunkRequest) -> proto::EditDraftChunkResponse; + undo_draft_chunk(proto::UndoDraftChunkRequest) -> proto::UndoDraftChunkResponse; + clear_draft_chunks(proto::ClearDraftChunksRequest) -> proto::ClearDraftChunksResponse; + get_draft_history(proto::GetDraftHistoryRequest) -> proto::GetDraftHistoryResponse; + issue_sandbox_token(proto::IssueSandboxTokenRequest) -> proto::IssueSandboxTokenResponse; + refresh_sandbox_token(proto::RefreshSandboxTokenRequest) -> proto::RefreshSandboxTokenResponse; + create_workspace(proto::CreateWorkspaceRequest) -> proto::CreateWorkspaceResponse; + get_workspace(proto::GetWorkspaceRequest) -> proto::GetWorkspaceResponse; + list_workspaces(proto::ListWorkspacesRequest) -> proto::ListWorkspacesResponse; + delete_workspace(proto::DeleteWorkspaceRequest) -> proto::DeleteWorkspaceResponse; + add_workspace_member(proto::AddWorkspaceMemberRequest) -> proto::AddWorkspaceMemberResponse; + remove_workspace_member(proto::RemoveWorkspaceMemberRequest) -> proto::RemoveWorkspaceMemberResponse; + list_workspace_members(proto::ListWorkspaceMembersRequest) -> proto::ListWorkspaceMembersResponse; + } + client_stream { + push_sandbox_logs(proto::PushSandboxLogsRequest) -> proto::PushSandboxLogsResponse; + } + server_stream { + watch_sandbox(proto::WatchSandboxRequest) -> WatchSandboxStream(proto::SandboxStreamEvent); + } + bidi { + forward_tcp(proto::TcpForwardFrame) -> ForwardTcpStream(proto::TcpForwardFrame); + connect_supervisor(proto::SupervisorMessage) -> ConnectSupervisorStream(proto::GatewayMessage); + relay_stream(proto::RelayFrame) -> RelayStreamStream(proto::RelayFrame); + peer_relay(proto::PeerRelayFrame) -> PeerRelayStream(proto::PeerRelayFrame); + } +} + +struct TestGateway { + endpoint: String, + config: tempfile::TempDir, + calls: Arc>, + finished: Arc, + task: JoinHandle<()>, +} + +impl Drop for TestGateway { + fn drop(&mut self) { + self.task.abort(); + } +} + +impl TestGateway { + async fn start(scenario: Scenario) -> Self { + let (ca, ca_key) = build_ca(); + let (server_cert, server_key) = build_server_cert(&ca, &ca_key); + let (client_cert, client_key) = build_client_cert(&ca, &ca_key); + let config = tempfile::tempdir().unwrap(); + let certs = config.path().join("openshell/gateways/test-gateway/mtls"); + std::fs::create_dir_all(&certs).unwrap(); + std::fs::write(certs.join("ca.crt"), ca.pem()).unwrap(); + std::fs::write(certs.join("tls.crt"), client_cert).unwrap(); + std::fs::write(certs.join("tls.key"), client_key).unwrap(); + let tls = ServerTlsConfig::new() + .identity(Identity::from_pem(server_cert, server_key)) + .client_ca_root(Certificate::from_pem(ca.pem())); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let endpoint = format!( + "https://localhost:{}", + listener.local_addr().unwrap().port() + ); + let calls = Arc::new(Mutex::new(Calls::default())); + let finished = Arc::new(Notify::new()); + let service = MockGateway { + scenario, + calls: Arc::clone(&calls), + finished: Arc::clone(&finished), + }; + let task = tokio::spawn(async move { + Server::builder() + .tls_config(tls) + .unwrap() + .add_service(OpenShellServer::new(service)) + .serve_with_incoming(TcpListenerStream::new(listener)) + .await + .unwrap(); + }); + Self { + endpoint, + config, + calls, + finished, + task, + } + } + + fn command(&self, executable: &Path, streaming: bool) -> Command { + let mut command = Command::new(executable); + command.args([ + "--gateway", + "test-gateway", + "--gateway-endpoint", + &self.endpoint, + "--workspace", + "default", + "--color", + "never", + "sandbox", + "exec", + "--name", + "test-sandbox", + ]); + if streaming { + command.arg("--stream-stdin"); + } + command + .args(["--no-tty", "--no-login-shell", "--", "test-command"]) + .env("XDG_CONFIG_HOME", self.config.path()) + .env( + "OPENSHELL_SYSTEM_GATEWAY_DIR", + self.config.path().join("system"), + ) + .env_remove("OPENSHELL_GATEWAY") + .env_remove("OPENSHELL_GATEWAY_ENDPOINT") + .env_remove("OPENSHELL_GATEWAY_INSECURE") + .env_remove("RUST_LOG") + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(true); + command + } + + fn spawn(&self, streaming: bool) -> Child { + self.command(Path::new(env!("CARGO_BIN_EXE_openshell")), streaming) + .spawn() + .unwrap() + } + + fn assert_one_stream(&self) { + let calls = self.calls.lock().unwrap(); + assert_eq!(calls.lookups, 1); + assert_eq!(calls.starts.len(), 1, "a command must not be relaunched"); + assert!( + calls.unary.is_empty(), + "streaming mode must never use unary exec" + ); + } + + async fn wait_for_stream_end(&self) { + timeout(DEADLINE, self.finished.notified()) + .await + .expect("gateway did not observe request completion or cancellation"); + } +} + +async fn finish(child: Child) -> Output { + timeout(DEADLINE, child.wait_with_output()) + .await + .expect("CLI did not terminate before the deadline") + .unwrap() +} + +async fn send_input(child: &mut Child, input: &[u8]) { + let mut stdin = child.stdin.take().unwrap(); + // An oversized input or failed relay may close the pipe before all bytes + // are written. The subprocess status and RPC observations decide success. + let _ = timeout(DEADLINE, stdin.write_all(input)) + .await + .expect("stdin write timed out"); +} + +#[tokio::test] +async fn streaming_exchanges_two_requests_before_eof_and_drains_final_output() { + let gateway = TestGateway::start(Scenario::Echo).await; + // This override is restricted to the regression reproducer: an older CLI + // has no --stream-stdin flag, but must still reach the same gateway lookup. + let baseline = std::env::var_os("OPENSHELL_STREAMING_TEST_BASELINE_CLI"); + let executable = baseline + .as_deref() + .map_or_else(|| Path::new(env!("CARGO_BIN_EXE_openshell")), Path::new); + let mut child = gateway + .command(executable, baseline.is_none()) + .spawn() + .unwrap(); + let mut stdin = child.stdin.take().unwrap(); + let mut stdout = BufReader::new(child.stdout.take().unwrap()); + for request in ["first request\n", "second request\n"] { + stdin.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + let read = timeout(DEADLINE, stdout.read_line(&mut response)).await; + if read.is_err() { + let calls = gateway.calls.lock().unwrap(); + assert_eq!(calls.lookups, 1, "CLI did not reach the gateway"); + panic!( + "no response before stdin EOF: {} starts, {} unary calls", + calls.starts.len(), + calls.unary.len() + ); + } + assert_eq!( + gateway.calls.lock().unwrap().lookups, + 1, + "CLI did not reach the gateway" + ); + assert_ne!(read.unwrap().unwrap(), 0, "CLI ended before response"); + assert_eq!(response, request); + assert!( + child.try_wait().unwrap().is_none(), + "CLI ended between exchanges" + ); + } + drop(stdin); + let mut final_stdout = String::new(); + timeout(DEADLINE, stdout.read_to_string(&mut final_stdout)) + .await + .unwrap() + .unwrap(); + let output = finish(child).await; + assert_eq!(output.status.code(), Some(7)); + assert_eq!(final_stdout, "after-eof\n"); + assert_eq!(output.stderr, b"remote-stderr\n"); + gateway.assert_one_stream(); + assert!(gateway.calls.lock().unwrap().request_ended); +} + +#[tokio::test] +async fn streaming_remote_exit_does_not_wait_for_idle_open_stdin() { + let gateway = TestGateway::start(Scenario::EarlyExit).await; + let mut child = gateway.spawn(true); + let held_open = child.stdin.take().unwrap(); + let output = finish(child).await; + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + gateway.assert_one_stream(); + assert!(!gateway.calls.lock().unwrap().request_ended); + drop(held_open); +} + +#[tokio::test] +async fn streaming_trailer_error_after_exit_is_not_success() { + let gateway = TestGateway::start(Scenario::ErrorAfterExit).await; + let mut child = gateway.spawn(true); + let held_open = child.stdin.take().unwrap(); + let output = finish(child).await; + assert!(!output.status.success()); + assert!(String::from_utf8_lossy(&output.stderr).contains("failure after exit")); + gateway.assert_one_stream(); + drop(held_open); +} + +#[tokio::test] +async fn default_large_pipe_checks_trailer_error_after_exit() { + let gateway = TestGateway::start(Scenario::ErrorAfterExit).await; + let baseline = std::env::var_os("OPENSHELL_STREAMING_TEST_BASELINE_CLI"); + let executable = baseline + .as_deref() + .map_or_else(|| Path::new(env!("CARGO_BIN_EXE_openshell")), Path::new); + let mut child = gateway.command(executable, false).spawn().unwrap(); + // The encoded request includes metadata, so a 1 MiB payload selects the + // existing streaming transport even without the new command-line flag. + send_input(&mut child, &vec![b'x'; 1024 * 1024]).await; + let output = finish(child).await; + gateway.assert_one_stream(); + assert!( + !output.status.success(), + "Exit must not hide a failing trailer" + ); + assert!(String::from_utf8_lossy(&output.stderr).contains("failure after exit")); +} + +#[tokio::test] +async fn streaming_missing_exit_is_not_success() { + let gateway = TestGateway::start(Scenario::MissingExit).await; + let mut child = gateway.spawn(true); + let held_open = child.stdin.take().unwrap(); + let output = finish(child).await; + assert!(!output.status.success()); + assert_eq!(output.stdout, b"incomplete\n"); + assert!(String::from_utf8_lossy(&output.stderr).contains("exit status")); + gateway.assert_one_stream(); + drop(held_open); +} + +#[tokio::test] +async fn streaming_disconnect_does_not_start_another_command() { + let gateway = TestGateway::start(Scenario::Disconnect).await; + let mut child = gateway.spawn(true); + let held_open = child.stdin.take().unwrap(); + let output = finish(child).await; + assert!(!output.status.success()); + let stderr = String::from_utf8_lossy(&output.stderr); + let normalized = stderr + .replace(['│', '×'], " ") + .split_whitespace() + .collect::>() + .join(" "); + assert!(normalized.contains("relay disconnected"), "{stderr}"); + gateway.assert_one_stream(); + drop(held_open); +} + +#[tokio::test] +async fn streaming_accepts_exact_input_limit() { + let gateway = TestGateway::start(Scenario::Count).await; + let mut child = gateway.spawn(true); + send_input(&mut child, &vec![b'x'; STDIN_LIMIT]).await; + let output = finish(child).await; + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + assert_eq!(output.stdout, format!("{STDIN_LIMIT}\n").as_bytes()); + gateway.assert_one_stream(); + assert!(gateway.calls.lock().unwrap().request_ended); +} + +#[tokio::test] +async fn streaming_rejects_excess_input_and_cancels_response() { + let gateway = TestGateway::start(Scenario::AwaitCancellation).await; + let mut child = gateway.spawn(true); + send_input(&mut child, &vec![b'x'; STDIN_LIMIT + 1]).await; + let output = finish(child).await; + assert!(!output.status.success()); + let stderr = String::from_utf8_lossy(&output.stderr); + assert!( + stderr.contains("streamed stdin exceeds the 4 MiB limit"), + "{stderr}" + ); + assert!(stderr.contains("partial input"), "{stderr}"); + gateway.wait_for_stream_end().await; + gateway.assert_one_stream(); + let calls = gateway.calls.lock().unwrap(); + assert!(calls.input_bytes <= STDIN_LIMIT); + assert!(calls.response_cancelled); +} + +#[tokio::test] +async fn streaming_reports_unreadable_stdin_and_cancels_response() { + let gateway = TestGateway::start(Scenario::AwaitCancellation).await; + let input = std::fs::File::open(gateway.config.path()).unwrap(); + let child = gateway + .command(Path::new(env!("CARGO_BIN_EXE_openshell")), true) + .stdin(Stdio::from(input)) + .spawn() + .unwrap(); + let output = finish(child).await; + assert!(!output.status.success()); + let stderr = String::from_utf8_lossy(&output.stderr); + assert!(stderr.contains("Is a directory"), "{stderr}"); + gateway.wait_for_stream_end().await; + gateway.assert_one_stream(); + let calls = gateway.calls.lock().unwrap(); + assert_eq!(calls.input_bytes, 0); + assert!(calls.response_cancelled); +} + +#[tokio::test] +async fn streaming_cancelled_request_ends_after_delayed_poll() { + let gateway = TestGateway::start(Scenario::ReadAfterCancellation).await; + let input = std::fs::File::open(gateway.config.path()).unwrap(); + let child = gateway + .command(Path::new(env!("CARGO_BIN_EXE_openshell")), true) + .stdin(Stdio::from(input)) + .spawn() + .unwrap(); + let output = finish(child).await; + assert!(!output.status.success()); + assert!(String::from_utf8_lossy(&output.stderr).contains("Is a directory")); + gateway.wait_for_stream_end().await; + gateway.assert_one_stream(); + let calls = gateway.calls.lock().unwrap(); + assert_eq!(calls.input_bytes, 0); + assert!(calls.response_cancelled); + assert!( + calls.request_ended || calls.request_error, + "cancellation must terminate the request, whether as EOF or a transport error" + ); +} + +#[tokio::test] +async fn cancelled_tonic_request_can_decode_as_end_of_input() { + struct NoFrames; + impl tonic::codec::Decoder for NoFrames { + type Item = (); + type Error = Status; + + fn decode(&mut self, _: &mut tonic::codec::DecodeBuf<'_>) -> Result, Status> { + Err(Status::internal("the cancelled body contains no message")) + } + } + + // Inject cancellation at the decoder boundary rather than relying on the + // HTTP/2 transport to choose the same terminal outcome on every platform. + let cancelled: Result, Status> = + Err(Status::cancelled("synthetic request cancellation")); + let body = http_body_util::StreamBody::new(futures::stream::iter([cancelled])); + let mut request = tonic::Streaming::new_request(NoFrames, body, None, None); + assert!(request.message().await.unwrap().is_none()); + + let unavailable: Result, Status> = + Err(Status::unavailable("synthetic transport failure")); + let body = http_body_util::StreamBody::new(futures::stream::iter([unavailable])); + let mut request = tonic::Streaming::new_request(NoFrames, body, None, None); + assert_eq!( + request.message().await.unwrap_err().code(), + tonic::Code::Unavailable, + ); +} + +#[tokio::test] +async fn default_small_pipe_retains_unary_exec() { + let gateway = TestGateway::start(Scenario::Echo).await; + let mut child = gateway.spawn(false); + send_input(&mut child, b"finite input").await; + let output = finish(child).await; + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + assert_eq!(output.stdout, b"finite input"); + let calls = gateway.calls.lock().unwrap(); + assert_eq!(calls.unary.len(), 1); + assert!(calls.starts.is_empty()); + assert!(!calls.unary[0].tty); +} + +#[tokio::test] +async fn default_large_finite_pipe_retains_streaming_transport() { + let gateway = TestGateway::start(Scenario::Count).await; + let mut child = gateway.spawn(false); + send_input(&mut child, &vec![b'x'; STDIN_LIMIT]).await; + let output = finish(child).await; + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + assert_eq!(output.stdout, format!("{STDIN_LIMIT}\n").as_bytes()); + gateway.assert_one_stream(); + assert!(gateway.calls.lock().unwrap().request_ended); +} + +#[tokio::test] +async fn default_oversized_pipe_is_rejected_before_launch() { + let gateway = TestGateway::start(Scenario::Count).await; + let mut child = gateway.spawn(false); + send_input(&mut child, &vec![b'x'; STDIN_LIMIT + 1]).await; + let output = finish(child).await; + assert!(!output.status.success()); + assert!( + String::from_utf8_lossy(&output.stderr).contains("piped stdin exceeds the 4 MiB limit") + ); + let calls = gateway.calls.lock().unwrap(); + assert_eq!(calls.lookups, 1); + assert!(calls.unary.is_empty()); + assert!(calls.starts.is_empty()); +} diff --git a/docs/how-it-works/sandboxes/overview.mdx b/docs/how-it-works/sandboxes/overview.mdx index b20756c2e3..d199e569d4 100644 --- a/docs/how-it-works/sandboxes/overview.mdx +++ b/docs/how-it-works/sandboxes/overview.mdx @@ -317,19 +317,19 @@ Pipe stdin into the command: echo "hello" | openshell sandbox exec -n my-sandbox -- cat ``` -The CLI sends small piped input in one request for compatibility with older -gateways. It streams larger input in bounded frames, including for commands -without a TTY, so input is not limited by the gateway's per-message request -size. Piped input is limited to 4 MiB; use `sandbox upload` for larger files. -The CLI closes remote stdin when the pipe reaches EOF and continues -reading command output until the command finishes. - -The command's exit code is propagated to the CLI, so `exec` works in scripts that check return codes. -Large stdout and stderr streams are delivered before a successful exit. A slow -reader backpressures the command. If output delivery fails, `exec` returns a -failure instead of reporting success with incomplete output. -If a background process keeps stdout or stderr open for more than 30 seconds -after the command exits, `exec` reports an output delivery failure. +By default, the CLI reads piped input to EOF before starting the command. It sends small input in one request and larger input in bounded frames. Piped input is limited to 4 MiB; the CLI rejects larger input before launch. Use `sandbox upload` for larger files. + +For commands that must respond while stdin remains open, add `--stream-stdin`: + +```shell +openshell sandbox exec -n my-sandbox --stream-stdin -- cat +``` + +This mode starts the command before stdin reaches EOF and forwards input as it arrives. It disables TTY allocation and keeps stdout and stderr separate. You cannot combine it with `--tty`; `--no-tty` is allowed. Closing stdin ends the command's input while the CLI continues reading output until completion. + +The 4 MiB limit applies to total stdin bytes for the command, including with `--stream-stdin`. Exceeding the limit cancels the execution and returns an error. The command may already have processed earlier input, and cancellation does not undo that work. + +The command's exit code is propagated to the CLI, so `exec` works in scripts that check return codes. Large stdout and stderr streams are delivered before a successful exit. A slow reader backpressures the command. If output delivery fails, `exec` returns a failure instead of reporting success with incomplete output. If a background process keeps stdout or stderr open for more than 30 seconds after the command exits, `exec` reports an output delivery failure. A failed final gRPC status also makes the CLI report failure, even after exit code zero. The CLI does not automatically retry the execution. Run an interactive shell with a TTY: @@ -348,6 +348,7 @@ OpenShell allocates a TTY automatically when both stdin and stdout are terminals | `--timeout` | Command timeout in seconds. `0` disables the timeout. | | `--tty` | Force TTY allocation. | | `--no-tty` | Disable TTY allocation even when attached to a terminal. | +| `--stream-stdin` | Start before stdin EOF without a TTY; retain the 4 MiB total input limit. | | `--no-login-shell`| Run the command without sourcing shell login startup files. | | `--env` | Set an environment variable for the command (`KEY=VALUE`, repeatable). | diff --git a/e2e/rust/tests/sandbox_lifecycle.rs b/e2e/rust/tests/sandbox_lifecycle.rs index a7c306fec9..65746ea656 100644 --- a/e2e/rust/tests/sandbox_lifecycle.rs +++ b/e2e/rust/tests/sandbox_lifecycle.rs @@ -116,7 +116,7 @@ async fn delete_sandbox(name: &str) { #[serial(sandbox_lifecycle)] async fn sandbox_exec_large_output_is_complete() { const BYTES: usize = 8 * 1024 * 1024; - let mut sandbox = SandboxGuard::create(&[]) + let mut sandbox = SandboxGuard::create_with_gateway_default(&[]) .await .expect("create sandbox for large output"); @@ -221,7 +221,7 @@ async fn sandbox_exec_large_output_is_complete() { #[tokio::test] #[serial(sandbox_lifecycle)] async fn piped_exec_stdin_crosses_grpc_message_limit() { - let mut sandbox = SandboxGuard::create(&[]) + let mut sandbox = SandboxGuard::create_with_gateway_default(&[]) .await .expect("create sandbox for streamed stdin"); @@ -363,6 +363,255 @@ async fn piped_exec_stdin_crosses_grpc_message_limit() { sandbox.cleanup().await; } +fn streaming_exec_command(sandbox_name: &str, argv: &[&str]) -> tokio::process::Command { + let mut command = openshell_cmd(); + command + .args([ + "sandbox", + "exec", + "--name", + sandbox_name, + "--stream-stdin", + "--no-login-shell", + "--", + ]) + .args(argv) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()); + command +} + +#[tokio::test] +#[serial(sandbox_lifecycle)] +async fn streaming_exec_exchanges_before_stdin_eof() { + let mut sandbox = SandboxGuard::create_with_gateway_default(&[]) + .await + .expect("create sandbox for streaming exchange"); + let script = "printf 'ready:%s\\n' \"$$\"; \ + while IFS= read -r request; do printf '%s:%s\\n' \"$$\" \"$request\"; done; \ + printf 'final-stdout\\n'; printf 'final-stderr\\n' >&2; exit 7"; + let mut child = streaming_exec_command(&sandbox.name, &["sh", "-c", script]) + .spawn() + .expect("spawn streaming exchange"); + let mut input = child.stdin.take().expect("streaming stdin"); + let mut output = BufReader::new(child.stdout.take().expect("streaming stdout")).lines(); + let mut errors = BufReader::new(child.stderr.take().expect("streaming stderr")).lines(); + + // Receiving each response while retaining stdin proves that neither the CLI + // nor the gateway waits for EOF before starting or servicing the command. + tokio::time::timeout(Duration::from_secs(30), async { + let ready = output + .next_line() + .await + .expect("read readiness") + .expect("readiness line"); + let pid = ready.strip_prefix("ready:").expect("remote PID"); + assert!(pid.parse::().is_ok(), "invalid remote PID: {ready}"); + for request in ["first", "second"] { + input + .write_all(format!("{request}\n").as_bytes()) + .await + .expect("write request"); + let response = output + .next_line() + .await + .expect("read response") + .expect("response line"); + assert_eq!(response, format!("{pid}:{request}")); + } + assert!(child.try_wait().expect("poll streaming command").is_none()); + drop(input); + + let (stdout, stderr, status) = tokio::join!( + async { + assert_eq!( + output + .next_line() + .await + .expect("read final stdout") + .as_deref(), + Some("final-stdout") + ); + output.next_line().await.expect("read stdout EOF") + }, + async { + assert_eq!( + errors + .next_line() + .await + .expect("read final stderr") + .as_deref(), + Some("final-stderr") + ); + errors.next_line().await.expect("read stderr EOF") + }, + child.wait(), + ); + assert!(stdout.is_none(), "unexpected trailing stdout: {stdout:?}"); + assert!(stderr.is_none(), "unexpected trailing stderr: {stderr:?}"); + assert_eq!(status.expect("wait for streaming exchange").code(), Some(7)); + }) + .await + .expect("streaming exchange timed out before or after stdin EOF"); + sandbox.cleanup().await; +} + +#[tokio::test] +#[serial(sandbox_lifecycle)] +async fn streaming_exec_exits_while_stdin_remains_open() { + let mut sandbox = SandboxGuard::create_with_gateway_default(&[]) + .await + .expect("create sandbox for idle streaming stdin"); + let mut child = streaming_exec_command( + &sandbox.name, + &[ + "sh", + "-c", + "printf 'stdout\\n'; printf 'stderr\\n' >&2; exit 9", + ], + ) + .spawn() + .expect("spawn command with idle stdin"); + let input = child.stdin.take().expect("idle stdin remains open"); + let output = tokio::time::timeout(Duration::from_secs(30), child.wait_with_output()) + .await + .expect("command waited for idle stdin to close") + .expect("wait for early command exit"); + assert_eq!(output.status.code(), Some(9)); + assert_eq!(output.stdout, b"stdout\n"); + assert_eq!(output.stderr, b"stderr\n"); + drop(input); + sandbox.cleanup().await; +} + +#[tokio::test] +#[serial(sandbox_lifecycle)] +async fn streaming_exec_enforces_cumulative_stdin_limit() { + const LIMIT: usize = 4 * 1024 * 1024; + let mut sandbox = SandboxGuard::create_with_gateway_default(&[]) + .await + .expect("create sandbox for streaming input limit"); + for size in [LIMIT, LIMIT + 1] { + let mut child = streaming_exec_command(&sandbox.name, &["wc", "-c"]) + .spawn() + .expect("spawn streaming byte count"); + let mut input = child.stdin.take().expect("streaming byte input"); + let (write_result, output) = tokio::time::timeout(Duration::from_secs(30), async { + tokio::join!( + async { + let result = input.write_all(&vec![b'x'; size]).await; + drop(input); + result + }, + child.wait_with_output(), + ) + }) + .await + .expect("streaming input limit test timed out"); + let output = output.expect("wait for streaming input limit"); + if size == LIMIT { + write_result.expect("write input at the limit"); + assert!( + output.status.success(), + "streaming count failed: {}", + String::from_utf8_lossy(&output.stderr) + ); + assert_eq!( + String::from_utf8_lossy(&output.stdout).trim(), + LIMIT.to_string() + ); + assert!(output.stderr.is_empty()); + } else { + // Streaming can execute a prefix before discovering the excess byte. + // The guarantee is an explicit error, not rollback of remote effects. + if let Err(error) = write_result { + assert_eq!(error.kind(), std::io::ErrorKind::BrokenPipe); + } + assert!( + !output.status.success(), + "oversized streaming input was accepted" + ); + assert!( + String::from_utf8_lossy(&output.stderr).contains("4 MiB limit"), + "missing input-limit diagnostic: {}", + String::from_utf8_lossy(&output.stderr) + ); + } + } + sandbox.cleanup().await; +} + +#[tokio::test] +#[serial(sandbox_lifecycle)] +async fn streaming_exec_disconnect_stops_command_without_relaunch() { + const STARTS: &str = "/tmp/stream-stdin-starts"; + let mut sandbox = SandboxGuard::create_with_gateway_default(&[]) + .await + .expect("create sandbox for streaming disconnect"); + let script = format!( + "printf '%s\\n' \"$$\" >> {STARTS}; printf 'ready:%s\\n' \"$$\"; while IFS= read -r request; do :; done" + ); + let mut child = streaming_exec_command(&sandbox.name, &["sh", "-c", &script]) + .spawn() + .expect("spawn command for disconnect"); + let input = child.stdin.take().expect("keep disconnect stdin open"); + let mut output = BufReader::new(child.stdout.take().expect("disconnect stdout")).lines(); + let ready = tokio::time::timeout(Duration::from_secs(30), output.next_line()) + .await + .expect("disconnect command did not start") + .expect("read disconnect readiness") + .expect("disconnect readiness line"); + let pid: u32 = ready + .strip_prefix("ready:") + .expect("disconnect remote PID") + .parse() + .expect("numeric remote PID"); + // The fixture waits in the shell's builtin read, so this checks the exec + // process itself without assuming cancellation kills detached descendants. + let probe = format!( + "if kill -0 {pid} 2>/dev/null; then printf 'running\\n'; else printf 'stopped\\n'; fi; wc -l < {STARTS}" + ); + let before = sandbox + .exec(&["sh", "-c", &probe]) + .await + .expect("inspect live streaming exec"); + assert_eq!( + before.lines().map(str::trim).collect::>(), + ["running", "1"], + "the cleanup probe must see the running command before disconnect" + ); + child.kill().await.expect("disconnect streaming client"); + drop(input); + + tokio::time::timeout(Duration::from_secs(30), async { + loop { + let state = sandbox + .exec(&["sh", "-c", &probe]) + .await + .expect("inspect disconnected exec"); + let lines: Vec<_> = state.lines().map(str::trim).collect(); + assert_eq!(lines.get(1), Some(&"1"), "exec was relaunched: {state}"); + if lines.first() == Some(&"stopped") { + break; + } + sleep(Duration::from_millis(100)).await; + } + }) + .await + .expect("remote exec remained running after disconnect"); + sleep(Duration::from_secs(2)).await; + let state = sandbox + .exec(&["sh", "-c", &probe]) + .await + .expect("check no delayed relaunch"); + assert_eq!( + state.lines().map(str::trim).collect::>(), + ["stopped", "1"] + ); + sandbox.cleanup().await; +} + async fn run_sandbox_lifecycle_command(operation: &str, name: &str) -> String { let mut cmd = openshell_cmd(); cmd.args(["sandbox", operation, name]) diff --git a/skills/openshell-cli/SKILL.md b/skills/openshell-cli/SKILL.md index 4db30ebac0..0c71fdf772 100644 --- a/skills/openshell-cli/SKILL.md +++ b/skills/openshell-cli/SKILL.md @@ -423,6 +423,8 @@ Check whether the command ran before retrying work with side effects. Use Use `--env` only for non-secret values. Attach credentials to the sandbox with a provider instead of passing API keys, tokens, or other secrets to `sandbox exec`. +When a client must read responses before closing stdin, check `openshell sandbox exec --help` for `--stream-stdin`. Use that mode to start the command immediately without a TTY and keep stdout and stderr separate. It conflicts with `--tty`. Total stdin remains limited to 4 MiB; exceeding the limit cancels the command after it may have processed earlier input. Without this flag, piped input is read to EOF and checked before launch. After a stream failure, inspect the command's effects before retrying it. + ### Change attached providers ```bash