diff --git a/codex-rs/exec-server/src/client.rs b/codex-rs/exec-server/src/client.rs index a686606210..6f794dd171 100644 --- a/codex-rs/exec-server/src/client.rs +++ b/codex-rs/exec-server/src/client.rs @@ -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; #[derive(Clone)] pub(crate) struct LazyRemoteExecServerClient { - pub(crate) transport_params: ExecServerTransportParams, + transport_params: Option, 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 diff --git a/codex-rs/exec-server/src/client/accepted.rs b/codex-rs/exec-server/src/client/accepted.rs new file mode 100644 index 0000000000..535eecb694 --- /dev/null +++ b/codex-rs/exec-server/src/client/accepted.rs @@ -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, + replacements_rx: Mutex>, + replacement_slots: Arc, +} + +/// 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, + 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 { + 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 { + 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 { + 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, + 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(¤t_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 + } +} diff --git a/codex-rs/exec-server/src/client/accepted_tests.rs b/codex-rs/exec-server/src/client/accepted_tests.rs new file mode 100644 index 0000000000..f284f1cf39 --- /dev/null +++ b/codex-rs/exec-server/src/client/accepted_tests.rs @@ -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"); +} diff --git a/codex-rs/exec-server/src/client_recovery.rs b/codex-rs/exec-server/src/client_recovery.rs index 97724180a1..0c0999cb18 100644 --- a/codex-rs/exec-server/src/client_recovery.rs +++ b/codex-rs/exec-server/src/client_recovery.rs @@ -322,7 +322,7 @@ impl Inner { } } - fn request_recovery( + pub(super) fn request_recovery( self: &Arc, failed_rpc_client: Arc, 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, session_id: &str, - ) -> Result, ExecServerError> { + ) -> Result<(Arc, Option), 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( diff --git a/codex-rs/exec-server/src/client_transport.rs b/codex-rs/exec-server/src/client_transport.rs index 62fff714fd..318f1adc62 100644 --- a/codex-rs/exec-server/src/client_transport.rs +++ b/codex-rs/exec-server/src/client_transport.rs @@ -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, +} + +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, + ) { + (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, @@ -68,13 +116,14 @@ impl ExecServerReconnectStrategy { pub(crate) async fn resume( &self, session_id: &str, - ) -> Result<(JsonRpcConnection, ExecServerClientConnectOptions), ExecServerError> { + ) -> Result { 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)) } } } diff --git a/codex-rs/exec-server/src/environment.rs b/codex-rs/exec-server/src/environment.rs index b1e73dca3b..e1ea7674da 100644 --- a/codex-rs/exec-server/src/environment.rs +++ b/codex-rs/exec-server/src/environment.rs @@ -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, + ) -> Self { let exec_backend: Arc = Arc::new(RemoteProcess::new(client.clone())); let filesystem: Arc = Arc::new(RemoteFileSystem::new(client.clone())); diff --git a/codex-rs/exec-server/src/environment/accepted.rs b/codex-rs/exec-server/src/environment/accepted.rs new file mode 100644 index 0000000000..600512b080 --- /dev/null +++ b/codex-rs/exec-server/src/environment/accepted.rs @@ -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 { + 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 + } +} diff --git a/codex-rs/exec-server/tests/accepted_websocket.rs b/codex-rs/exec-server/tests/accepted_websocket.rs new file mode 100644 index 0000000000..217ad8e74c --- /dev/null +++ b/codex-rs/exec-server/tests/accepted_websocket.rs @@ -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::(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::(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::(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, + 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>, +) -> 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, + session_id: &str, +) -> Result<( + WebSocketStream>, + 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, +) -> Result<( + WebSocketStream>, + 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( + websocket: &mut WebSocketStream, + 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::( + 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(websocket: &mut WebSocketStream, 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(websocket: &mut WebSocketStream) -> Result +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)?); + } + } +}