From 15fdeab3ca3969d2152b5c9bb146066a76b1f0d7 Mon Sep 17 00:00:00 2001 From: Evan Lezar Date: Tue, 14 Jul 2026 14:27:13 +0200 Subject: [PATCH] fix(cli): retry transient sandbox sync disconnects Signed-off-by: Evan Lezar --- crates/openshell-cli/src/ssh.rs | 357 +++++++++++++++++++++++++------- crates/openshell-cli/src/tls.rs | 18 +- 2 files changed, 294 insertions(+), 81 deletions(-) diff --git a/crates/openshell-cli/src/ssh.rs b/crates/openshell-cli/src/ssh.rs index f8de99a692..1886116304 100644 --- a/crates/openshell-cli/src/ssh.rs +++ b/crates/openshell-cli/src/ssh.rs @@ -3,8 +3,8 @@ //! SSH connection and proxy utilities. -use crate::tls::{TlsOptions, grpc_client}; -use miette::{IntoDiagnostic, Result, WrapErr}; +use crate::tls::{GrpcConnectError, TlsOptions, grpc_client}; +use miette::{IntoDiagnostic, Report, Result, WrapErr}; #[cfg(unix)] use nix::sys::signal::{SaFlags, SigAction, SigHandler, SigSet, Signal, sigaction}; use openshell_core::ObjectId; @@ -40,6 +40,56 @@ const FORWARD_LISTENER_PROBE_INTERVAL: Duration = Duration::from_millis(50); /// grace period. const FORWARD_LISTENER_CONNECT_TIMEOUT: Duration = Duration::from_millis(200); +const SYNC_RETRY_ATTEMPTS: usize = 4; +const SYNC_RETRY_DELAY: Duration = Duration::from_secs(2); + +#[derive(Debug, thiserror::Error, miette::Diagnostic)] +#[error("ssh {operation} exited with status {status}")] +struct SshCommandExitError { + operation: &'static str, + status: String, + code: Option, +} + +#[derive(Debug, thiserror::Error, miette::Diagnostic)] +#[error(transparent)] +struct SyncGrpcStatusError(#[from] tonic::Status); + +#[derive(Debug, thiserror::Error, miette::Diagnostic)] +#[error("{context}: {source}")] +struct SyncIoError { + context: String, + #[source] + source: std::io::Error, +} + +fn sync_io_error(context: impl Into, source: std::io::Error) -> Report { + Report::new(SyncIoError { + context: context.into(), + source, + }) +} + +trait SyncIoResultExt { + fn with_sync_context(self, context: impl Into) -> Result; +} + +impl SyncIoResultExt for std::io::Result { + fn with_sync_context(self, context: impl Into) -> Result { + self.map_err(|source| sync_io_error(context, source)) + } +} + +impl SshCommandExitError { + fn from_status(operation: &'static str, status: std::process::ExitStatus) -> Self { + Self { + operation, + status: status.to_string(), + code: status.code(), + } + } +} + #[derive(Clone, Copy, Debug)] pub enum Editor { Vscode, @@ -88,7 +138,7 @@ async fn ssh_session_config( workspace: workspace.to_string(), }) .await - .into_diagnostic()? + .map_err(|error| Report::new(SyncGrpcStatusError(error)))? .into_inner() .sandbox .ok_or_else(|| miette::miette!("sandbox not found"))?; @@ -98,7 +148,7 @@ async fn ssh_session_config( sandbox_id: sandbox.object_id().to_string(), }) .await - .into_diagnostic()?; + .map_err(|error| Report::new(SyncGrpcStatusError(error)))?; let session = response.into_inner(); validate_ssh_session_response(&session) .map_err(|err| miette::miette!("gateway returned invalid SSH session response: {err}"))?; @@ -618,6 +668,7 @@ pub(crate) async fn sandbox_exec_without_exec( } /// What to pack into the tar archive streamed to the sandbox. +#[derive(Clone)] enum UploadSource { /// A single local file or directory. `tar_name` controls the entry name /// inside the archive (e.g. the target basename for file-to-file uploads). @@ -656,7 +707,9 @@ fn write_upload_archive(writer: W, source: UploadSource) -> Result<()> } } } - archive.finish().into_diagnostic()?; + archive + .finish() + .with_sync_context("failed to finish upload tar archive")?; Ok(()) } @@ -670,9 +723,7 @@ fn append_upload_path( Ok(metadata) => metadata, Err(err) if skip_missing && err.kind() == std::io::ErrorKind::NotFound => return Ok(()), Err(err) => { - return Err(err) - .into_diagnostic() - .wrap_err_with(|| format!("failed to stat {}", local_path.display())); + return Err(err).with_sync_context(format!("failed to stat {}", local_path.display())); } }; let file_type = metadata.file_type(); @@ -680,8 +731,10 @@ fn append_upload_path( if file_type.is_file() { archive .append_path_with_name(local_path, archive_path) - .into_diagnostic() - .wrap_err_with(|| format!("failed to add {} to tar archive", archive_path.display()))?; + .with_sync_context(format!( + "failed to add {} to tar archive", + archive_path.display() + ))?; return Ok(()); } @@ -689,13 +742,10 @@ fn append_upload_path( let dir_archive_path = upload_archive_dir_entry_path(archive_path); archive .append_dir(&dir_archive_path, local_path) - .into_diagnostic() - .wrap_err_with(|| { - format!( - "failed to add directory {} to tar archive", - archive_path.display() - ) - })?; + .with_sync_context(format!( + "failed to add directory {} to tar archive", + archive_path.display() + ))?; append_upload_dir_contents(archive, local_path, archive_path)?; return Ok(()); } @@ -723,11 +773,9 @@ fn append_upload_dir_contents( archive_path: &Path, ) -> Result<()> { let mut entries = fs::read_dir(local_path) - .into_diagnostic() - .wrap_err_with(|| format!("failed to read directory {}", local_path.display()))? + .with_sync_context(format!("failed to read directory {}", local_path.display()))? .collect::>>() - .into_diagnostic() - .wrap_err_with(|| format!("failed to read directory {}", local_path.display()))?; + .with_sync_context(format!("failed to read directory {}", local_path.display()))?; entries.sort_by_key(fs::DirEntry::file_name); for entry in entries { @@ -746,8 +794,7 @@ fn append_upload_symlink( metadata: &fs::Metadata, ) -> Result<()> { let target = fs::read_link(local_path) - .into_diagnostic() - .wrap_err_with(|| format!("failed to read symlink {}", local_path.display()))?; + .with_sync_context(format!("failed to read symlink {}", local_path.display()))?; let mut header = tar::Header::new_gnu(); header.set_metadata(metadata); header.set_entry_type(tar::EntryType::Symlink); @@ -755,13 +802,10 @@ fn append_upload_symlink( header.set_cksum(); archive .append_link(&mut header, archive_path, target) - .into_diagnostic() - .wrap_err_with(|| { - format!( - "failed to add symlink {} to tar archive", - archive_path.display() - ) - })?; + .with_sync_context(format!( + "failed to add symlink {} to tar archive", + archive_path.display() + ))?; Ok(()) } @@ -822,9 +866,10 @@ async fn ssh_tar_upload( .into_diagnostic()?; if !status.success() { - return Err(miette::miette!( - "ssh tar extract exited with status {status}" - )); + return Err(Report::new(SshCommandExitError::from_status( + "tar extract", + status, + ))); } Ok(()) @@ -986,18 +1031,15 @@ pub async fn sandbox_sync_up_files( if files.is_empty() { return Ok(()); } - ssh_tar_upload( - server, - name, - dest, - UploadSource::FileList { - base_dir: base_dir.to_path_buf(), - files: files.to_vec(), - archive_prefix: file_list_archive_prefix(local_path), - }, - tls, - workspace, - ) + let source = UploadSource::FileList { + base_dir: base_dir.to_path_buf(), + files: files.to_vec(), + archive_prefix: file_list_archive_prefix(local_path), + }; + retry_sandbox_sync("upload", || { + let source = source.clone(); + async move { ssh_tar_upload(server, name, dest, source, tls, workspace).await } + }) .await } @@ -1032,17 +1074,14 @@ pub async fn sandbox_sync_up( { let (parent, target_name) = split_sandbox_path(path); if parent != "/" { - return ssh_tar_upload( - server, - name, - Some(parent), - UploadSource::SinglePath { - local_path: local_path.to_path_buf(), - tar_name: target_name.into(), - }, - tls, - workspace, - ) + let source = UploadSource::SinglePath { + local_path: local_path.to_path_buf(), + tar_name: target_name.into(), + }; + return retry_sandbox_sync("upload", || { + let source = source.clone(); + async move { ssh_tar_upload(server, name, Some(parent), source, tls, workspace).await } + }) .await; } } @@ -1060,17 +1099,14 @@ pub async fn sandbox_sync_up( directory_upload_prefix(local_path) }; - ssh_tar_upload( - server, - name, - sandbox_path, - UploadSource::SinglePath { - local_path: local_path.to_path_buf(), - tar_name, - }, - tls, - workspace, - ) + let source = UploadSource::SinglePath { + local_path: local_path.to_path_buf(), + tar_name, + }; + retry_sandbox_sync("upload", || { + let source = source.clone(); + async move { ssh_tar_upload(server, name, sandbox_path, source, tls, workspace).await } + }) .await } @@ -1119,10 +1155,10 @@ async fn ssh_run_capture_stdout(session: &SshSessionConfig, command: &str) -> Re .into_diagnostic()? .into_diagnostic()?; if !output.status.success() { - return Err(miette::miette!( - "ssh probe exited with status {}", - output.status - )); + return Err(Report::new(SshCommandExitError::from_status( + "probe", + output.status, + ))); } Ok(String::from_utf8_lossy(&output.stdout).trim().to_string()) } @@ -1178,6 +1214,20 @@ pub async fn sandbox_sync_down( dest: &str, tls: &TlsOptions, workspace: &str, +) -> Result<()> { + retry_sandbox_sync("download", || async { + sandbox_sync_down_once(server, name, sandbox_path, dest, tls, workspace).await + }) + .await +} + +async fn sandbox_sync_down_once( + server: &str, + name: &str, + sandbox_path: &str, + dest: &str, + tls: &TlsOptions, + workspace: &str, ) -> Result<()> { let sandbox_path = validate_sandbox_source_path(sandbox_path)?; let session = ssh_session_config(server, name, tls, workspace).await?; @@ -1192,6 +1242,70 @@ pub async fn sandbox_sync_down( } } +async fn retry_sandbox_sync(operation: &str, run: F) -> Result<()> +where + F: FnMut() -> Fut, + Fut: Future>, +{ + retry_sandbox_sync_with_delay(operation, SYNC_RETRY_ATTEMPTS, SYNC_RETRY_DELAY, run).await +} + +async fn retry_sandbox_sync_with_delay( + operation: &str, + max_attempts: usize, + delay: Duration, + mut run: F, +) -> Result<()> +where + F: FnMut() -> Fut, + Fut: Future>, +{ + debug_assert!(max_attempts > 0); + let mut attempt = 1; + loop { + match run().await { + Ok(()) => return Ok(()), + Err(err) if attempt < max_attempts && sync_error_is_retryable(&err) => { + tracing::warn!( + operation, + attempt, + max_attempts, + error = %err, + "sandbox sync operation failed; retrying" + ); + tokio::time::sleep(delay).await; + attempt += 1; + } + Err(err) => return Err(err), + } + } +} + +fn sync_error_is_retryable(err: &Report) -> bool { + let mut source: Option<&(dyn std::error::Error + 'static)> = Some(err.as_ref()); + while let Some(error) = source { + if let Some(exit) = error.downcast_ref::() { + return exit.code == Some(255); + } + if let Some(status) = error.downcast_ref::() { + return status.0.code() == tonic::Code::Unavailable; + } + if error.downcast_ref::().is_some() { + return true; + } + if let Some(io_error) = error.downcast_ref::() { + return matches!( + io_error.source.kind(), + std::io::ErrorKind::BrokenPipe + | std::io::ErrorKind::ConnectionReset + | std::io::ErrorKind::UnexpectedEof + ); + } + source = error.source(); + } + false +} + /// Stream a tar archive from the sandbox and extract it into a fresh /// destination directory. The source is always wrapped on the sandbox side so /// the host can pick a basename when needed. @@ -1221,8 +1335,7 @@ async fn stream_sandbox_tar( let mut archive = tar::Archive::new(stdout); archive .unpack(&extract_into) - .into_diagnostic() - .wrap_err("failed to extract tar archive from sandbox")?; + .with_sync_context("failed to extract tar archive from sandbox")?; Ok(()) }) .await @@ -1234,9 +1347,10 @@ async fn stream_sandbox_tar( .into_diagnostic()?; if !status.success() { - return Err(miette::miette!( - "ssh tar create exited with status {status}" - )); + return Err(Report::new(SshCommandExitError::from_status( + "tar create", + status, + ))); } Ok(()) } @@ -1677,6 +1791,97 @@ pub fn print_ssh_config(gateway: &str, name: &str, workspace: &str) { mod tests { use super::*; use crate::TEST_ENV_LOCK; + use std::sync::{Arc, atomic::AtomicUsize, atomic::Ordering}; + + fn io_sync_error(kind: std::io::ErrorKind) -> Report { + sync_io_error("test I/O failure", std::io::Error::from(kind)) + } + + fn ssh_exit_error(code: i32) -> Report { + Report::new(SshCommandExitError { + operation: "test", + status: format!("exit status: {code}"), + code: Some(code), + }) + } + + #[tokio::test] + async fn sandbox_sync_retry_succeeds_after_transient_disconnects() { + let attempts = Arc::new(AtomicUsize::new(0)); + let observed_attempts = Arc::clone(&attempts); + retry_sandbox_sync_with_delay("test", 4, Duration::ZERO, move || { + let attempts = Arc::clone(&attempts); + async move { + if attempts.fetch_add(1, Ordering::SeqCst) < 2 { + Err(io_sync_error(std::io::ErrorKind::BrokenPipe)) + } else { + Ok(()) + } + } + }) + .await + .expect("third attempt should succeed"); + + assert_eq!(observed_attempts.load(Ordering::SeqCst), 3); + } + + #[tokio::test] + async fn sandbox_sync_retry_stops_after_exhausting_attempts() { + let attempts = Arc::new(AtomicUsize::new(0)); + let observed_attempts = Arc::clone(&attempts); + let err = retry_sandbox_sync_with_delay("test", 4, Duration::ZERO, move || { + let attempts = Arc::clone(&attempts); + async move { + attempts.fetch_add(1, Ordering::SeqCst); + Err(io_sync_error(std::io::ErrorKind::ConnectionReset)) + } + }) + .await + .expect_err("all attempts should fail"); + + assert!(sync_error_is_retryable(&err)); + assert_eq!(observed_attempts.load(Ordering::SeqCst), 4); + } + + #[tokio::test] + async fn sandbox_sync_retry_returns_permanent_failure_immediately() { + let attempts = Arc::new(AtomicUsize::new(0)); + let observed_attempts = Arc::clone(&attempts); + let err = retry_sandbox_sync_with_delay("test", 4, Duration::ZERO, move || { + let attempts = Arc::clone(&attempts); + async move { + attempts.fetch_add(1, Ordering::SeqCst); + Err(ssh_exit_error(1)) + } + }) + .await + .expect_err("exit 1 should not be retried"); + + assert!(!sync_error_is_retryable(&err)); + assert_eq!(observed_attempts.load(Ordering::SeqCst), 1); + } + + #[test] + fn sandbox_sync_retry_classifies_typed_failures() { + let unavailable = Report::new(SyncGrpcStatusError(tonic::Status::unavailable( + "gateway restarting", + ))); + let invalid = Report::new(SyncGrpcStatusError(tonic::Status::invalid_argument( + "invalid path", + ))); + + assert!(sync_error_is_retryable(&unavailable)); + assert!(sync_error_is_retryable(&ssh_exit_error(255))); + assert!(sync_error_is_retryable(&io_sync_error( + std::io::ErrorKind::UnexpectedEof + ))); + assert!(!sync_error_is_retryable(&invalid)); + assert!(!sync_error_is_retryable(&ssh_exit_error(1))); + assert!(!sync_error_is_retryable(&ssh_exit_error(2))); + assert!(!sync_error_is_retryable(&io_sync_error( + std::io::ErrorKind::PermissionDenied + ))); + } #[test] fn upsert_host_block_appends_when_missing() { diff --git a/crates/openshell-cli/src/tls.rs b/crates/openshell-cli/src/tls.rs index 10df401a5b..608290b437 100644 --- a/crates/openshell-cli/src/tls.rs +++ b/crates/openshell-cli/src/tls.rs @@ -22,6 +22,14 @@ use tonic::service::interceptor::InterceptedService; use tonic::transport::{Certificate, Channel, ClientTlsConfig, Endpoint, Identity}; use tracing::debug; +#[derive(Debug, thiserror::Error, miette::Diagnostic)] +#[error(transparent)] +pub(crate) struct GrpcConnectError(#[from] tonic::transport::Error); + +fn grpc_connect_error(error: tonic::transport::Error) -> miette::Report { + miette::Report::new(GrpcConnectError(error)) +} + /// Concrete gRPC client type used by all commands. pub type GrpcClient = OpenShellClient>; /// Concrete inference client type. @@ -346,7 +354,7 @@ pub async fn build_channel(server: &str, tls: &TlsOptions) -> Result { .http2_adaptive_window(true) .http2_keep_alive_interval(Duration::from_secs(10)) .keep_alive_while_idle(true); - return endpoint.connect().await.into_diagnostic(); + return endpoint.connect().await.map_err(grpc_connect_error); } // When Cloudflare edge bearer auth is active and the server is HTTPS, @@ -367,7 +375,7 @@ pub async fn build_channel(server: &str, tls: &TlsOptions) -> Result { .http2_adaptive_window(true) .http2_keep_alive_interval(Duration::from_secs(10)) .keep_alive_while_idle(true); - return endpoint.connect().await.into_diagnostic(); + return endpoint.connect().await.map_err(grpc_connect_error); } if tls.gateway_insecure && server.starts_with("https://") { @@ -386,7 +394,7 @@ pub async fn build_channel(server: &str, tls: &TlsOptions) -> Result { return endpoint .connect_with_connector(connector) .await - .into_diagnostic(); + .map_err(grpc_connect_error); } let mut endpoint = Endpoint::from_shared(server.to_string()) @@ -419,14 +427,14 @@ pub async fn build_channel(server: &str, tls: &TlsOptions) -> Result { } else if tls.edge_token.is_some() { // Edge bearer mode — routed through tunnel above; if we reach here // the server is not HTTPS so connect plaintext. - return endpoint.connect().await.into_diagnostic(); + return endpoint.connect().await.map_err(grpc_connect_error); } else { // Standard mTLS: private CA + client cert. let materials = require_tls_materials(server, tls)?; build_tonic_tls_config(&materials) }; endpoint = endpoint.tls_config(tls_config).into_diagnostic()?; - endpoint.connect().await.into_diagnostic() + endpoint.connect().await.map_err(grpc_connect_error) } /// Build a gRPC [`OpenShellClient`].