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
138 changes: 80 additions & 58 deletions crates/openshell-cli/src/run.rs
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,9 @@ pub use crate::commands::provider_readiness::{ProviderWaitOptions, sandbox_provi

use crate::color::Colorize;
use crate::policy_update::build_policy_update_plan;
use crate::tls::{TlsOptions, grpc_client};
use crate::tls::{
TlsOptions, grpc_client, owner_replica, retry_unrouted, routed_to, with_owner_hint,
};
use futures::StreamExt;
use miette::{IntoDiagnostic, Result, WrapErr, miette};
use openshell_bootstrap::{
Expand Down Expand Up @@ -71,6 +73,7 @@ use std::io::{ErrorKind, IsTerminal, Read, Write};
use std::path::{Path, PathBuf};
use std::process::Command;
use std::time::{Duration, Instant};
use tonic::metadata::AsciiMetadataValue;
use tonic::{Code, Status};

const PROVISIONAL_CONTAINER_EXIT_RECONCILIATION_TIMEOUT: Duration = Duration::from_secs(5);
Expand Down Expand Up @@ -1981,15 +1984,17 @@ pub async fn sandbox_exec_grpc(
let mut client = grpc_client(server, tls).await?;

// Resolve sandbox name to id.
let sandbox = client
.get_sandbox(GetSandboxRequest {
let response = client
.get_sandbox(with_owner_hint(GetSandboxRequest {
name: name.to_string(),
workspace_scope: Some(openshell_core::proto::workspace_selector(
workspace.to_string(),
)),
})
}))
.await
.into_diagnostic()?
.into_diagnostic()?;
let replica = owner_replica(&response);
let sandbox = response
.into_inner()
.sandbox
.ok_or_else(|| miette::miette!("sandbox not found"))?;
Expand Down Expand Up @@ -2077,16 +2082,26 @@ pub async fn sandbox_exec_grpc(
stdin_is_terminal,
std::mem::take(&mut request.stdin),
stdin_rest,
replica.as_ref(),
)
.await;
}

// Make the streaming gRPC call.
let mut stream = client
.exec_sandbox(request)
.await
.into_diagnostic()?
.into_inner();
let retry = replica.is_some().then(|| request.clone());
let mut stream = match (
client
.exec_sandbox(routed_to(request, replica.as_ref()))
.await,
retry,
) {
(Err(status), Some(request)) if retry_unrouted(replica.as_ref(), &status) => {
client.exec_sandbox(request).await
}
(response, _) => response,
}
.into_diagnostic()?
.into_inner();

// Stream output to terminal in real-time.
let mut exit_code = 0i32;
Expand Down Expand Up @@ -2139,7 +2154,7 @@ pub async fn service_forward_tcp(
let (bind_addr, bind_port) = parse_tcp_forward_spec(local, target_port)?;
let mut client = grpc_client(server, tls).await?;

fetch_ready_sandbox_for_forward(&mut client, name, workspace).await?;
let mut replica = fetch_ready_sandbox_for_forward(&mut client, name, workspace).await?;

let listener = tokio::net::TcpListener::bind((bind_addr.as_str(), bind_port))
.await
Expand Down Expand Up @@ -2173,7 +2188,7 @@ pub async fn service_forward_tcp(
}

_ = health_check.tick() => {
fetch_ready_sandbox_for_forward(&mut client, name, workspace).await?;
replica = fetch_ready_sandbox_for_forward(&mut client, name, workspace).await?;
}

accepted = listener.accept() => {
Expand All @@ -2188,6 +2203,7 @@ pub async fn service_forward_tcp(
let service_id = format!("service-forward:{name}:{target_host}:{target_port}");
let fatal_tx = fatal_tx.clone();
let legacy_session_tokens = legacy_session_tokens.clone();
let replica = replica.clone();
tokio::spawn(async move {
if let Err(err) = forward_one_tcp_connection(
&mut client,
Expand All @@ -2198,6 +2214,7 @@ pub async fn service_forward_tcp(
target_port,
service_id,
legacy_session_tokens,
replica.as_ref(),
)
.await
{
Expand Down Expand Up @@ -2237,18 +2254,19 @@ fn forward_requires_session_token(status: &Status) -> bool {
&& status.message().contains("authorization_token is required")
}

/// Confirm the sandbox is still ready, returning the replica that owns it.
async fn fetch_ready_sandbox_for_forward(
client: &mut crate::tls::GrpcClient,
name: &str,
workspace: &str,
) -> Result<Sandbox> {
) -> Result<Option<AsciiMetadataValue>> {
let response = match client
.get_sandbox(GetSandboxRequest {
.get_sandbox(with_owner_hint(GetSandboxRequest {
name: name.to_string(),
workspace_scope: Some(openshell_core::proto::workspace_selector(
workspace.to_string(),
)),
})
}))
.await
{
Ok(response) => response,
Expand All @@ -2260,6 +2278,7 @@ async fn fetch_ready_sandbox_for_forward(
Err(status) => return Err(status).into_diagnostic(),
};

let replica = owner_replica(&response);
let sandbox = response
.into_inner()
.sandbox
Expand All @@ -2273,7 +2292,7 @@ async fn fetch_ready_sandbox_for_forward(
));
}

Ok(sandbox)
Ok(replica)
}

#[derive(Debug)]
Expand Down Expand Up @@ -2339,6 +2358,7 @@ async fn forward_one_tcp_connection(
target_port: u16,
service_id: String,
legacy_session_tokens: std::sync::Arc<std::sync::atomic::AtomicBool>,
replica: Option<&AsciiMetadataValue>,
) -> std::result::Result<(), ForwardTcpConnectionError> {
let mut init = TcpForwardInit {
sandbox: sandbox_name.clone(),
Expand Down Expand Up @@ -2369,7 +2389,7 @@ async fn forward_one_tcp_connection(
}
}

let opened = match open_forward_tcp_stream(client, init.clone()).await {
let opened = match open_forward_tcp_stream(client, init.clone(), replica).await {
Ok(opened) => opened,
Err(status) if session_token.is_none() && forward_requires_session_token(&status) => {
tracing::info!(
Expand All @@ -2387,7 +2407,7 @@ async fn forward_one_tcp_connection(
};
init.authorization_token.clone_from(&token);
session_token = Some(token);
match open_forward_tcp_stream(client, init).await {
match open_forward_tcp_stream(client, init, replica).await {
Ok(opened) => opened,
Err(status) => {
drain_and_shutdown_local_socket(socket).await;
Expand Down Expand Up @@ -2418,22 +2438,15 @@ type OpenForwardTcpStream = (
async fn open_forward_tcp_stream(
client: &mut crate::tls::GrpcClient,
init: TcpForwardInit,
replica: Option<&AsciiMetadataValue>,
) -> std::result::Result<OpenForwardTcpStream, Status> {
use tokio_stream::wrappers::ReceiverStream;

let (tx, rx) = tokio::sync::mpsc::channel::<TcpForwardFrame>(16);
tx.send(TcpForwardFrame {
let init = TcpForwardFrame {
payload: Some(openshell_core::proto::tcp_forward_frame::Payload::Init(
init,
)),
})
.await
.map_err(|_| Status::internal("failed to initialize forward stream"))?;
let response = client
.forward_tcp(ReceiverStream::new(rx))
.await?
.into_inner();
Ok((tx, response))
};
let (tx, response) = crate::ssh::open_forward_tcp(client, init, replica).await;
Ok((tx, response?.into_inner()))
}

async fn revoke_forward_session_token(client: &mut crate::tls::GrpcClient, token: Option<String>) {
Expand Down Expand Up @@ -2561,6 +2574,7 @@ async fn sandbox_exec_streaming_grpc(
stdin_is_terminal: bool,
stdin_prefix: Vec<u8>,
stdin_rest: Option<tokio::sync::mpsc::Receiver<std::io::Result<Vec<u8>>>>,
replica: Option<&AsciiMetadataValue>,
) -> Result<i32> {
#[cfg(unix)]
use openshell_core::proto::ExecSandboxWindowResize;
Expand All @@ -2573,36 +2587,44 @@ async fn sandbox_exec_streaming_grpc(
(0, 0)
};

let (input_tx, input_rx) = tokio::sync::mpsc::channel::<ExecSandboxInput>(64);

// Send the start message with exec metadata.
input_tx
.send(ExecSandboxInput {
payload: Some(exec_sandbox_input::Payload::Start(ExecSandboxRequest {
request_id: String::new(),
sandbox: sandbox.object_name().to_string(),
workspace_scope: Some(openshell_core::proto::workspace_selector(
(sandbox.object_workspace()).to_string(),
)),
command: command.to_vec(),
workdir: workdir.unwrap_or_default().to_string(),
environment: environment.clone(),
no_login_shell,
execution_timeout: proto_execution_timeout(timeout_seconds)?,
stdin: Vec::new(),
tty,
cols,
rows,
})),
})
.await
.into_diagnostic()?;
// The start message carries the exec metadata.
let start = ExecSandboxInput {
payload: Some(exec_sandbox_input::Payload::Start(ExecSandboxRequest {
request_id: String::new(),
sandbox: sandbox.object_name().to_string(),
workspace_scope: Some(openshell_core::proto::workspace_selector(
(sandbox.object_workspace()).to_string(),
)),
command: command.to_vec(),
workdir: workdir.unwrap_or_default().to_string(),
environment: environment.clone(),
no_login_shell,
execution_timeout: proto_execution_timeout(timeout_seconds)?,
stdin: Vec::new(),
tty,
cols,
rows,
})),
};

let mut stream = client
.exec_sandbox_interactive(ReceiverStream::new(input_rx))
let (mut input_tx, input_rx) = tokio::sync::mpsc::channel::<ExecSandboxInput>(64);
input_tx.send(start.clone()).await.into_diagnostic()?;
let mut stream = match client
.exec_sandbox_interactive(routed_to(ReceiverStream::new(input_rx), replica))
Comment thread
FrostGod marked this conversation as resolved.
.await
.into_diagnostic()?
.into_inner();
{
Err(status) if retry_unrouted(replica, &status) => {
let (tx, rx) = tokio::sync::mpsc::channel::<ExecSandboxInput>(64);
tx.send(start).await.into_diagnostic()?;
input_tx = tx;
client
.exec_sandbox_interactive(ReceiverStream::new(rx))
.await
}
response => response,
}
.into_diagnostic()?
.into_inner();

// Raw mode is only appropriate for an interactive terminal, not a pipe.
let raw_guard = if tty && stdin_is_terminal {
Expand Down
Loading
Loading