diff --git a/crates/targets/src/target/mqtt.rs b/crates/targets/src/target/mqtt.rs index 2b25f6a63..8a8560261 100644 --- a/crates/targets/src/target/mqtt.rs +++ b/crates/targets/src/target/mqtt.rs @@ -58,6 +58,13 @@ const DEFAULT_MQTT_TCP_PORT: u16 = 1883; const DEFAULT_MQTT_TLS_PORT: u16 = 8883; const DEFAULT_MQTT_WSS_PORT: u16 = 443; const MAX_MQTT_PACKET_SIZE_BYTES: u32 = 100 * 1024 * 1024; +/// Minimum delay before the supervisor rebuilds the client and event loop +/// after a session exits. Also the delay used right after a session that had +/// successfully connected, so a transient drop reconnects promptly. +const MQTT_RECONNECT_BACKOFF_MIN: Duration = Duration::from_secs(1); +/// Upper bound for the exponential reconnect backoff, so repeated fatal +/// failures never turn into a tight reconnect storm. +const MQTT_RECONNECT_BACKOFF_MAX: Duration = Duration::from_secs(30); const DEFAULT_MQTT_WS_PATH_ALLOWLIST: &[&str] = &["/", "/mqtt"]; const LOG_COMPONENT_TARGETS: &str = "targets"; const LOG_SUBSYSTEM_MQTT: &str = "mqtt"; @@ -633,24 +640,6 @@ where "mqtt target state" ); - // Use the latest MqttOptions (may have been updated by TLS reload coordinator). - let mqtt_options: MqttOptions = (**pending_mqtt_options.load()).clone(); - - let (new_client, eventloop) = AsyncClient::builder(mqtt_options).capacity(10).build(); - - if let Err(e) = new_client.subscribe(&args_clone.topic, args_clone.qos).await { - error!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id_clone, - state = "subscribe_failed", - error = %e, - "mqtt target state" - ); - return Err(TargetError::Network(format!("MQTT subscribe failed: {e}"))); - } - let mut rx_guard = bg_task_manager.initial_cancel_rx.lock().await; let cancel_rx = rx_guard.take().ok_or_else(|| { error!( @@ -665,18 +654,27 @@ where })?; drop(rx_guard); - *client_arc.lock().await = Some(new_client.clone()); - info!( event = EVENT_MQTT_TARGET_STATE, component = LOG_COMPONENT_TARGETS, subsystem = LOG_SUBSYSTEM_MQTT, target_id = %target_id_clone, - state = "event_loop_spawning", + state = "supervisor_spawning", "mqtt target state" ); - let task_handle = - tokio::spawn(run_mqtt_event_loop(eventloop, connected_arc.clone(), target_id_clone.clone(), cancel_rx)); + // Spawn a supervisor that owns the reconnect loop. Building the + // client/event loop, subscribing, and publishing the client to + // `client_arc` all happen per session inside the supervisor, so a + // fatal protocol error that ends one session is followed by a + // backoff and a fresh session instead of permanent silence. + let task_handle = tokio::spawn(supervise_mqtt_event_loop( + pending_mqtt_options, + args_clone, + client_arc, + connected_arc, + target_id_clone, + cancel_rx, + )); Ok(task_handle) }) .await @@ -885,12 +883,138 @@ where } } -async fn run_mqtt_event_loop( - mut eventloop: EventLoop, +/// Computes the next reconnect backoff by doubling the current delay, capped at +/// [`MQTT_RECONNECT_BACKOFF_MAX`]. Kept as a pure function so the backoff policy +/// can be unit tested without a live broker. +fn next_reconnect_backoff(current: Duration) -> Duration { + current.saturating_mul(2).min(MQTT_RECONNECT_BACKOFF_MAX) +} + +/// Drives the supervised reconnect loop: run a session, then wait a backoff +/// before restarting, until a cancellation signal arrives. Cancellation drops +/// the in-flight session future (the outer `select!`), so `close()` stops the +/// loop promptly without the session needing its own cancel channel. +/// +/// `run_session` returns whether its session connected at least once; a +/// connected session resets the backoff so a transient drop reconnects quickly, +/// while repeated immediate failures back off exponentially. +async fn reconnect_supervisor(mut cancel_rx: mpsc::Receiver<()>, mut run_session: F) +where + F: FnMut() -> Fut, + Fut: std::future::Future, +{ + let mut backoff = MQTT_RECONNECT_BACKOFF_MIN; + loop { + let connected = tokio::select! { + biased; + _ = cancel_rx.recv() => break, + connected = run_session() => connected, + }; + + if connected { + backoff = MQTT_RECONNECT_BACKOFF_MIN; + } + + tokio::select! { + biased; + _ = cancel_rx.recv() => break, + _ = tokio::time::sleep(backoff) => {} + } + + backoff = next_reconnect_backoff(backoff); + } +} + +/// Supervises the MQTT event loop for the lifetime of the target. Each session +/// rebuilds the client and event loop from the latest `MqttOptions`, so TLS +/// reloads are picked up on reconnect, and a session that exits (including on a +/// fatal protocol error) is restarted after a backoff instead of leaving the +/// target permanently wedged. +async fn supervise_mqtt_event_loop( + pending_mqtt_options: Arc>, + args: MQTTArgs, + client_arc: Arc>>, connected_status: Arc, target_id: TargetID, - mut cancel_rx: mpsc::Receiver<()>, + cancel_rx: mpsc::Receiver<()>, ) { + info!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "supervisor_started", + "mqtt target state" + ); + + reconnect_supervisor(cancel_rx, || { + let pending_mqtt_options = Arc::clone(&pending_mqtt_options); + let args = args.clone(); + let client_arc = Arc::clone(&client_arc); + let connected_status = Arc::clone(&connected_status); + let target_id = target_id.clone(); + async move { run_one_mqtt_session(pending_mqtt_options, args, client_arc, connected_status, target_id).await } + }) + .await; + + connected_status.store(false, Ordering::SeqCst); + info!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "supervisor_stopped", + "mqtt target state" + ); +} + +/// Builds a client and event loop, subscribes, publishes the client for +/// `send_body`, then runs the event loop until it exits. Returns whether the +/// session established a connection at least once. +async fn run_one_mqtt_session( + pending_mqtt_options: Arc>, + args: MQTTArgs, + client_arc: Arc>>, + connected_status: Arc, + target_id: TargetID, +) -> bool { + // Use the latest MqttOptions (may have been updated by TLS reload coordinator). + let mqtt_options: MqttOptions = (**pending_mqtt_options.load()).clone(); + let (new_client, eventloop) = AsyncClient::builder(mqtt_options).capacity(10).build(); + + if let Err(e) = new_client.subscribe(&args.topic, args.qos).await { + error!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "subscribe_failed", + error = %e, + "mqtt target state" + ); + return false; + } + + *client_arc.lock().await = Some(new_client); + connected_status.store(false, Ordering::SeqCst); + + info!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "event_loop_spawning", + "mqtt target state" + ); + + run_mqtt_event_loop(eventloop, connected_status, target_id).await +} + +/// Runs a single MQTT event-loop session until it exits (fatal protocol error +/// or `RequestsDone`). Returns whether the session connected at least once. +/// Cancellation is handled by the supervisor dropping this future, so no cancel +/// channel is needed here. +async fn run_mqtt_event_loop(mut eventloop: EventLoop, connected_status: Arc, target_id: TargetID) -> bool { info!( event = EVENT_MQTT_TARGET_STATE, component = LOG_COMPONENT_TARGETS, @@ -902,171 +1026,154 @@ async fn run_mqtt_event_loop( let mut initial_connection_established = false; loop { - tokio::select! { - biased; - _ = cancel_rx.recv() => { - info!( + let polled_event_result = if !initial_connection_established || !connected_status.load(Ordering::SeqCst) { + match tokio::time::timeout(EVENT_LOOP_POLL_TIMEOUT, eventloop.poll()).await { + Ok(result) => Some(result), + Err(_) => { + debug!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "poll_timeout", + "mqtt target state" + ); + connected_status.store(false, Ordering::SeqCst); + None + } + } + } else { + Some(eventloop.poll().await) + }; + + match polled_event_result { + Some(Ok(notification)) => { + trace!(target_id = %target_id, event = ?notification, "Received MQTT event"); + match notification { + rumqttc::Event::Incoming(Incoming::ConnAck(_conn_ack)) => { + info!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "connack_received", + "mqtt target state" + ); + connected_status.store(true, Ordering::SeqCst); + initial_connection_established = true; + } + rumqttc::Event::Incoming(Incoming::Publish(publish)) => { + debug!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "publish_received", + topic = ?publish.topic, + payload_len = publish.payload.len(), + "mqtt target state" + ); + } + rumqttc::Event::Incoming(Incoming::Disconnect(_)) => { + info!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "broker_disconnected", + "mqtt target state" + ); + connected_status.store(false, Ordering::SeqCst); + } + rumqttc::Event::Incoming(Incoming::PingResp(_)) => { + trace!(target_id = %target_id, "Received PingResp from broker. Connection is alive."); + } + rumqttc::Event::Incoming(Incoming::SubAck(suback)) => { + trace!(target_id = %target_id, "Received SubAck for pkid: {}", suback.pkid); + } + rumqttc::Event::Incoming(Incoming::PubAck(puback)) => { + trace!(target_id = %target_id, "Received PubAck for pkid: {}", puback.pkid); + } + // Process other incoming packet types as needed (PubRec, PubRel, PubComp, UnsubAck) + rumqttc::Event::Outgoing(Outgoing::Disconnect) => { + info!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "client_disconnect_requested", + "mqtt target state" + ); + connected_status.store(false, Ordering::SeqCst); + } + rumqttc::Event::Outgoing(Outgoing::PingReq) => { + trace!(target_id = %target_id, "Client sent PingReq to broker."); + } + // Other Outgoing events (Subscribe, Unsubscribe, Publish) usually do not need to handle connection status here, + // Because they are actions initiated by the client. + _ => { + // Log other unspecified MQTT events that are not handled, which helps debug + trace!(target_id = %target_id, "Unhandled or generic MQTT event: {:?}", notification); + } + } + } + Some(Err(e)) => { + connected_status.store(false, Ordering::SeqCst); + error!( event = EVENT_MQTT_TARGET_STATE, component = LOG_COMPONENT_TARGETS, subsystem = LOG_SUBSYSTEM_MQTT, target_id = %target_id, - state = "cancellation_received", + state = "poll_failed", + error = %e, "mqtt target state" ); - break; - } - polled_event_result = async { - if !initial_connection_established || !connected_status.load(Ordering::SeqCst) { - match tokio::time::timeout(EVENT_LOOP_POLL_TIMEOUT, eventloop.poll()).await { - Ok(result) => Some(result), - Err(_) => { - debug!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "poll_timeout", - "mqtt target state" - ); - connected_status.store(false, Ordering::SeqCst); - None - } - } - } else { - Some(eventloop.poll().await) - } - } => { - match polled_event_result { - Some(Ok(notification)) => { - trace!(target_id = %target_id, event = ?notification, "Received MQTT event"); - match notification { - rumqttc::Event::Incoming(Incoming::ConnAck(_conn_ack)) => { - info!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "connack_received", - "mqtt target state" - ); - connected_status.store(true, Ordering::SeqCst); - initial_connection_established = true; - } - rumqttc::Event::Incoming(Incoming::Publish(publish)) => { - debug!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "publish_received", - topic = ?publish.topic, - payload_len = publish.payload.len(), - "mqtt target state" - ); - } - rumqttc::Event::Incoming(Incoming::Disconnect(_)) => { - info!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "broker_disconnected", - "mqtt target state" - ); - connected_status.store(false, Ordering::SeqCst); - } - rumqttc::Event::Incoming(Incoming::PingResp(_)) => { - trace!(target_id = %target_id, "Received PingResp from broker. Connection is alive."); - } - rumqttc::Event::Incoming(Incoming::SubAck(suback)) => { - trace!(target_id = %target_id, "Received SubAck for pkid: {}", suback.pkid); - } - rumqttc::Event::Incoming(Incoming::PubAck(puback)) => { - trace!(target_id = %target_id, "Received PubAck for pkid: {}", puback.pkid); - } - // Process other incoming packet types as needed (PubRec, PubRel, PubComp, UnsubAck) - rumqttc::Event::Outgoing(Outgoing::Disconnect) => { - info!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "client_disconnect_requested", - "mqtt target state" - ); - connected_status.store(false, Ordering::SeqCst); - } - rumqttc::Event::Outgoing(Outgoing::PingReq) => { - trace!(target_id = %target_id, "Client sent PingReq to broker."); - } - // Other Outgoing events (Subscribe, Unsubscribe, Publish) usually do not need to handle connection status here, - // Because they are actions initiated by the client. - _ => { - // Log other unspecified MQTT events that are not handled, which helps debug - trace!(target_id = %target_id, "Unhandled or generic MQTT event: {:?}", notification); - } - } - } - Some(Err(e)) => { - connected_status.store(false, Ordering::SeqCst); - error!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "poll_failed", - error = %e, - "mqtt target state" - ); - if matches!(e, - ConnectionError::Io(_) | - ConnectionError::Timeout(_) | - ConnectionError::ConnectionRefused(_) | - ConnectionError::Tls(_) - ) { - warn!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "reconnect_pending", - error = %e, - "mqtt target state" - ); - } - // Here you can decide whether to break loops based on the error type. - // For example, for some unrecoverable errors. - if is_fatal_mqtt_error(&e) { - error!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "fatal_error", - error = %e, - "mqtt target state" - ); - break; - } - // rumqttc's eventloop.poll() may return Err and terminate after some errors, - // Or it will handle reconnection internally. To continue here will make select! wait again. - // If the error is temporary and rumqttc is handling reconnection, poll() should eventually succeed or return a different error again. - // Sleep briefly to avoid busy cycles in case of rapid failure. - tokio::time::sleep(Duration::from_secs(1)).await; - } - None => { - warn!( - event = EVENT_MQTT_TARGET_STATE, - component = LOG_COMPONENT_TARGETS, - subsystem = LOG_SUBSYSTEM_MQTT, - target_id = %target_id, - state = "poll_retry_scheduled", - "mqtt target state" - ); - continue; - } + if matches!( + e, + ConnectionError::Io(_) + | ConnectionError::Timeout(_) + | ConnectionError::ConnectionRefused(_) + | ConnectionError::Tls(_) + ) { + warn!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "reconnect_pending", + error = %e, + "mqtt target state" + ); } + // Fatal protocol errors end this session; the supervisor rebuilds + // the client and event loop after a backoff. Non-fatal errors are + // usually handled by rumqttc's internal reconnection, so keep + // polling after a short pause to avoid a busy loop on rapid failure. + if is_fatal_mqtt_error(&e) { + error!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "fatal_error", + error = %e, + "mqtt target state" + ); + break; + } + tokio::time::sleep(Duration::from_secs(1)).await; + } + None => { + warn!( + event = EVENT_MQTT_TARGET_STATE, + component = LOG_COMPONENT_TARGETS, + subsystem = LOG_SUBSYSTEM_MQTT, + target_id = %target_id, + state = "poll_retry_scheduled", + "mqtt target state" + ); + continue; } } } @@ -1079,6 +1186,8 @@ async fn run_mqtt_event_loop( state = "event_loop_finished", "mqtt target state" ); + + initial_connection_established } /// Check whether the given MQTT connection error should be considered a fatal error, @@ -1429,21 +1538,20 @@ where ); } - // Wait for the task to finish if it was initialized - if let Some(_task_handle) = self.bg_task_manager.init_cell.get() { + // The cancel signal above makes the supervisor's `select!` drop the + // in-flight session and stop the reconnect loop. The `JoinHandle` lives + // in a `OnceCell` shared across `clone_target()` clones, so it cannot be + // taken out to be joined here; we rely on the cancel signal for a prompt, + // graceful stop. + if self.bg_task_manager.init_cell.get().is_some() { debug!( event = EVENT_MQTT_TARGET_STATE, component = LOG_COMPONENT_TARGETS, subsystem = LOG_SUBSYSTEM_MQTT, target_id = %self.id, - state = "waiting_for_event_loop", + state = "supervisor_stop_signalled", "mqtt target state" ); - // It's tricky to await here if close is called from a sync context or Drop - // For async close, this is fine. Consider a timeout. - // let _ = tokio::time::timeout(Duration::from_secs(5), task_handle.await).await; - // If task_handle.await is directly used, ensure it's not awaited multiple times if close can be called multiple times. - // For now, we rely on the signal and the task's self-termination. } if let Some(client_instance) = self.client.lock().await.take() { @@ -1527,11 +1635,91 @@ where #[cfg(test)] mod tests { - use super::{MQTTArgs, MQTTTlsConfig, QoS, validate_mqtt_broker_url}; + use super::{ + MQTT_RECONNECT_BACKOFF_MAX, MQTT_RECONNECT_BACKOFF_MIN, MQTTArgs, MQTTTlsConfig, QoS, next_reconnect_backoff, + reconnect_supervisor, validate_mqtt_broker_url, + }; use crate::target::{REDACTED_SECRET, TargetType}; + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; use std::time::Duration; + use tokio::sync::mpsc; use url::Url; + #[test] + fn next_reconnect_backoff_doubles_until_capped() { + let mut backoff = MQTT_RECONNECT_BACKOFF_MIN; + // Doubles on each step. + backoff = next_reconnect_backoff(backoff); + assert_eq!(backoff, MQTT_RECONNECT_BACKOFF_MIN * 2); + backoff = next_reconnect_backoff(backoff); + assert_eq!(backoff, MQTT_RECONNECT_BACKOFF_MIN * 4); + + // Never exceeds the cap, even from a huge starting point. + assert_eq!(next_reconnect_backoff(MQTT_RECONNECT_BACKOFF_MAX), MQTT_RECONNECT_BACKOFF_MAX); + assert_eq!(next_reconnect_backoff(Duration::from_secs(3600)), MQTT_RECONNECT_BACKOFF_MAX); + } + + #[tokio::test(start_paused = true)] + async fn supervisor_restarts_session_until_cancelled() { + // A session that exits immediately (as after a fatal protocol error) + // must be restarted by the supervisor rather than leaving the target + // permanently silent. Time is paused so the reconnect backoff advances + // automatically without real waits. + let (cancel_tx, cancel_rx) = mpsc::channel(1); + let (attempt_tx, mut attempt_rx) = mpsc::unbounded_channel(); + let attempts = Arc::new(AtomicUsize::new(0)); + + let attempts_in_task = Arc::clone(&attempts); + let handle = tokio::spawn(reconnect_supervisor(cancel_rx, move || { + let attempt_tx = attempt_tx.clone(); + let attempts_in_task = Arc::clone(&attempts_in_task); + async move { + attempts_in_task.fetch_add(1, Ordering::SeqCst); + let _ = attempt_tx.send(()); + // Session exits immediately and never connected. + false + } + })); + + // Observe several automatic restarts driven purely by the supervisor. + for _ in 0..4 { + attempt_rx.recv().await.expect("supervisor should restart the session"); + } + + cancel_tx.send(()).await.expect("cancel signal should be delivered"); + handle.await.expect("supervisor task should stop cleanly"); + + assert!(attempts.load(Ordering::SeqCst) >= 4, "session should have been restarted repeatedly"); + } + + #[tokio::test(start_paused = true)] + async fn supervisor_stops_promptly_on_cancel() { + // A connected session that stays up must be torn down by cancellation + // (the supervisor drops the in-flight session future). + let (cancel_tx, cancel_rx) = mpsc::channel(1); + let started = Arc::new(AtomicUsize::new(0)); + + let started_in_task = Arc::clone(&started); + let handle = tokio::spawn(reconnect_supervisor(cancel_rx, move || { + let started_in_task = Arc::clone(&started_in_task); + async move { + started_in_task.fetch_add(1, Ordering::SeqCst); + // Long-lived, "connected" session that never returns on its own. + std::future::pending::().await + } + })); + + // Let the session start, then cancel; the supervisor must stop. + while started.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + cancel_tx.send(()).await.expect("cancel signal should be delivered"); + handle.await.expect("supervisor task should stop cleanly on cancel"); + + assert_eq!(started.load(Ordering::SeqCst), 1, "session should not be restarted after cancel"); + } + #[test] fn validate_mqtt_broker_url_rejects_non_websocket_path() { let url = Url::parse("mqtt://broker.example.com:1883/custom").expect("valid url");