Support host-accepted exec-server WebSockets (#39786)

## What changed

- Add `EnvironmentManager::from_accepted_websocket` so embedding hosts can
  construct a remote environment from an already accepted and authenticated
  Axum WebSocket.
- Add `replace_accepted_websocket` to retire the current transport and resume
  the same exec-server session on a host-supplied replacement connection.
- Serialize replacement handoffs, reject overlapping replacements, and release
  the handoff claim when a replacement attempt is cancelled or fails.

## Testing

- Cover initial connection validation and immediate environment readiness.
- Verify replacement retry behavior and recovery of a running process and its
  output after reconnecting.

GitOrigin-RevId: 1f2ab7bcf7b5abbbece5c101801432dc84a8058d
This commit is contained in:
ostepanian
2026-08-20 19:16:15 +00:00
committed by copyberry
parent 0cc80b8db5
commit 9e680a52e7
8 changed files with 928 additions and 25 deletions

View File

@@ -125,6 +125,8 @@ use crate::rpc::RpcClient;
use crate::rpc_server_requests::MAX_IN_FLIGHT_SERVER_CALLS;
use codex_http_client::HttpClientFactory;
#[path = "client/accepted.rs"]
pub(crate) mod accepted;
pub(crate) mod http_client;
mod network_policy_audit;
#[path = "client_recovery.rs"]
@@ -347,7 +349,7 @@ type ConnectionAttempt = OnceCell<ConnectionResult>;
#[derive(Clone)]
pub(crate) struct LazyRemoteExecServerClient {
pub(crate) transport_params: ExecServerTransportParams,
transport_params: Option<ExecServerTransportParams>,
http_client_factory: HttpClientFactory,
recovery_policy: RecoveryPolicy,
// Saves the first startup result so callers share it; retryable failures use reconnect.
@@ -364,7 +366,7 @@ impl LazyRemoteExecServerClient {
http_client_factory: HttpClientFactory,
) -> Self {
Self {
transport_params,
transport_params: Some(transport_params),
http_client_factory,
recovery_policy: RecoveryPolicy::Wait,
startup: Arc::new(ConnectionAttempt::new()),
@@ -385,7 +387,7 @@ impl LazyRemoteExecServerClient {
// Stdio starts a process, so keep it lazy until the environment is used.
if matches!(
self.transport_params,
ExecServerTransportParams::StdioCommand { .. }
Some(ExecServerTransportParams::StdioCommand { .. })
) {
return None;
}
@@ -567,16 +569,23 @@ impl LazyRemoteExecServerClient {
fn can_reconnect(&self) -> bool {
matches!(
self.transport_params,
ExecServerTransportParams::Deferred(_)
| ExecServerTransportParams::WebSocketUrl { .. }
| ExecServerTransportParams::NoiseRendezvous { .. }
Some(
ExecServerTransportParams::Deferred(_)
| ExecServerTransportParams::WebSocketUrl { .. }
| ExecServerTransportParams::NoiseRendezvous { .. }
)
)
}
#[tracing::instrument(name = "codex.exec_server.remote.connect", skip_all)]
async fn connect_once(&self) -> ConnectionResult {
let transport_params = self.transport_params.as_ref().ok_or_else(|| {
Arc::new(ExecServerError::Protocol(
"missing transport params for lazy exec-server connection".to_string(),
))
})?;
let result = ExecServerClient::connect_for_transport(
self.transport_params.clone(),
transport_params.clone(),
self.http_client_factory.clone(),
)
.await

View File

@@ -0,0 +1,262 @@
use std::sync::Arc;
use axum::extract::ws::WebSocket;
use futures::lock::Mutex;
use tokio::sync::OnceCell;
use tokio::sync::OwnedSemaphorePermit;
use tokio::sync::Semaphore;
use tokio::sync::mpsc;
use tokio::sync::watch;
use super::ConnectionStatus;
use super::ExecServerClient;
use super::Inner;
use super::LazyRemoteExecServerClient;
use crate::EnvironmentConnectionState;
use crate::ExecServerClientConnectOptions;
use crate::ExecServerError;
use crate::client_transport::ExecServerReconnectStrategy;
use crate::client_transport::ReconnectAttempt;
use crate::connection::JsonRpcConnection;
use codex_http_client::HttpClientFactory;
struct AcceptedReplacement {
connection: JsonRpcConnection,
permit: OwnedSemaphorePermit,
}
struct AcceptedConnectionSourceInner {
replacements_tx: mpsc::UnboundedSender<AcceptedReplacement>,
replacements_rx: Mutex<mpsc::UnboundedReceiver<AcceptedReplacement>>,
replacement_slots: Arc<Semaphore>,
}
/// Receives authenticated connections supplied by an embedding host.
///
/// The source owns serialization and cancellation cleanup for replacement
/// handoffs. The reconnect loop only asks it for the next connection.
#[derive(Clone)]
pub(crate) struct AcceptedConnectionSource {
inner: Arc<AcceptedConnectionSourceInner>,
options: ExecServerClientConnectOptions,
}
struct AcceptedReplacementSubmission {
source: AcceptedConnectionSource,
permit: OwnedSemaphorePermit,
}
impl AcceptedConnectionSource {
fn new(options: ExecServerClientConnectOptions) -> Self {
let (replacements_tx, replacements_rx) = mpsc::unbounded_channel();
Self {
inner: Arc::new(AcceptedConnectionSourceInner {
replacements_tx,
replacements_rx: Mutex::new(replacements_rx),
replacement_slots: Arc::new(Semaphore::new(1)),
}),
options,
}
}
fn begin_replacement(&self) -> Result<AcceptedReplacementSubmission, ExecServerError> {
let permit = Arc::clone(&self.inner.replacement_slots)
.try_acquire_owned()
.map_err(|_| {
ExecServerError::Protocol(
"an accepted exec-server replacement is already in progress".to_string(),
)
})?;
Ok(AcceptedReplacementSubmission {
source: self.clone(),
permit,
})
}
pub(crate) async fn next_connection(
&self,
session_id: &str,
) -> Result<ReconnectAttempt, ExecServerError> {
let replacement = self
.inner
.replacements_rx
.lock()
.await
.recv()
.await
.ok_or_else(|| {
ExecServerError::Disconnected(
"accepted exec-server replacement channel closed".to_string(),
)
})?;
let mut options = self.options.clone();
options.resume_session_id = Some(session_id.to_string());
Ok(ReconnectAttempt::with_attempt_permit(
replacement.connection,
options,
replacement.permit,
))
}
}
impl AcceptedReplacementSubmission {
fn submit(self, connection: JsonRpcConnection) -> Result<(), ExecServerError> {
self.source
.inner
.replacements_tx
.send(AcceptedReplacement {
connection,
permit: self.permit,
})
.map_err(|_| {
ExecServerError::Disconnected(
"accepted exec-server connection is no longer awaiting replacements"
.to_string(),
)
})
}
}
impl ExecServerClient {
/// Initializes an exec-server client over a WebSocket accepted by an Axum handler.
///
/// The caller owns accepting and authenticating replacement WebSockets.
pub(crate) async fn connect_accepted_websocket(
websocket: WebSocket,
options: ExecServerClientConnectOptions,
) -> Result<Self, ExecServerError> {
if options.resume_session_id.is_some() {
return Err(ExecServerError::Protocol(
"accepted exec-server initial connection cannot resume a session".to_string(),
));
}
let connection_source = AcceptedConnectionSource::new(options.clone());
Self::connect_with_recovery(
JsonRpcConnection::from_axum_websocket(
websocket,
"accepted exec-server websocket".to_string(),
),
options,
Some(ExecServerReconnectStrategy::Accepted(connection_source)),
)
.await
}
/// Supplies an authenticated replacement WebSocket for this accepted client.
///
/// Retires the old transport before resuming the saved session. Returns
/// after handoff; recovery continues asynchronously.
pub(crate) async fn replace_accepted_websocket(
&self,
websocket: WebSocket,
) -> Result<(), ExecServerError> {
self.inner
.accept_replacement_connection(JsonRpcConnection::from_axum_websocket(
websocket,
"accepted exec-server replacement websocket".to_string(),
))
.await
}
}
impl Inner {
/// Hands a replacement connection from the host to the existing accepted client.
///
/// This method coordinates the handoff in this order:
///
/// 1. Verify that this client uses the accepted connection source.
/// 2. Reserve the source so concurrent handoffs are rejected.
/// 3. If the old RPC transport is still connected, move the client into
/// recovery and close that transport before attaching the same session to
/// the replacement.
/// 4. Queue the raw connection for the recovery task. That task creates the
/// new RPC client, runs the initialize/resume handshake with the saved
/// session ID, and recovers the existing processes.
async fn accept_replacement_connection(
self: &Arc<Self>,
connection: JsonRpcConnection,
) -> Result<(), ExecServerError> {
if self.session_id.get().is_none() {
return Err(ExecServerError::Protocol(
"accepted exec-server connection is missing its session ID".to_string(),
));
}
let Some(ExecServerReconnectStrategy::Accepted(connection_source)) =
&self.reconnect_strategy
else {
return Err(ExecServerError::Protocol(
"only an accepted exec-server connection can be replaced directly".to_string(),
));
};
let (current_rpc_client, replacement_submission) = {
let connection = self
.connection
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let current_rpc_client = match &connection.status {
ConnectionStatus::Failed(message) => {
return Err(ExecServerError::Disconnected(message.clone()));
}
ConnectionStatus::Connected(rpc_client) => Some(Arc::clone(rpc_client)),
ConnectionStatus::Recovering => None,
};
let replacement_submission = connection_source.begin_replacement()?;
(current_rpc_client, replacement_submission)
};
if let Some(current_rpc_client) = current_rpc_client {
self.request_recovery(
Arc::clone(&current_rpc_client),
"exec-server connection replaced".to_string(),
);
current_rpc_client.close_transport().await;
}
// Synchronize the enqueue with terminal recovery. Recovery can time out
// while the handoff is waiting for the old transport to close; in that
// case the host must not receive success for a socket that no task will
// consume. Holding the connection lock through the synchronous send
// makes either the failure or the enqueue win the race unambiguously.
let connection_state = self
.connection
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let ConnectionStatus::Failed(message) = &connection_state.status {
return Err(ExecServerError::Disconnected(message.clone()));
}
replacement_submission.submit(connection)
}
}
#[cfg(test)]
#[path = "accepted_tests.rs"]
mod tests;
impl LazyRemoteExecServerClient {
pub(crate) fn from_connected(
client: ExecServerClient,
http_client_factory: HttpClientFactory,
) -> Self {
let environment_connection_state_tx =
watch::channel(EnvironmentConnectionState::Connected).0;
client.attach_environment_connection_state(environment_connection_state_tx.clone());
Self {
transport_params: None,
http_client_factory,
recovery_policy: super::RecoveryPolicy::Wait,
startup: std::sync::Arc::new(OnceCell::new_with(Some(Ok(client)))),
current_client: std::sync::Arc::new(std::sync::Mutex::new(None)),
reconnect: std::sync::Arc::new(std::sync::Mutex::new(None)),
environment_connection_state_tx,
}
}
pub(crate) async fn replace_accepted_websocket(
&self,
websocket: WebSocket,
) -> Result<(), ExecServerError> {
self.get()
.await?
.replace_accepted_websocket(websocket)
.await
}
}

View File

@@ -0,0 +1,34 @@
use tokio::sync::oneshot;
use super::AcceptedConnectionSource;
use crate::ExecServerClientConnectOptions;
#[tokio::test]
async fn replacement_claim_is_released_when_handoff_is_cancelled() {
let source = AcceptedConnectionSource::new(ExecServerClientConnectOptions::default());
let (claimed_tx, claimed_rx) = oneshot::channel();
let (_release_tx, release_rx) = oneshot::channel::<()>();
let handoff = tokio::spawn({
let source = source.clone();
async move {
let _submission = source
.begin_replacement()
.expect("the first replacement should claim the handoff");
claimed_tx
.send(())
.expect("the test should wait for the claim");
let _ = release_rx.await;
}
});
claimed_rx.await.expect("the handoff should be claimed");
assert!(source.begin_replacement().is_err());
handoff.abort();
let _ = handoff.await;
source
.begin_replacement()
.expect("a new replacement should be accepted after cancellation");
}

View File

@@ -322,7 +322,7 @@ impl Inner {
}
}
fn request_recovery(
pub(super) fn request_recovery(
self: &Arc<Self>,
failed_rpc_client: Arc<RpcClient>,
disconnect_message: String,
@@ -388,8 +388,8 @@ impl Inner {
let mut registry_retry_attempt = 0;
let last_error = loop {
match timeout_at(deadline, self.resume_once(&session_id)).await {
Ok(Ok(candidate)) => {
if !candidate.is_disconnected() && self.install_recovered_client(candidate) {
Ok(Ok((rpc_client, _attempt))) => {
if !rpc_client.is_disconnected() && self.install_recovered_client(rpc_client) {
return;
}
}
@@ -466,12 +466,13 @@ impl Inner {
async fn resume_once(
self: &Arc<Self>,
session_id: &str,
) -> Result<Arc<RpcClient>, ExecServerError> {
) -> Result<(Arc<RpcClient>, Option<tokio::sync::OwnedSemaphorePermit>), ExecServerError> {
let reconnect_strategy = self
.reconnect_strategy
.as_ref()
.ok_or_else(|| ExecServerError::Protocol("missing reconnect strategy".to_string()))?;
let (connection, options) = reconnect_strategy.resume(session_id).await?;
let attempt = reconnect_strategy.resume(session_id).await?;
let (connection, options, attempt_permit) = attempt.into_parts();
let (rpc_client, events_rx) = RpcClient::new(connection);
let rpc_client = Arc::new(rpc_client);
let client = ExecServerClient {
@@ -486,7 +487,7 @@ impl Inner {
client.initialize_rpc(&rpc_client, options).await?;
self.recover_processes(&rpc_client).await?;
Ok(rpc_client)
Ok((rpc_client, attempt_permit))
}
async fn recover_processes(

View File

@@ -5,6 +5,7 @@ use std::time::Duration;
use tokio::io::AsyncBufReadExt;
use tokio::io::BufReader;
use tokio::process::Command;
use tokio::sync::OwnedSemaphorePermit;
use tokio::time::Instant;
use tokio::time::sleep;
use tokio::time::timeout;
@@ -21,6 +22,7 @@ use codex_websocket_client::WebSocketTlsMode;
use crate::ExecServerClient;
use crate::ExecServerError;
use crate::client::accepted::AcceptedConnectionSource;
use crate::client::is_retryable_registry_error;
use crate::client::registry_recovery_retry_delay;
use crate::client_api::DEFAULT_REMOTE_EXEC_SERVER_CONNECT_TIMEOUT;
@@ -46,6 +48,51 @@ const INITIAL_REGISTRY_MAX_RETRIES: u32 = 4;
const INITIAL_REGISTRY_REQUEST_TIMEOUT: Duration = Duration::from_secs(6);
const INITIAL_REGISTRY_OPERATION_TIMEOUT: Duration = Duration::from_secs(14);
/// Everything the recovery loop needs for one connection attempt.
///
/// An attempt may also carry a permit whose lifetime must extend until the
/// attempt finishes.
pub(crate) struct ReconnectAttempt {
connection: JsonRpcConnection,
options: ExecServerClientConnectOptions,
attempt_permit: Option<OwnedSemaphorePermit>,
}
impl ReconnectAttempt {
pub(crate) fn new(
connection: JsonRpcConnection,
options: ExecServerClientConnectOptions,
) -> Self {
Self {
connection,
options,
attempt_permit: None,
}
}
pub(crate) fn with_attempt_permit(
connection: JsonRpcConnection,
options: ExecServerClientConnectOptions,
attempt_permit: OwnedSemaphorePermit,
) -> Self {
Self {
connection,
options,
attempt_permit: Some(attempt_permit),
}
}
pub(crate) fn into_parts(
self,
) -> (
JsonRpcConnection,
ExecServerClientConnectOptions,
Option<OwnedSemaphorePermit>,
) {
(self.connection, self.options, self.attempt_permit)
}
}
/// Reopens the transport for one logical exec-server client session.
///
/// URL connections reuse their configured endpoint. Noise connections retain
@@ -53,6 +100,7 @@ const INITIAL_REGISTRY_OPERATION_TIMEOUT: Duration = Duration::from_secs(14);
/// every physical connection attempt.
#[derive(Clone)]
pub(crate) enum ExecServerReconnectStrategy {
Accepted(AcceptedConnectionSource),
WebSocket(RemoteExecServerConnectArgs),
NoiseRendezvous {
provider: Arc<dyn NoiseRendezvousConnectProvider>,
@@ -68,13 +116,14 @@ impl ExecServerReconnectStrategy {
pub(crate) async fn resume(
&self,
session_id: &str,
) -> Result<(JsonRpcConnection, ExecServerClientConnectOptions), ExecServerError> {
) -> Result<ReconnectAttempt, ExecServerError> {
match self {
Self::Accepted(source) => source.next_connection(session_id).await,
Self::WebSocket(args) => {
let mut args = args.clone();
args.resume_session_id = Some(session_id.to_string());
let connection = ExecServerClient::open_websocket_connection(&args).await?;
Ok((connection, args.into()))
Ok(ReconnectAttempt::new(connection, args.into()))
}
Self::NoiseRendezvous {
provider,
@@ -85,16 +134,19 @@ impl ExecServerReconnectStrategy {
http_client_factory,
} => {
let bundle = provider.connect_bundle(identity.public_key()).await?;
ExecServerClient::open_noise_rendezvous_connection(NoiseRendezvousConnectArgs {
bundle,
harness_identity: identity.clone(),
client_name: client_name.clone(),
connect_timeout: *connect_timeout,
initialize_timeout: *initialize_timeout,
resume_session_id: Some(session_id.to_string()),
http_client_factory: http_client_factory.clone(),
})
.await
let (connection, options) = ExecServerClient::open_noise_rendezvous_connection(
NoiseRendezvousConnectArgs {
bundle,
harness_identity: identity.clone(),
client_name: client_name.clone(),
connect_timeout: *connect_timeout,
initialize_timeout: *initialize_timeout,
resume_session_id: Some(session_id.to_string()),
http_client_factory: http_client_factory.clone(),
},
)
.await?;
Ok(ReconnectAttempt::new(connection, options))
}
}
}

View File

@@ -46,6 +46,9 @@ use tokio_util::task::AbortOnDropHandle;
use tracing::Instrument;
use tracing::instrument::WithSubscriber;
#[path = "environment/accepted.rs"]
mod accepted;
pub const CODEX_EXEC_SERVER_URL_ENV_VAR: &str = "CODEX_EXEC_SERVER_URL";
pub const CODEX_EXEC_SERVER_NOISE_REGISTRY_URL_ENV_VAR: &str =
"CODEX_EXEC_SERVER_NOISE_REGISTRY_URL";
@@ -770,6 +773,13 @@ impl Environment {
http_client_factory: HttpClientFactory,
) -> Self {
let client = LazyRemoteExecServerClient::new(remote_transport, http_client_factory);
Self::remote_with_client(client, local_runtime_paths)
}
pub(crate) fn remote_with_client(
client: LazyRemoteExecServerClient,
local_runtime_paths: Option<ExecServerRuntimePaths>,
) -> Self {
let exec_backend: Arc<dyn ExecBackend> = Arc::new(RemoteProcess::new(client.clone()));
let filesystem: Arc<dyn ExecutorFileSystem> =
Arc::new(RemoteFileSystem::new(client.clone()));

View File

@@ -0,0 +1,64 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::RwLock;
use super::Environment;
use super::EnvironmentManager;
use super::validate_environment_id;
use crate::ExecServerClient;
use crate::ExecServerClientConnectOptions;
use crate::ExecServerError;
use crate::client::LazyRemoteExecServerClient;
use axum::extract::ws::WebSocket;
use codex_http_client::HttpClientFactory;
impl EnvironmentManager {
/// Builds a manager around a WebSocket already accepted and authenticated by its host.
///
/// The manager owns client construction, session initialization, and later
/// recovery. The host only supplies the initial socket and authenticated
/// replacement sockets through [`Self::replace_accepted_websocket`].
pub async fn from_accepted_websocket(
environment_id: String,
websocket: WebSocket,
options: ExecServerClientConnectOptions,
http_client_factory: HttpClientFactory,
) -> Result<Self, ExecServerError> {
validate_environment_id(&environment_id)?;
let client = ExecServerClient::connect_accepted_websocket(websocket, options).await?;
let client =
LazyRemoteExecServerClient::from_connected(client, http_client_factory.clone());
let environment = Arc::new(Environment::remote_with_client(
client, /*local_runtime_paths*/ None,
));
Ok(Self {
default_environment: Some(environment_id.clone()),
environments: RwLock::new(HashMap::from([(environment_id, environment)])),
local_environment: None,
local_runtime_paths: None,
http_client_factory,
})
}
/// Hands a replacement WebSocket to an existing accepted environment.
/// Returns after handoff; recovery continues asynchronously.
pub async fn replace_accepted_websocket(
&self,
environment_id: &str,
websocket: WebSocket,
) -> Result<(), ExecServerError> {
let environment = self.get_environment(environment_id).ok_or_else(|| {
ExecServerError::Protocol(format!("environment `{environment_id}` is not configured"))
})?;
environment
.remote_client
.as_ref()
.ok_or_else(|| {
ExecServerError::Protocol(
"local environment does not have a replaceable exec-server client".to_string(),
)
})?
.replace_accepted_websocket(websocket)
.await
}
}

View File

@@ -0,0 +1,471 @@
use std::collections::HashMap;
use std::time::Duration;
use anyhow::Context;
use anyhow::Result;
use axum::Router;
use axum::extract::State;
use axum::extract::WebSocketUpgrade;
use axum::response::IntoResponse;
use axum::routing::any;
use codex_exec_server::EnvironmentManager;
use codex_exec_server::EnvironmentObservedStatus;
use codex_exec_server::EnvironmentStatus;
use codex_exec_server::EnvironmentStatusKind;
use codex_exec_server::ExecParams;
use codex_exec_server::ExecResponse;
use codex_exec_server::ExecServerClientConnectOptions;
use codex_exec_server::InitializeParams;
use codex_exec_server::InitializeResponse;
use codex_exec_server::ProcessId;
use codex_exec_server::ReadParams;
use codex_exec_server::ReadResponse;
use codex_exec_server_protocol::JSONRPCError;
use codex_exec_server_protocol::JSONRPCErrorError;
use codex_exec_server_protocol::JSONRPCMessage;
use codex_exec_server_protocol::JSONRPCNotification;
use codex_exec_server_protocol::JSONRPCRequest;
use codex_exec_server_protocol::JSONRPCResponse;
use codex_http_client::HttpClientFactory;
use codex_http_client::OutboundProxyPolicy;
use codex_utils_path_uri::PathUri;
use futures::SinkExt;
use futures::StreamExt;
use pretty_assertions::assert_eq;
use tokio::net::TcpListener;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio::time::timeout;
use tokio_tungstenite::MaybeTlsStream;
use tokio_tungstenite::WebSocketStream;
use tokio_tungstenite::connect_async;
use tokio_tungstenite::tungstenite::Message;
const TEST_TIMEOUT: Duration = Duration::from_secs(5);
type AcceptedSocket = axum::extract::ws::WebSocket;
const SESSION_ALREADY_ATTACHED_ERROR_CODE: i64 = -32010;
#[tokio::test]
async fn accepted_websocket_rejects_initial_resume_session_id() -> Result<()> {
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
let (_socket, _) = connect_async(&websocket_url).await?;
let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv())
.await?
.context("accepted websocket channel should remain open")?;
let mut options = accepted_options();
options.resume_session_id = Some("session-1".to_string());
let error = EnvironmentManager::from_accepted_websocket(
"environment-1".to_string(),
accepted_websocket,
options,
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
)
.await
.expect_err("initial accepted websocket should reject a resume session ID");
assert!(
error
.to_string()
.contains("initial connection cannot resume a session"),
"unexpected error: {error}"
);
server_task.abort();
let _ = server_task.await;
Ok(())
}
#[tokio::test]
async fn accepted_websocket_environment_is_ready_immediately() -> Result<()> {
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
let (mut socket, manager) =
connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?;
let status_task =
tokio::spawn(async move { manager.get_environment_status("environment-1").await });
let request = receive_jsonrpc(&mut socket).await?;
let JSONRPCMessage::Request(JSONRPCRequest { id, method, .. }) = request else {
anyhow::bail!("expected environment status request, got {request:?}");
};
assert_eq!(method, "environment/status");
send_jsonrpc(
&mut socket,
JSONRPCMessage::Response(JSONRPCResponse {
id,
result: serde_json::to_value(EnvironmentStatus {
status: EnvironmentStatusKind::Ready,
})?,
}),
)
.await?;
assert_eq!(
timeout(TEST_TIMEOUT, status_task).await??,
Some(EnvironmentObservedStatus::Ready)
);
server_task.abort();
let _ = server_task.await;
Ok(())
}
#[tokio::test]
async fn accepted_websocket_replacement_retires_old_socket_and_retries() -> Result<()> {
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
let (mut first_socket, manager) =
connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?;
let (mut rejected_socket, rejected_websocket) =
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
manager
.replace_accepted_websocket("environment-1", rejected_websocket)
.await?;
let previous_socket_event = timeout(TEST_TIMEOUT, first_socket.next())
.await
.context("the previous accepted websocket should be retired before replacement")?;
assert!(
matches!(
previous_socket_event,
None | Some(Ok(Message::Close(_))) | Some(Err(_))
),
"the previous accepted websocket should close before replacement: {previous_socket_event:?}"
);
let initialize = receive_jsonrpc(&mut rejected_socket).await?;
let JSONRPCMessage::Request(JSONRPCRequest { id, method, .. }) = initialize else {
anyhow::bail!("expected replacement initialize request, got {initialize:?}");
};
assert_eq!(method, "initialize");
let (_overlapping_socket, overlapping_websocket) =
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
manager
.replace_accepted_websocket("environment-1", overlapping_websocket)
.await
.expect_err("an overlapping replacement should be rejected");
send_jsonrpc(
&mut rejected_socket,
JSONRPCMessage::Error(JSONRPCError {
id,
error: JSONRPCErrorError {
code: SESSION_ALREADY_ATTACHED_ERROR_CODE,
message: "session session-1 is already attached to another connection".to_string(),
data: None,
},
}),
)
.await?;
let rejected_socket_event = timeout(TEST_TIMEOUT, rejected_socket.next())
.await
.context("rejected replacement websocket should close")?;
assert!(
matches!(
rejected_socket_event,
None | Some(Ok(Message::Close(_))) | Some(Err(_))
),
"rejected replacement websocket should close: {rejected_socket_event:?}"
);
let (mut replacement_socket, replacement_websocket) =
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
manager
.replace_accepted_websocket("environment-1", replacement_websocket)
.await?;
complete_initialize(&mut replacement_socket, "session-1", Some("session-1")).await?;
server_task.abort();
let _ = server_task.await;
Ok(())
}
#[tokio::test]
async fn accepted_websocket_reconnect_recovers_running_process_and_output() -> Result<()> {
let (websocket_url, mut accepted_sockets, server_task) = start_acceptor().await?;
let (mut first_socket, manager) =
connect_executor(&websocket_url, &mut accepted_sockets, "session-1").await?;
let environment = manager
.default_environment()
.context("default environment should be installed")?;
let backend = environment.get_exec_backend();
let process_id = ProcessId::from("process-1");
let process_task = tokio::spawn({
let process_id = process_id.clone();
async move {
backend
.start(ExecParams {
process_id,
argv: vec!["test-command".to_string()],
cwd: PathUri::parse("file:///workspace")?,
env_policy: None,
env: HashMap::new(),
tty: false,
pipe_stdin: false,
arg0: None,
sandbox: None,
enforce_managed_network: false,
managed_network: None,
network_proxy: None,
shell_snapshot: None,
})
.await
.map_err(anyhow::Error::from)
}
});
let request = receive_jsonrpc(&mut first_socket).await?;
let JSONRPCMessage::Request(JSONRPCRequest {
id, method, params, ..
}) = request
else {
anyhow::bail!("expected process start request, got {request:?}");
};
assert_eq!(method, "process/start");
assert_eq!(
serde_json::from_value::<ExecParams>(params.context("process params should exist")?)?
.process_id,
process_id
);
send_jsonrpc(
&mut first_socket,
JSONRPCMessage::Response(JSONRPCResponse {
id,
result: serde_json::to_value(ExecResponse {
process_id: process_id.clone(),
sandbox_type: None,
})?,
}),
)
.await?;
let process = timeout(TEST_TIMEOUT, process_task).await???.process;
first_socket.close(/*close_frame*/ None).await?;
let (mut replacement_socket, replacement_websocket) =
connect_replacement_executor(&websocket_url, &mut accepted_sockets).await?;
manager
.replace_accepted_websocket("environment-1", replacement_websocket)
.await?;
complete_initialize(&mut replacement_socket, "session-1", Some("session-1")).await?;
let request = receive_jsonrpc(&mut replacement_socket).await?;
let JSONRPCMessage::Request(JSONRPCRequest {
id, method, params, ..
}) = request
else {
anyhow::bail!("expected recovery process read request, got {request:?}");
};
assert_eq!(method, "process/read");
assert_eq!(
serde_json::from_value::<ReadParams>(params.context("read params should exist")?)?,
ReadParams {
process_id: process_id.clone(),
after_seq: Some(0),
max_bytes: None,
wait_ms: Some(0),
}
);
send_jsonrpc(
&mut replacement_socket,
JSONRPCMessage::Response(JSONRPCResponse {
id,
result: serde_json::to_value(ReadResponse {
chunks: Vec::new(),
next_seq: 1,
exited: false,
exit_code: None,
closed: false,
failure: None,
sandbox_denied: false,
})?,
}),
)
.await?;
let read_task = tokio::spawn(async move {
process.read(Some(0), /*max_bytes*/ None, Some(0)).await
});
let request = receive_jsonrpc(&mut replacement_socket).await?;
let JSONRPCMessage::Request(JSONRPCRequest {
id, method, params, ..
}) = request
else {
anyhow::bail!("expected existing process read request, got {request:?}");
};
assert_eq!(method, "process/read");
assert_eq!(
serde_json::from_value::<ReadParams>(params.context("read params should exist")?)?,
ReadParams {
process_id,
after_seq: Some(0),
max_bytes: None,
wait_ms: Some(0),
}
);
let response = ReadResponse {
chunks: Vec::new(),
next_seq: 1,
exited: false,
exit_code: None,
closed: false,
failure: None,
sandbox_denied: false,
};
send_jsonrpc(
&mut replacement_socket,
JSONRPCMessage::Response(JSONRPCResponse {
id,
result: serde_json::to_value(&response)?,
}),
)
.await?;
assert_eq!(timeout(TEST_TIMEOUT, read_task).await???, response);
server_task.abort();
let _ = server_task.await;
Ok(())
}
async fn start_acceptor() -> Result<(
String,
mpsc::UnboundedReceiver<AcceptedSocket>,
JoinHandle<()>,
)> {
let listener = TcpListener::bind("127.0.0.1:0").await?;
let local_addr = listener.local_addr()?;
let (accepted_tx, accepted_rx) = mpsc::unbounded_channel();
let app = Router::new()
.route("/", any(accept_websocket))
.with_state(accepted_tx);
let server_task = tokio::spawn(async move {
let result = axum::serve(listener, app).await;
assert!(
result.is_ok(),
"accepted websocket test server should run: {result:?}"
);
});
Ok((format!("ws://{local_addr}/"), accepted_rx, server_task))
}
async fn accept_websocket(
websocket: WebSocketUpgrade,
State(accepted_tx): State<mpsc::UnboundedSender<AcceptedSocket>>,
) -> impl IntoResponse {
websocket.on_upgrade(move |websocket| async move {
let _ = accepted_tx.send(websocket);
})
}
async fn connect_executor(
websocket_url: &str,
accepted_sockets: &mut mpsc::UnboundedReceiver<AcceptedSocket>,
session_id: &str,
) -> Result<(
WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>,
EnvironmentManager,
)> {
let (mut websocket, _) = connect_async(websocket_url).await?;
let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv())
.await?
.context("accepted websocket channel should remain open")?;
let manager_task = tokio::spawn(EnvironmentManager::from_accepted_websocket(
"environment-1".to_string(),
accepted_websocket,
accepted_options(),
HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault),
));
complete_initialize(&mut websocket, session_id, /*resume_session_id*/ None).await?;
let manager = timeout(TEST_TIMEOUT, manager_task).await???;
Ok((websocket, manager))
}
async fn connect_replacement_executor(
websocket_url: &str,
accepted_sockets: &mut mpsc::UnboundedReceiver<AcceptedSocket>,
) -> Result<(
WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>,
AcceptedSocket,
)> {
let (websocket, _) = connect_async(websocket_url).await?;
let accepted_websocket = timeout(TEST_TIMEOUT, accepted_sockets.recv())
.await?
.context("accepted websocket channel should remain open")?;
Ok((websocket, accepted_websocket))
}
fn accepted_options() -> ExecServerClientConnectOptions {
ExecServerClientConnectOptions {
client_name: "host-test".to_string(),
initialize_timeout: TEST_TIMEOUT,
resume_session_id: None,
}
}
async fn complete_initialize<S>(
websocket: &mut WebSocketStream<S>,
session_id: &str,
resume_session_id: Option<&str>,
) -> Result<()>
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
let initialize = receive_jsonrpc(&mut *websocket).await?;
let JSONRPCMessage::Request(JSONRPCRequest {
id, method, params, ..
}) = initialize
else {
anyhow::bail!("expected initialize request, got {initialize:?}");
};
assert_eq!(method, "initialize");
assert_eq!(
serde_json::from_value::<InitializeParams>(
params.context("initialize request should contain params")?
)?,
InitializeParams {
client_name: "host-test".to_string(),
resume_session_id: resume_session_id.map(str::to_string),
}
);
send_jsonrpc(
&mut *websocket,
JSONRPCMessage::Response(JSONRPCResponse {
id,
result: serde_json::to_value(InitializeResponse {
session_id: session_id.to_string(),
})?,
}),
)
.await?;
let initialized = receive_jsonrpc(&mut *websocket).await?;
assert_eq!(
initialized,
JSONRPCMessage::Notification(JSONRPCNotification {
method: "initialized".to_string(),
params: Some(serde_json::json!({})),
})
);
Ok(())
}
async fn send_jsonrpc<S>(websocket: &mut WebSocketStream<S>, message: JSONRPCMessage) -> Result<()>
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
websocket
.send(Message::Text(serde_json::to_string(&message)?.into()))
.await?;
Ok(())
}
async fn receive_jsonrpc<S>(websocket: &mut WebSocketStream<S>) -> Result<JSONRPCMessage>
where
S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin,
{
loop {
let message = websocket
.next()
.await
.context("accepted websocket should remain open")??;
if let Message::Text(text) = message {
return Ok(serde_json::from_str(&text)?);
}
}
}