From 59af4a730cda33974053afecf62e6cae00c269a4 Mon Sep 17 00:00:00 2001 From: Ruslan Nigmatullin Date: Tue, 7 Apr 2026 11:36:17 -0700 Subject: [PATCH] app-server: Allow enabling remote control in runtime (#16973) Refresh the feature flag on writes to the config. --- codex-rs/app-server/src/in_process.rs | 1 + codex-rs/app-server/src/lib.rs | 28 ++-- codex-rs/app-server/src/message_processor.rs | 30 +++- .../src/message_processor/tracing_tests.rs | 1 + codex-rs/app-server/src/transport/mod.rs | 1 + .../src/transport/remote_control/mod.rs | 43 +++++- .../src/transport/remote_control/tests.rs | 117 ++++++++++++--- .../src/transport/remote_control/websocket.rs | 133 +++++++++++++----- 8 files changed, 285 insertions(+), 69 deletions(-) diff --git a/codex-rs/app-server/src/in_process.rs b/codex-rs/app-server/src/in_process.rs index 5d8cf052c..7bddd8a71 100644 --- a/codex-rs/app-server/src/in_process.rs +++ b/codex-rs/app-server/src/in_process.rs @@ -398,6 +398,7 @@ fn start_uninitialized(args: InProcessStartArgs) -> InProcessClientHandle { session_source: args.session_source, auth_manager, rpc_transport: AppServerRpcTransport::InProcess, + remote_control_handle: None, }); let mut thread_created_rx = processor.thread_created_receiver(); let mut session = ConnectionSessionState::default(); diff --git a/codex-rs/app-server/src/lib.rs b/codex-rs/app-server/src/lib.rs index 5e02cdb94..a837d9c75 100644 --- a/codex-rs/app-server/src/lib.rs +++ b/codex-rs/app-server/src/lib.rs @@ -564,25 +564,26 @@ pub async fn run_main_with_transport( let auth_manager = AuthManager::shared_from_config(&config, /*enable_codex_api_key_env*/ false); - if config.features.enabled(Feature::RemoteControl) { - let accept_handle = start_remote_control( - config.chatgpt_base_url.clone(), - state_db.clone(), - auth_manager.clone(), - transport_event_tx.clone(), - transport_shutdown_token.clone(), - app_server_client_name_rx, - ) - .await?; - transport_accept_handles.push(accept_handle); - } - if transport_accept_handles.is_empty() { + let remote_control_enabled = config.features.enabled(Feature::RemoteControl); + if transport_accept_handles.is_empty() && !remote_control_enabled { return Err(std::io::Error::new( ErrorKind::InvalidInput, "no transport configured; use --listen or enable remote control", )); } + let (remote_control_accept_handle, remote_control_handle) = start_remote_control( + config.chatgpt_base_url.clone(), + state_db.clone(), + auth_manager.clone(), + transport_event_tx.clone(), + transport_shutdown_token.clone(), + app_server_client_name_rx, + remote_control_enabled, + ) + .await?; + transport_accept_handles.push(remote_control_accept_handle); + let outbound_handle = tokio::spawn(async move { let mut outbound_connections = HashMap::::new(); loop { @@ -659,6 +660,7 @@ pub async fn run_main_with_transport( session_source, auth_manager, rpc_transport: analytics_rpc_transport(transport), + remote_control_handle: Some(remote_control_handle), }); let mut thread_created_rx = processor.thread_created_receiver(); let mut running_turn_count_rx = processor.subscribe_running_assistant_turn_count(); diff --git a/codex-rs/app-server/src/message_processor.rs b/codex-rs/app-server/src/message_processor.rs index 15df5a2a5..fbc1bd6d1 100644 --- a/codex-rs/app-server/src/message_processor.rs +++ b/codex-rs/app-server/src/message_processor.rs @@ -19,6 +19,7 @@ use crate::outgoing_message::ConnectionRequestId; use crate::outgoing_message::OutgoingMessageSender; use crate::outgoing_message::RequestContext; use crate::transport::AppServerTransport; +use crate::transport::RemoteControlHandle; use async_trait::async_trait; use codex_analytics::AnalyticsEventsClient; use codex_analytics::AppServerRpcTransport; @@ -170,6 +171,7 @@ pub(crate) struct MessageProcessor { config: Arc, config_warnings: Arc>, rpc_transport: AppServerRpcTransport, + remote_control_handle: Option, } #[derive(Clone, Debug, Default)] @@ -195,6 +197,7 @@ pub(crate) struct MessageProcessorArgs { pub(crate) session_source: SessionSource, pub(crate) auth_manager: Arc, pub(crate) rpc_transport: AppServerRpcTransport, + pub(crate) remote_control_handle: Option, } impl MessageProcessor { @@ -215,6 +218,7 @@ impl MessageProcessor { session_source, auth_manager, rpc_transport, + remote_control_handle, } = args; auth_manager.set_external_auth(Arc::new(ExternalAuthRefreshBridge { outgoing: outgoing.clone(), @@ -285,6 +289,7 @@ impl MessageProcessor { config, config_warnings: Arc::new(config_warnings), rpc_transport, + remote_control_handle, } } @@ -969,13 +974,36 @@ impl MessageProcessor { ) { match result { Ok(response) => { - self.codex_message_processor.handle_config_mutation(); + self.handle_config_mutation().await; self.outgoing.send_response(request_id, response).await; } Err(error) => self.outgoing.send_error(request_id, error).await, } } + async fn handle_config_mutation(&self) { + self.codex_message_processor.handle_config_mutation(); + let Some(remote_control_handle) = &self.remote_control_handle else { + return; + }; + + match self + .config_api + .load_latest_config(/*fallback_cwd*/ None) + .await + { + Ok(config) => { + remote_control_handle.set_enabled(config.features.enabled(Feature::RemoteControl)); + } + Err(error) => { + tracing::warn!( + "failed to load config for remote control enablement refresh after config mutation: {}", + error.message + ); + } + } + } + async fn handle_config_requirements_read(&self, request_id: ConnectionRequestId) { match self.config_api.config_requirements_read().await { Ok(response) => self.outgoing.send_response(request_id, response).await, diff --git a/codex-rs/app-server/src/message_processor/tracing_tests.rs b/codex-rs/app-server/src/message_processor/tracing_tests.rs index 2e8781606..ef88364a3 100644 --- a/codex-rs/app-server/src/message_processor/tracing_tests.rs +++ b/codex-rs/app-server/src/message_processor/tracing_tests.rs @@ -251,6 +251,7 @@ fn build_test_processor( session_source: SessionSource::VSCode, auth_manager, rpc_transport: AppServerRpcTransport::Stdio, + remote_control_handle: None, }); (processor, outgoing_rx) } diff --git a/codex-rs/app-server/src/transport/mod.rs b/codex-rs/app-server/src/transport/mod.rs index 7e1512a79..92383cb78 100644 --- a/codex-rs/app-server/src/transport/mod.rs +++ b/codex-rs/app-server/src/transport/mod.rs @@ -33,6 +33,7 @@ mod remote_control; mod stdio; mod websocket; +pub(crate) use remote_control::RemoteControlHandle; pub(crate) use remote_control::start_remote_control; pub(crate) use stdio::start_stdio_connection; pub(crate) use websocket::start_websocket_acceptor; diff --git a/codex-rs/app-server/src/transport/remote_control/mod.rs b/codex-rs/app-server/src/transport/remote_control/mod.rs index 6d9d65e8a..1ea89bb64 100644 --- a/codex-rs/app-server/src/transport/remote_control/mod.rs +++ b/codex-rs/app-server/src/transport/remote_control/mod.rs @@ -19,6 +19,7 @@ use std::io; use std::sync::Arc; use tokio::sync::mpsc; use tokio::sync::oneshot; +use tokio::sync::watch; use tokio::task::JoinHandle; use tokio_util::sync::CancellationToken; @@ -29,6 +30,21 @@ pub(super) struct QueuedServerEnvelope { pub(super) write_complete_tx: Option>, } +#[derive(Clone)] +pub(crate) struct RemoteControlHandle { + enabled_tx: Arc>, +} + +impl RemoteControlHandle { + pub(crate) fn set_enabled(&self, enabled: bool) { + self.enabled_tx.send_if_modified(|state| { + let changed = *state != enabled; + *state = enabled; + changed + }); + } +} + pub(crate) async fn start_remote_control( remote_control_url: String, state_db: Option>, @@ -36,21 +52,38 @@ pub(crate) async fn start_remote_control( transport_event_tx: mpsc::Sender, shutdown_token: CancellationToken, app_server_client_name_rx: Option>, -) -> io::Result> { - let remote_control_target = normalize_remote_control_url(&remote_control_url)?; - validate_remote_control_auth(&auth_manager).await?; + initial_enabled: bool, +) -> io::Result<(JoinHandle<()>, RemoteControlHandle)> { + let remote_control_target = if initial_enabled { + Some(normalize_remote_control_url(&remote_control_url)?) + } else { + None + }; + if initial_enabled { + validate_remote_control_auth(&auth_manager).await?; + } - Ok(tokio::spawn(async move { + let (enabled_tx, enabled_rx) = watch::channel(initial_enabled); + let join_handle = tokio::spawn(async move { RemoteControlWebsocket::new( + remote_control_url, remote_control_target, state_db, auth_manager, transport_event_tx, shutdown_token, + enabled_rx, ) .run(app_server_client_name_rx) .await; - })) + }); + + Ok(( + join_handle, + RemoteControlHandle { + enabled_tx: Arc::new(enabled_tx), + }, + )) } pub(crate) async fn validate_remote_control_auth( diff --git a/codex-rs/app-server/src/transport/remote_control/tests.rs b/codex-rs/app-server/src/transport/remote_control/tests.rs index b403fac48..9d430ccfb 100644 --- a/codex-rs/app-server/src/transport/remote_control/tests.rs +++ b/codex-rs/app-server/src/transport/remote_control/tests.rs @@ -125,13 +125,14 @@ async fn remote_control_transport_manages_virtual_clients_and_routes_messages() let (transport_event_tx, mut transport_event_rx) = mpsc::channel::(CHANNEL_CAPACITY); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(remote_control_state_runtime(&codex_home).await), remote_control_auth_manager(), transport_event_tx, shutdown_token.clone(), /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, ) .await .expect("remote control should start"); @@ -376,7 +377,7 @@ async fn remote_control_transport_manages_virtual_clients_and_routes_messages() ); shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; } #[tokio::test] @@ -389,13 +390,14 @@ async fn remote_control_transport_reconnects_after_disconnect() { let (transport_event_tx, mut transport_event_rx) = mpsc::channel::(CHANNEL_CAPACITY); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(remote_control_state_runtime(&codex_home).await), remote_control_auth_manager(), transport_event_tx, shutdown_token.clone(), /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, ) .await .expect("remote control should start"); @@ -452,7 +454,84 @@ async fn remote_control_transport_reconnects_after_disconnect() { } shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; +} + +#[tokio::test] +async fn remote_control_start_allows_remote_control_invalid_url_when_disabled() { + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, _remote_handle) = start_remote_control( + "https://internal.example.com/backend-api/".to_string(), + /*state_db*/ None, + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + /*initial_enabled*/ false, + ) + .await + .expect("disabled remote control should not validate the URL at startup"); + + shutdown_token.cancel(); + timeout(Duration::from_secs(1), remote_task) + .await + .expect("remote control task should stop") + .expect("remote control task should join"); +} + +#[tokio::test] +async fn remote_control_handle_set_enabled_stops_and_restarts_connections() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let remote_control_url = remote_control_url_for_listener(&listener); + let codex_home = TempDir::new().expect("temp dir should create"); + let (transport_event_tx, _transport_event_rx) = + mpsc::channel::(CHANNEL_CAPACITY); + let shutdown_token = CancellationToken::new(); + let (remote_task, remote_handle) = start_remote_control( + remote_control_url, + Some(remote_control_state_runtime(&codex_home).await), + remote_control_auth_manager(), + transport_event_tx, + shutdown_token.clone(), + /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, + ) + .await + .expect("remote control should start"); + + let enroll_request = accept_http_request(&listener).await; + assert_eq!( + enroll_request.request_line, + "POST /backend-api/wham/remote/control/server/enroll HTTP/1.1" + ); + respond_with_json( + enroll_request.stream, + json!({ "server_id": "srv_e_test", "environment_id": "env_test" }), + ) + .await; + let mut first_websocket = accept_remote_control_connection(&listener).await; + + remote_handle.set_enabled(/*enabled*/ false); + timeout(Duration::from_secs(1), first_websocket.next()) + .await + .expect("disabling remote control should close the websocket"); + timeout(Duration::from_millis(100), listener.accept()) + .await + .expect_err("disabled remote control should not reconnect"); + + remote_handle.set_enabled(/*enabled*/ true); + let mut second_websocket = accept_remote_control_connection(&listener).await; + second_websocket + .close(None) + .await + .expect("second websocket should close"); + + shutdown_token.cancel(); + let _ = remote_task.await; } #[tokio::test] @@ -465,13 +544,14 @@ async fn remote_control_transport_clears_outgoing_buffer_when_backend_acks() { let (transport_event_tx, mut transport_event_rx) = mpsc::channel::(CHANNEL_CAPACITY); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(remote_control_state_runtime(&codex_home).await), remote_control_auth_manager(), transport_event_tx, shutdown_token.clone(), /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, ) .await .expect("remote control should start"); @@ -617,7 +697,7 @@ async fn remote_control_transport_clears_outgoing_buffer_when_backend_acks() { ); shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; } #[tokio::test] @@ -631,13 +711,14 @@ async fn remote_control_http_mode_enrolls_before_connecting() { mpsc::channel::(CHANNEL_CAPACITY); let expected_server_name = gethostname().to_string_lossy().trim().to_string(); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(remote_control_state_runtime(&codex_home).await), remote_control_auth_manager(), transport_event_tx, shutdown_token.clone(), /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, ) .await .expect("remote control should start"); @@ -815,7 +896,7 @@ async fn remote_control_http_mode_enrolls_before_connecting() { ); shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; } #[tokio::test] @@ -847,13 +928,14 @@ async fn remote_control_http_mode_reuses_persisted_enrollment_before_reenrolling let (transport_event_tx, _transport_event_rx) = mpsc::channel::(CHANNEL_CAPACITY); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(state_db.clone()), remote_control_auth_manager_with_home(&codex_home), transport_event_tx, shutdown_token.clone(), /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, ) .await .expect("remote control should start"); @@ -879,7 +961,7 @@ async fn remote_control_http_mode_reuses_persisted_enrollment_before_reenrolling ); shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; } #[tokio::test] @@ -913,13 +995,14 @@ async fn remote_control_stdio_mode_waits_for_client_name_before_connecting() { mpsc::channel::(CHANNEL_CAPACITY); let (app_server_client_name_tx, app_server_client_name_rx) = oneshot::channel::(); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(state_db.clone()), remote_control_auth_manager_with_home(&codex_home), transport_event_tx, shutdown_token.clone(), Some(app_server_client_name_rx), + /*initial_enabled*/ true, ) .await .expect("remote control should start"); @@ -936,7 +1019,7 @@ async fn remote_control_stdio_mode_waits_for_client_name_before_connecting() { ); shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; } #[tokio::test] @@ -969,13 +1052,14 @@ async fn remote_control_waits_for_account_id_before_enrolling() { let (transport_event_tx, _transport_event_rx) = mpsc::channel::(CHANNEL_CAPACITY); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(state_db.clone()), auth_manager, transport_event_tx, shutdown_token.clone(), /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, ) .await .expect("remote control should start before account id is available"); @@ -1012,7 +1096,7 @@ async fn remote_control_waits_for_account_id_before_enrolling() { ); shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; } #[tokio::test] @@ -1051,13 +1135,14 @@ async fn remote_control_http_mode_clears_stale_persisted_enrollment_after_404() let (transport_event_tx, _transport_event_rx) = mpsc::channel::(CHANNEL_CAPACITY); let shutdown_token = CancellationToken::new(); - let remote_handle = start_remote_control( + let (remote_task, _remote_handle) = start_remote_control( remote_control_url, Some(state_db.clone()), remote_control_auth_manager_with_home(&codex_home), transport_event_tx, shutdown_token.clone(), /*app_server_client_name_rx*/ None, + /*initial_enabled*/ true, ) .await .expect("remote control should start"); @@ -1104,7 +1189,7 @@ async fn remote_control_http_mode_clears_stale_persisted_enrollment_after_404() ); shutdown_token.cancel(); - let _ = remote_handle.await; + let _ = remote_task.await; } #[derive(Debug)] diff --git a/codex-rs/app-server/src/transport/remote_control/websocket.rs b/codex-rs/app-server/src/transport/remote_control/websocket.rs index 3f42a7bbb..a0387ef6c 100644 --- a/codex-rs/app-server/src/transport/remote_control/websocket.rs +++ b/codex-rs/app-server/src/transport/remote_control/websocket.rs @@ -111,7 +111,8 @@ struct WebsocketState { } pub(crate) struct RemoteControlWebsocket { - remote_control_target: RemoteControlTarget, + remote_control_url: String, + remote_control_target: Option, state_db: Option>, auth_manager: Arc, shutdown_token: CancellationToken, @@ -122,15 +123,24 @@ pub(crate) struct RemoteControlWebsocket { state: Arc>, server_event_rx: Arc>>, used_rx: watch::Receiver, + enabled_rx: watch::Receiver, +} + +enum ConnectOutcome { + Connected(Box>>), + Disabled, + Shutdown, } impl RemoteControlWebsocket { pub(crate) fn new( - remote_control_target: RemoteControlTarget, + remote_control_url: String, + remote_control_target: Option, state_db: Option>, auth_manager: Arc, transport_event_tx: mpsc::Sender, shutdown_token: CancellationToken, + enabled_rx: watch::Receiver, ) -> Self { let shutdown_token = shutdown_token.child_token(); let (server_event_tx, server_event_rx) = mpsc::channel(super::CHANNEL_CAPACITY); @@ -140,6 +150,7 @@ impl RemoteControlWebsocket { let auth_recovery = auth_manager.unauthorized_recovery(); Self { + remote_control_url, remote_control_target, state_db, auth_manager, @@ -155,6 +166,7 @@ impl RemoteControlWebsocket { })), server_event_rx: Arc::new(Mutex::new(server_event_rx)), used_rx, + enabled_rx, } } @@ -174,13 +186,18 @@ impl RemoteControlWebsocket { }; loop { + if !self.wait_until_enabled().await { + break; + } + let shutdown_token = self.shutdown_token.child_token(); let websocket_connection = match self .connect(&shutdown_token, app_server_client_name.as_deref()) .await { - Some(websocket_connection) => websocket_connection, - None => break, + ConnectOutcome::Connected(websocket_connection) => *websocket_connection, + ConnectOutcome::Disabled => continue, + ConnectOutcome::Shutdown => break, }; self.run_connection(websocket_connection, shutdown_token) @@ -208,50 +225,93 @@ impl RemoteControlWebsocket { } } + async fn wait_until_enabled(&mut self) -> bool { + tokio::select! { + _ = self.shutdown_token.cancelled() => false, + enabled = self.enabled_rx.wait_for(|enabled| *enabled) => enabled.is_ok(), + } + } + async fn connect( &mut self, shutdown_token: &CancellationToken, app_server_client_name: Option<&str>, - ) -> Option>> { + ) -> ConnectOutcome { + let remote_control_target = match self.remote_control_target.as_ref() { + Some(remote_control_target) => remote_control_target.clone(), + None => match super::protocol::normalize_remote_control_url(&self.remote_control_url) { + Ok(remote_control_target) => { + self.remote_control_target = Some(remote_control_target.clone()); + remote_control_target + } + Err(err) => { + warn!("remote control is enabled but the URL is invalid: {err}"); + tokio::select! { + _ = shutdown_token.cancelled() => return ConnectOutcome::Shutdown, + changed = self.enabled_rx.wait_for(|enabled| !*enabled) => { + if changed.is_err() { + return ConnectOutcome::Shutdown; + } + return ConnectOutcome::Disabled; + } + } + } + }, + }; + loop { let subscribe_cursor = self.state.lock().await.subscribe_cursor.clone(); - tokio::select! { - _ = shutdown_token.cancelled() => return None, + let connect_result = tokio::select! { + _ = shutdown_token.cancelled() => return ConnectOutcome::Shutdown, + changed = self.enabled_rx.wait_for(|enabled| !*enabled) => { + if changed.is_err() { + return ConnectOutcome::Shutdown; + } + return ConnectOutcome::Disabled; + } connect_result = connect_remote_control_websocket( - &self.remote_control_target, + &remote_control_target, self.state_db.as_deref(), &self.auth_manager, &mut self.auth_recovery, &mut self.enrollment, subscribe_cursor.as_deref(), app_server_client_name, - ) => { - match connect_result { - Ok((websocket_connection, response)) => { - self.reconnect_attempt = 0; - self.auth_recovery = self.auth_manager.unauthorized_recovery(); - info!( - "connected to app-server remote control websocket: {}, {}", - self.remote_control_target.websocket_url, - format_headers(response.headers()) - ); - return Some(websocket_connection); - } - Err(err) => { - let reconnect_delay = if err.kind() == ErrorKind::WouldBlock { - info!("{err}"); - REMOTE_CONTROL_ACCOUNT_ID_RETRY_INTERVAL - } else { - warn!("{err}"); - let reconnect_delay = backoff(self.reconnect_attempt); - self.reconnect_attempt += 1; - reconnect_delay - }; - tokio::select! { - _ = shutdown_token.cancelled() => return None, - _ = tokio::time::sleep(reconnect_delay) => {} + ) => connect_result, + }; + + match connect_result { + Ok((websocket_connection, response)) => { + self.reconnect_attempt = 0; + self.auth_recovery = self.auth_manager.unauthorized_recovery(); + info!( + "connected to app-server remote control websocket: {}, {}", + remote_control_target.websocket_url, + format_headers(response.headers()) + ); + return ConnectOutcome::Connected(Box::new(websocket_connection)); + } + Err(err) => { + let reconnect_delay = if err.kind() == ErrorKind::WouldBlock { + REMOTE_CONTROL_ACCOUNT_ID_RETRY_INTERVAL + } else { + warn!( + "failed to connect to app-server remote control websocket: {}, err: {}", + remote_control_target.websocket_url, err + ); + let reconnect_delay = backoff(self.reconnect_attempt); + self.reconnect_attempt += 1; + reconnect_delay + }; + tokio::select! { + _ = shutdown_token.cancelled() => return ConnectOutcome::Shutdown, + changed = self.enabled_rx.wait_for(|enabled| !*enabled) => { + if changed.is_err() { + return ConnectOutcome::Shutdown; } + return ConnectOutcome::Disabled; } + _ = tokio::time::sleep(reconnect_delay) => {} } } } @@ -282,8 +342,10 @@ impl RemoteControlWebsocket { shutdown_token.clone(), )); + let mut enabled_rx = self.enabled_rx.clone(); tokio::select! { _ = shutdown_token.cancelled() => {} + _ = enabled_rx.wait_for(|enabled| !*enabled) => shutdown_token.cancel(), _ = join_set.join_next() => shutdown_token.cancel(), } @@ -1124,15 +1186,18 @@ mod tests { let (transport_event_tx, transport_event_rx) = mpsc::channel(1); drop(transport_event_rx); let shutdown_token = CancellationToken::new(); + let (_enabled_tx, enabled_rx) = watch::channel(true); let websocket_task = tokio::spawn({ let shutdown_token = shutdown_token.clone(); async move { RemoteControlWebsocket::new( - remote_control_target, + remote_control_url, + Some(remote_control_target), /*state_db*/ None, remote_control_auth_manager(), transport_event_tx, shutdown_token, + enabled_rx, ) .run(/*app_server_client_name_rx*/ None) .await