#[cfg(feature = "metrics")]
pub mod metrics;
use async_trait::async_trait;
use derive_builder::Builder;
#[cfg(feature = "lwt")]
use jsonpath_rust::JsonPath;
#[cfg(feature = "lwt")]
use rumqttc::mqttbytes::v5::{LastWill, LastWillProperties};
use rumqttc::mqttbytes::v5::{Packet, PublishProperties, SubscribeProperties};
use rumqttc::mqttbytes::QoS;
use rumqttc::{AsyncClient, Broker, Event, EventLoop, MqttOptions};
#[cfg(feature = "use-rustls")]
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
#[cfg(feature = "use-rustls")]
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
#[cfg(feature = "use-rustls")]
use rustls::{DigitallySignedStruct, SignatureScheme};
use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
#[cfg(feature = "lwt")]
use stinger_mqtt_trait::availability_trait::AvailabilityHelper as AvailabilityHelperTrait;
#[cfg(all(feature = "lwt", test))]
use stinger_mqtt_trait::concrete::GenericAvailability;
use stinger_mqtt_trait::message::{MqttMessage, MqttMessageBuilder, QoS as StingerQoS};
use stinger_mqtt_trait::{Mqtt5PubSub, Mqtt5PubSubError, MqttConnectionState, MqttPublishSuccess};
use thiserror::Error;
use tokio::sync::{broadcast, mpsc, oneshot, watch, Mutex, RwLock};
use tracing::{debug, error, info, warn};
use uuid::Uuid;
#[derive(Error, Debug)]
pub enum MqttierError {
#[error("MQTT connection error: {0}")]
ConnectionError(#[from] rumqttc::ConnectionError),
#[error("MQTT client error: {0}")]
ClientError(#[from] rumqttc::ClientError),
#[error("Channel send error")]
ChannelSendError,
#[error("Invalid QoS value: {0}")]
InvalidQos(u8),
#[error("Subscribe error: {0}")]
SubscribeError(String),
}
type Result<T> = std::result::Result<T, MqttierError>;
type PublishCompletion = oneshot::Sender<std::result::Result<MqttPublishSuccess, Mqtt5PubSubError>>;
#[derive(Debug, Clone)]
struct QueuedSubscription {
topic: String,
qos: QoS,
props: SubscribeProperties,
}
#[derive(Debug)]
struct QueuedMessage {
message: MqttMessage,
completion: Option<PublishCompletion>,
#[cfg(feature = "metrics")]
start_timestamp: std::time::Instant,
}
struct PublishState {
publish_queue_rx: mpsc::Receiver<QueuedMessage>,
pub ack_timeout_ms: u64,
}
#[derive(Clone)]
pub struct TcpConnection {
pub hostname: String,
pub port: u16,
}
impl TcpConnection {
pub fn from_env_with_defaults(default_hostname: impl Into<String>, default_port: u16) -> Self {
let hostname = std::env::var("MQTT_HOSTNAME").unwrap_or_else(|_| default_hostname.into());
let port = std::env::var("MQTT_PORT")
.ok()
.and_then(|v| v.parse::<u16>().ok())
.unwrap_or(default_port);
Self { hostname, port }
}
}
#[cfg(feature = "use-rustls")]
#[derive(Clone)]
pub enum ServerCertificateVerification {
Insecure,
CaCertPath(String),
}
#[cfg(feature = "use-rustls")]
#[derive(Clone)]
pub struct TlsConnection {
pub hostname: String,
pub port: u16,
pub certificate_verification: ServerCertificateVerification,
pub client_cert_path: Option<String>,
pub client_key_path: Option<String>,
}
#[derive(Clone)]
pub enum Connection {
TcpLocalhost(u16), UnixSocket(String), Tcp(TcpConnection), #[cfg(feature = "use-rustls")]
TcpWithTls(TlsConnection), }
#[derive(Clone)]
pub struct Credentials {
pub username: String,
pub password: String,
}
#[derive(Clone, Builder)]
#[builder(setter(into))]
pub struct MqttierOptions {
#[builder(default = "Connection::TcpLocalhost(1883)")]
pub connection: Connection,
#[builder(default = "Uuid::new_v4().to_string()")]
pub client_id: String,
#[builder(default = "5000")]
pub ack_timeout_ms: u64,
#[builder(default = "60")]
pub keepalive_secs: u16,
#[builder(default = "1200")]
pub session_expiry_interval_secs: u16,
#[cfg(feature = "lwt")]
#[builder(default = "None")]
pub availability_helper: Option<Arc<dyn AvailabilityHelperTrait + Send + Sync>>,
#[builder(default = "10")]
pub lwt_delay_secs: u16, #[builder(default = "128")]
pub publish_queue_size: u16, #[builder(default = "32")]
pub max_inflight_messages: u16, #[builder(default = "(10 * 1024)")]
pub max_incoming_packet_size: u32,
#[builder(default = "None")]
pub credentials: Option<Credentials>,
}
#[cfg(feature = "use-rustls")]
#[derive(Debug)]
struct NoopCertVerifier;
#[cfg(feature = "use-rustls")]
impl ServerCertVerifier for NoopCertVerifier {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: UnixTime,
) -> std::result::Result<ServerCertVerified, rustls::Error> {
Ok(ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &DigitallySignedStruct,
) -> std::result::Result<HandshakeSignatureValid, rustls::Error> {
Ok(HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &DigitallySignedStruct,
) -> std::result::Result<HandshakeSignatureValid, rustls::Error> {
Ok(HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
vec![
SignatureScheme::RSA_PKCS1_SHA1,
SignatureScheme::ECDSA_SHA1_Legacy,
SignatureScheme::RSA_PKCS1_SHA256,
SignatureScheme::ECDSA_NISTP256_SHA256,
SignatureScheme::RSA_PKCS1_SHA384,
SignatureScheme::ECDSA_NISTP384_SHA384,
SignatureScheme::RSA_PKCS1_SHA512,
SignatureScheme::ECDSA_NISTP521_SHA512,
SignatureScheme::RSA_PSS_SHA256,
SignatureScheme::RSA_PSS_SHA384,
SignatureScheme::RSA_PSS_SHA512,
SignatureScheme::ED25519,
SignatureScheme::ED448,
]
}
}
#[derive(Clone)]
pub struct MqttierClient {
pub client_id: String,
client: AsyncClient,
next_subscription_id: Arc<AtomicUsize>,
is_running: Arc<Mutex<bool>>,
eventloop: Arc<Mutex<Option<EventLoop>>>,
is_connected: Arc<RwLock<bool>>,
subscriptions: Arc<Mutex<HashMap<usize, Vec<broadcast::Sender<MqttMessage>>>>>,
topic_to_subscription_ids: Arc<Mutex<HashMap<String, Vec<usize>>>>,
queued_subscriptions: Arc<Mutex<Vec<QueuedSubscription>>>,
publish_queue_tx: mpsc::Sender<QueuedMessage>,
publish_state: Arc<Mutex<PublishState>>,
connection_state_tx: watch::Sender<MqttConnectionState>,
connection_state_rx: watch::Receiver<MqttConnectionState>,
#[cfg(feature = "lwt")]
availability_helper: Option<Arc<dyn AvailabilityHelperTrait + Send + Sync>>,
#[cfg(feature = "metrics")]
metrics: Arc<metrics::Metrics>,
}
impl MqttierClient {
pub fn new(mqttier_options: MqttierOptions) -> Result<Self> {
let client_id = mqttier_options.client_id;
let (publish_queue_tx, publish_queue_rx) =
mpsc::channel::<QueuedMessage>(mqttier_options.publish_queue_size as usize);
let initial_publish_state = PublishState {
publish_queue_rx,
ack_timeout_ms: mqttier_options.ack_timeout_ms,
};
let broker = match &mqttier_options.connection {
Connection::TcpLocalhost(tcp_port) => Broker::tcp("localhost", *tcp_port),
Connection::Tcp(conn) => Broker::tcp(&conn.hostname, conn.port),
Connection::UnixSocket(path) => Broker::unix(path),
#[cfg(feature = "use-rustls")]
Connection::TcpWithTls(tls_conn) => Broker::tcp(&tls_conn.hostname, tls_conn.port),
};
let mut mqttoptions = MqttOptions::new(client_id.clone(), broker);
mqttoptions.set_keep_alive(mqttier_options.keepalive_secs);
mqttoptions.set_clean_start(true);
mqttoptions.set_max_packet_size(Some(mqttier_options.max_incoming_packet_size));
mqttoptions.set_outgoing_inflight_upper_limit(mqttier_options.max_inflight_messages);
if let Some(credentials) = &mqttier_options.credentials {
mqttoptions.set_credentials(credentials.username.clone(), credentials.password.clone());
}
#[cfg(feature = "use-rustls")]
if let Connection::TcpWithTls(tls_conn) = &mqttier_options.connection {
let client_auth = match (&tls_conn.client_cert_path, &tls_conn.client_key_path) {
(Some(cert_path), Some(key_path)) => {
match (std::fs::read(cert_path), std::fs::read(key_path)) {
(Ok(cert), Ok(key)) => {
debug!("Mutual TLS: loaded client cert and key");
Some((cert, key))
}
_ => {
error!("Failed to read client cert or key for mutual TLS");
None
}
}
}
_ => None,
};
let transport = match &tls_conn.certificate_verification {
ServerCertificateVerification::Insecure => {
warn!("TLS configured as Insecure: certificate verification is DISABLED");
let config = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(NoopCertVerifier))
.with_no_client_auth();
rumqttc::Transport::tls_with_config(rumqttc::TlsConfiguration::Rustls(
Arc::new(config),
))
}
ServerCertificateVerification::CaCertPath(ca_cert_path) => {
match std::fs::read(ca_cert_path) {
Ok(ca) => {
debug!("TLS configured with CA certificate: {}", ca_cert_path);
rumqttc::Transport::tls(ca, client_auth, None)
}
Err(e) => {
error!("Failed to read CA certificate '{}': {}", ca_cert_path, e);
rumqttc::Transport::tls_with_default_config()
}
}
}
};
mqttoptions.set_transport(transport);
}
#[cfg(feature = "lwt")]
let availability_helper = mqttier_options.availability_helper.clone();
#[cfg(feature = "lwt")]
if let Some(helper) = availability_helper.as_ref() {
let offline_message = helper.get_client_offline_message();
let will_qos = match offline_message.qos {
StingerQoS::AtMostOnce => QoS::AtMostOnce,
StingerQoS::AtLeastOnce => QoS::AtLeastOnce,
StingerQoS::ExactlyOnce => QoS::ExactlyOnce,
};
let last_will_properties = LastWillProperties {
delay_interval: Some(mqttier_options.lwt_delay_secs as u32),
payload_format_indicator: None,
message_expiry_interval: offline_message.message_expiry_interval,
content_type: offline_message.content_type.clone(),
response_topic: offline_message.response_topic.clone(),
correlation_data: offline_message.correlation_data.clone(),
user_properties: offline_message
.user_properties
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
};
let will = LastWill::new(
offline_message.topic.clone(),
offline_message.payload.clone(),
will_qos,
offline_message.retain,
Some(last_will_properties),
);
mqttoptions.set_last_will(will);
}
let (client, eventloop) = AsyncClient::new(mqttoptions.clone(), 10);
let (connection_state_tx, connection_state_rx) =
watch::channel(MqttConnectionState::Disconnected);
Ok(Self {
client_id: client_id.clone(),
client,
next_subscription_id: Arc::new(AtomicUsize::new(5)),
is_running: Arc::new(Mutex::new(false)),
eventloop: Arc::new(Mutex::new(Some(eventloop))),
is_connected: Arc::new(RwLock::new(false)),
subscriptions: Arc::new(Mutex::new(HashMap::new())),
topic_to_subscription_ids: Arc::new(Mutex::new(HashMap::new())),
queued_subscriptions: Arc::new(Mutex::new(Vec::new())),
publish_queue_tx,
publish_state: Arc::new(Mutex::new(initial_publish_state)),
connection_state_tx,
connection_state_rx,
#[cfg(feature = "lwt")]
availability_helper,
#[cfg(feature = "metrics")]
metrics: Arc::new(metrics::Metrics::new()),
})
}
fn next_subscription_id(&self) -> usize {
self.next_subscription_id.fetch_add(1, Ordering::SeqCst)
}
#[cfg(feature = "metrics")]
pub fn get_metrics(&self) -> metrics::MetricsSnapshot {
self.metrics.snapshot()
}
#[cfg(feature = "metrics")]
pub fn reset_metrics(&self) {
self.metrics.reset();
}
pub async fn subscribe(
&mut self,
topic: String,
qos: u8,
received_message_tx: broadcast::Sender<MqttMessage>,
) -> Result<usize> {
let stinger_qos = match qos {
0 => StingerQoS::AtMostOnce,
1 => StingerQoS::AtLeastOnce,
2 => StingerQoS::ExactlyOnce,
_ => return Err(MqttierError::InvalidQos(qos)),
};
let subscription_id =
<Self as Mqtt5PubSub>::subscribe(self, topic, stinger_qos, received_message_tx)
.await
.map_err(|e| MqttierError::SubscribeError(e.to_string()))?;
Ok(subscription_id as usize)
}
pub async fn run_loop(&self) -> Result<()> {
let mut is_running = self.is_running.lock().await;
if *is_running {
debug!("Run loop is already running");
return Ok(());
}
*is_running = true;
drop(is_running);
let client = self.client.clone();
let is_connected = self.is_connected.clone();
let subscriptions = self.subscriptions.clone();
let queued_subscriptions = self.queued_subscriptions.clone();
let eventloop = self.eventloop.clone();
let publish_state = self.publish_state.clone();
let connection_state_tx = self.connection_state_tx.clone();
#[cfg(feature = "metrics")]
let metrics = self.metrics.clone();
let client_for_publish = client.clone();
let is_connected_for_publish = is_connected.clone();
#[cfg(feature = "metrics")]
let metrics_for_publish = metrics.clone();
tokio::spawn(async move {
Self::publish_loop(
client_for_publish,
is_connected_for_publish,
publish_state,
#[cfg(feature = "metrics")]
metrics_for_publish,
)
.await;
error!("Publish loop has exited unexpectedly");
});
tokio::spawn(async move {
loop {
info!("Starting MQTT connection loop");
#[cfg(feature = "metrics")]
metrics.increment_connection_attempts();
let mut eventloop_guard = eventloop.lock().await;
if let Some(mut el) = eventloop_guard.take() {
drop(eventloop_guard);
let result = Self::handle_connection(
&client,
&mut el,
is_connected.clone(),
subscriptions.clone(),
queued_subscriptions.clone(),
connection_state_tx.clone(),
#[cfg(feature = "metrics")]
metrics.clone(),
)
.await;
match result {
Ok(_) => {
info!("MQTT connection loop ended normally");
}
Err(e) => {
error!("MQTT connection error: {}", e);
#[cfg(feature = "metrics")]
metrics.record_failed_connection();
}
}
let mut eventloop_guard = eventloop.lock().await;
*eventloop_guard = Some(el);
} else {
error!("EventLoop not available");
#[cfg(feature = "metrics")]
metrics.record_failed_connection();
break;
}
{
let mut cg = is_connected.write().await;
*cg = false;
}
#[cfg(feature = "metrics")]
metrics.record_disconnection();
let _ = connection_state_tx.send(MqttConnectionState::Disconnected);
warn!("Reconnecting in 5 seconds...");
#[cfg(feature = "metrics")]
metrics.increment_reconnection_count();
tokio::time::sleep(Duration::from_secs(5)).await;
}
});
#[cfg(feature = "lwt")]
if let Some(availability_helper) = self.availability_helper.clone() {
let this_client = self.clone();
tokio::spawn(async move {
let online_message = availability_helper.get_client_online_message();
let _ = this_client.clone().publish_nowait(online_message);
if let Some(interval) = availability_helper.get_republish_interval() {
loop {
tokio::time::sleep(interval).await;
let online_message = availability_helper.get_client_online_message();
let _ = this_client.clone().publish_nowait(online_message);
}
}
});
}
Ok(())
}
async fn wait_for_connection(is_connected: Arc<RwLock<bool>>) {
let mut i = 0;
loop {
if *is_connected.read().await {
break;
}
if (i % 20) == 0 {
debug!("Waiting for mqtt connection");
}
i += 1;
tokio::time::sleep(Duration::from_millis(100)).await;
}
debug!("MQTT connection good.");
}
async fn publish_loop(
client: AsyncClient,
is_connected: Arc<RwLock<bool>>,
publish_state: Arc<Mutex<PublishState>>,
#[cfg(feature = "metrics")] metrics: Arc<metrics::Metrics>,
) {
debug!("Starting publish loop");
let mut pub_state = publish_state.lock().await;
let ack_timeout_ms = pub_state.ack_timeout_ms;
while let Some(queued_message) = pub_state.publish_queue_rx.recv().await {
#[cfg(feature = "metrics")]
let start_ts = queued_message.start_timestamp;
MqttierClient::wait_for_connection(is_connected.clone()).await;
let topic = queued_message.message.topic.clone();
#[cfg(feature = "metrics")]
let payload_size = queued_message.message.payload.len();
debug!("Publishing message to topic: {}", topic);
let qos = {
match queued_message.message.qos {
StingerQoS::AtMostOnce => QoS::AtMostOnce,
StingerQoS::AtLeastOnce => QoS::AtLeastOnce,
StingerQoS::ExactlyOnce => QoS::ExactlyOnce,
}
};
#[cfg(feature = "metrics")]
let qos_u8 = match qos {
QoS::AtMostOnce => 0,
QoS::AtLeastOnce => 1,
QoS::ExactlyOnce => 2,
};
let mut pub_props = PublishProperties::default();
if let Some(resp_topic) = queued_message.message.response_topic {
pub_props.response_topic = Some(resp_topic.clone());
}
if let Some(corr_data) = queued_message.message.correlation_data {
pub_props.correlation_data = Some(corr_data.clone());
}
if !queued_message.message.user_properties.is_empty() {
let mut user_props_vec: Vec<(String, String)> = Vec::new();
for (k, v) in queued_message.message.user_properties.iter() {
user_props_vec.push((k.clone(), v.clone()));
}
pub_props.user_properties = user_props_vec;
}
if queued_message.message.content_type.is_some() {
pub_props.content_type = queued_message.message.content_type.clone();
}
let completion = queued_message.completion;
match client
.publish_with_properties_tracked(
queued_message.message.topic,
qos,
queued_message.message.retain,
queued_message.message.payload,
pub_props,
)
.await
{
Ok(notice) => {
#[cfg(feature = "metrics")]
let metrics_spawn = metrics.clone();
tokio::spawn(async move {
match tokio::time::timeout(
Duration::from_millis(ack_timeout_ms),
notice.wait_async(),
)
.await
{
Ok(Ok(())) => {
#[cfg(feature = "metrics")]
{
let latency_us = start_ts.elapsed().as_micros() as u64;
metrics_spawn.record_publish_latency(latency_us);
metrics_spawn.record_message_published(qos_u8, payload_size);
}
if let Some(c) = completion {
let success = match qos {
QoS::AtMostOnce => MqttPublishSuccess::Sent,
QoS::AtLeastOnce => MqttPublishSuccess::Acknowledged,
QoS::ExactlyOnce => MqttPublishSuccess::Completed,
};
let _ = c.send(Ok(success));
}
}
Ok(Err(e)) => {
error!("Publish notice error: {}", e);
#[cfg(feature = "metrics")]
metrics_spawn.increment_publish_failures();
if let Some(c) = completion {
let _ = c.send(Err(Mqtt5PubSubError::PublishError(format!(
"{}",
e
))));
}
}
Err(_) => {
error!(
"Publish acknowledgment timed out after {}ms",
ack_timeout_ms
);
#[cfg(feature = "metrics")]
metrics_spawn.increment_publish_failures();
if let Some(c) = completion {
let _ = c.send(Err(Mqtt5PubSubError::TimeoutError(format!(
"Publish ack timeout after {}ms",
ack_timeout_ms
))));
}
}
}
});
}
Err(e) => {
error!("Failed to publish message: {}", e);
#[cfg(feature = "metrics")]
metrics.increment_publish_failures();
if let Some(c) = completion {
let _ = c.send(Err(Mqtt5PubSubError::PublishError(format!("{}", e))));
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
async fn handle_connection(
client: &AsyncClient,
eventloop: &mut EventLoop,
is_connected: Arc<RwLock<bool>>,
subscriptions: Arc<Mutex<HashMap<usize, Vec<broadcast::Sender<MqttMessage>>>>>,
queued_subscriptions: Arc<Mutex<Vec<QueuedSubscription>>>,
connection_state_tx: watch::Sender<MqttConnectionState>,
#[cfg(feature = "metrics")] metrics: Arc<metrics::Metrics>,
) -> Result<()> {
loop {
let poll_result = eventloop.poll().await;
match poll_result {
Ok(Event::Outgoing(_)) => {}
Ok(Event::Incoming(Packet::ConnAck(_))) => {
{
let mut conn_guard = is_connected.write().await;
*conn_guard = true;
info!("CONNACK: Connected to MQTT broker");
}
#[cfg(feature = "metrics")]
metrics.record_successful_connection();
let _ = connection_state_tx.send(MqttConnectionState::Connected);
let c = client.clone();
let qs_arc = queued_subscriptions.clone();
#[cfg(feature = "metrics")]
let m = metrics.clone();
tokio::spawn(async move {
loop {
let next_subscription = {
let mut q_guard = qs_arc.lock().await;
q_guard.pop()
};
info!("Got next queued subscription: {:?}", next_subscription);
if let Some(subscription) = next_subscription {
debug!(
"Processing queued subscription for topic: {}",
subscription.topic
);
#[cfg(feature = "metrics")]
m.increment_subscription_requests();
if let Err(e) = c
.subscribe_with_properties(
&subscription.topic,
subscription.qos,
subscription.props,
)
.await
{
error!("Failed to subscribe to {}: {}", subscription.topic, e);
#[cfg(feature = "metrics")]
m.increment_subscription_failures();
} else {
#[cfg(feature = "metrics")]
m.increment_active_subscriptions();
}
} else {
debug!("Finished processing queued subscriptions");
break;
}
}
});
}
Ok(Event::Incoming(Packet::Publish(publish))) => {
let topic_str = String::from_utf8_lossy(&publish.topic).to_string();
debug!("Received message on topic: {}", topic_str);
if let Some(pub_props) = publish.properties {
let subscription_ids = pub_props.subscription_identifiers;
let correlation_data = pub_props.correlation_data;
let response_topic = pub_props.response_topic;
let content_type = pub_props.content_type;
#[cfg(feature = "metrics")]
{
let payload_size = publish.payload.len();
let qos_u8 = match publish.qos {
QoS::AtMostOnce => 0,
QoS::AtLeastOnce => 1,
QoS::ExactlyOnce => 2,
};
metrics.record_message_received(qos_u8, payload_size);
}
for subscription_id in subscription_ids {
let mut user_props_map: HashMap<String, String> = HashMap::new();
for (k, v) in pub_props.user_properties.iter() {
user_props_map.insert(k.clone(), v.clone());
}
let message = MqttMessageBuilder::default()
.topic(&topic_str)
.payload(publish.payload.clone())
.subscription_id(Some(subscription_id as u32))
.response_topic(response_topic.clone())
.correlation_data(correlation_data.clone())
.content_type(content_type.clone())
.user_properties(user_props_map)
.retain(publish.retain)
.qos(match publish.qos {
QoS::AtMostOnce => StingerQoS::AtMostOnce,
QoS::AtLeastOnce => StingerQoS::AtLeastOnce,
QoS::ExactlyOnce => StingerQoS::AtLeastOnce,
})
.build()
.unwrap();
let subs_guard = subscriptions.lock().await;
if let Some(senders) = subs_guard.get(&subscription_id) {
for sender in senders.iter() {
if let Err(e) = sender.send(message.clone()) {
warn!(
"Failed to send message to subscription {}: {}",
subscription_id, e
);
}
}
}
}
}
}
Ok(_) => {}
Err(e) => {
error!("Event loop error: {}", e);
return Err(MqttierError::ConnectionError(e));
}
}
}
}
}
#[async_trait]
impl Mqtt5PubSub for MqttierClient {
fn get_client_id(&self) -> String {
self.client_id.clone()
}
fn get_state(&self) -> watch::Receiver<MqttConnectionState> {
self.connection_state_rx.clone()
}
async fn subscribe(
&mut self,
topic: String,
qos: stinger_mqtt_trait::message::QoS,
tx: broadcast::Sender<MqttMessage>,
) -> std::result::Result<u32, Mqtt5PubSubError> {
let rumqttc_qos = match qos {
stinger_mqtt_trait::message::QoS::AtMostOnce => QoS::AtMostOnce,
stinger_mqtt_trait::message::QoS::AtLeastOnce => QoS::AtLeastOnce,
stinger_mqtt_trait::message::QoS::ExactlyOnce => QoS::ExactlyOnce,
};
let existing_id = {
let tmap = self.topic_to_subscription_ids.lock().await;
tmap.get(&topic).and_then(|ids| ids.first().copied())
};
if let Some(existing_id) = existing_id {
let mut subs = self.subscriptions.lock().await;
subs.entry(existing_id).or_insert_with(Vec::new).push(tx);
return Ok(existing_id as u32);
}
let subscription_id = self.next_subscription_id();
{
let mut subs = self.subscriptions.lock().await;
subs.entry(subscription_id)
.or_insert_with(Vec::new)
.push(tx);
}
{
let mut tmap = self.topic_to_subscription_ids.lock().await;
tmap.entry(topic.clone())
.or_insert_with(Vec::new)
.push(subscription_id);
}
let subscription_props = SubscribeProperties {
id: Some(subscription_id),
user_properties: Vec::new(),
};
#[cfg(feature = "metrics")]
self.metrics.increment_subscription_requests();
let connected = { *self.is_connected.read().await };
if connected {
debug!("Subscribing to topic: {} with QoS: {:?}", topic, qos);
let result = self
.client
.subscribe_with_properties(&topic, rumqttc_qos, subscription_props)
.await;
match result {
Ok(_) => {
#[cfg(feature = "metrics")]
self.metrics.increment_active_subscriptions();
}
Err(e) => {
#[cfg(feature = "metrics")]
self.metrics.increment_subscription_failures();
return Err(Mqtt5PubSubError::SubscriptionError(format!("{}", e)));
}
}
} else {
debug!(
"Queueing subscription for topic: {} with QoS: {:?}",
topic, qos
);
let mut queued = self.queued_subscriptions.lock().await;
queued.push(QueuedSubscription {
topic,
qos: rumqttc_qos,
props: subscription_props,
});
}
Ok(subscription_id as u32)
}
async fn unsubscribe(&mut self, topic: String) -> std::result::Result<(), Mqtt5PubSubError> {
let ids = {
let mut tmap = self.topic_to_subscription_ids.lock().await;
tmap.remove(&topic).unwrap_or_default()
};
#[cfg(feature = "metrics")]
let num_removed = ids.len();
{
let mut subs = self.subscriptions.lock().await;
for id in ids.iter() {
subs.remove(id);
}
}
{
let mut queued = self.queued_subscriptions.lock().await;
queued.retain(|s| s.topic != topic);
}
let connected = { *self.is_connected.read().await };
if connected {
self.client
.unsubscribe(&topic)
.await
.map_err(|e| Mqtt5PubSubError::UnsubscribeError(format!("{}", e)))?;
}
#[cfg(feature = "metrics")]
for _ in 0..num_removed {
self.metrics.decrement_active_subscriptions();
}
Ok(())
}
async fn publish(
&mut self,
message: MqttMessage,
) -> std::result::Result<MqttPublishSuccess, Mqtt5PubSubError> {
let (completion_tx, completion_rx) =
oneshot::channel::<std::result::Result<MqttPublishSuccess, Mqtt5PubSubError>>();
Self::wait_for_connection(self.is_connected.clone()).await;
debug!(
"Sending message to publish queue for topic: {} with QoS: {:?}",
message.topic, message.qos
);
let topic = message.topic.clone();
let queued_message = QueuedMessage {
message,
completion: Some(completion_tx),
#[cfg(feature = "metrics")]
start_timestamp: std::time::Instant::now(),
};
match self.publish_queue_tx.send(queued_message).await {
Ok(_) => {
debug!("Message to {} queued for publish", topic);
}
Err(e) => {
let mut returned = e.0;
if let Some(sender) = returned.completion.take() {
let _ = sender.send(Err(Mqtt5PubSubError::Other(
"Channel send error".to_string(),
)));
}
}
}
match tokio::time::timeout(Duration::from_millis(50000), completion_rx).await {
Ok(Ok(result)) => result,
Ok(Err(_)) => Err(Mqtt5PubSubError::PublishError(
"Completion channel closed".to_string(),
)),
Err(_) => Err(Mqtt5PubSubError::TimeoutError(
"Publish completion timeout after 50000ms".to_string(),
)),
}
}
async fn publish_noblock(
&mut self,
message: MqttMessage,
) -> oneshot::Receiver<std::result::Result<MqttPublishSuccess, Mqtt5PubSubError>> {
let (completion_tx, completion_rx) =
oneshot::channel::<std::result::Result<MqttPublishSuccess, Mqtt5PubSubError>>();
Self::wait_for_connection(self.is_connected.clone()).await;
debug!(
"Sending message to publish queue for topic: {} with QoS: {:?}",
message.topic, message.qos
);
let topic = message.topic.clone();
let queued_message = QueuedMessage {
message,
completion: Some(completion_tx),
#[cfg(feature = "metrics")]
start_timestamp: std::time::Instant::now(),
};
match self.publish_queue_tx.send(queued_message).await {
Ok(_) => {
debug!("Message to {} queued for publish", topic);
}
Err(e) => {
let mut returned = e.0;
if let Some(sender) = returned.completion.take() {
let _ = sender.send(Err(Mqtt5PubSubError::PublishError(
"Channel send error".to_string(),
)));
}
}
}
completion_rx
}
fn publish_nowait(
&mut self,
message: MqttMessage,
) -> std::result::Result<MqttPublishSuccess, Mqtt5PubSubError> {
debug!(
"Queueing message for fire-and-forget publish to topic: {} with QoS: {:?}",
message.topic, message.qos
);
let queued_message = QueuedMessage {
message,
completion: None, #[cfg(feature = "metrics")]
start_timestamp: std::time::Instant::now(),
};
self.publish_queue_tx
.try_send(queued_message)
.map_err(|e| {
Mqtt5PubSubError::PublishError(format!("Failed to queue message: {}", e))
})?;
Ok(MqttPublishSuccess::Sent)
}
#[cfg(feature = "lwt")]
fn get_availability_helper(&mut self) -> Option<Box<dyn AvailabilityHelperTrait>> {
self.availability_helper.as_ref().map(|helper| {
Box::new(ArcAvailabilityHelper {
inner: Arc::clone(helper),
}) as Box<dyn AvailabilityHelperTrait>
})
}
}
#[cfg(feature = "lwt")]
struct ArcAvailabilityHelper {
inner: Arc<dyn AvailabilityHelperTrait + Send + Sync>,
}
#[cfg(feature = "lwt")]
impl AvailabilityHelperTrait for ArcAvailabilityHelper {
fn get_client_online_message(&self) -> MqttMessage {
self.inner.get_client_online_message()
}
fn get_client_offline_message(&self) -> MqttMessage {
self.inner.get_client_offline_message()
}
fn get_online_json_path(&self) -> JsonPath {
self.inner.get_online_json_path()
}
fn get_republish_interval(&self) -> Option<Duration> {
self.inner.get_republish_interval()
}
}
impl MqttierClient {
pub async fn connect(&mut self) -> std::result::Result<(), Mqtt5PubSubError> {
let is_running = {
let guard = self.is_running.lock().await;
*guard
};
if !is_running {
self.run_loop().await.map_err(|e| {
Mqtt5PubSubError::Other(format!("Failed to start connection loop: {}", e))
})?;
}
let mut state_rx = self.connection_state_rx.clone();
let timeout_duration = Duration::from_secs(30);
let start_time = std::time::Instant::now();
loop {
if (state_rx.changed().await).is_ok() {
let state = *state_rx.borrow();
if matches!(state, MqttConnectionState::Connected) {
return Ok(());
}
}
if start_time.elapsed() > timeout_duration {
return Err(Mqtt5PubSubError::TimeoutError(
"Connection timeout".to_string(),
));
}
}
}
pub async fn disconnect(&mut self) -> std::result::Result<(), Mqtt5PubSubError> {
let _ = self
.connection_state_tx
.send(MqttConnectionState::Disconnected);
{
let mut is_connected_guard = self.is_connected.write().await;
*is_connected_guard = false;
}
self.client
.disconnect()
.await
.map_err(|e| Mqtt5PubSubError::Other(format!("Disconnection error: {}", e)))?;
Ok(())
}
pub async fn start(&mut self) -> std::result::Result<(), Mqtt5PubSubError> {
self.run_loop()
.await
.map_err(|e| Mqtt5PubSubError::Other(format!("Connection error: {}", e)))
}
pub async fn clean_stop(&mut self) -> std::result::Result<(), Mqtt5PubSubError> {
{
let mut is_running = self.is_running.lock().await;
*is_running = false;
}
self.disconnect().await?;
Ok(())
}
pub async fn force_stop(&mut self) -> std::result::Result<(), Mqtt5PubSubError> {
{
let mut is_running = self.is_running.lock().await;
*is_running = false;
}
let _ = self
.connection_state_tx
.send(MqttConnectionState::Disconnected);
{
let mut is_connected_guard = self.is_connected.write().await;
*is_connected_guard = false;
}
Ok(())
}
pub async fn reconnect(
&mut self,
_clean_start: bool,
) -> std::result::Result<(), Mqtt5PubSubError> {
{
let connected = { *self.is_connected.read().await };
if connected {
self.disconnect().await?;
}
}
Ok(())
}
}
#[cfg(all(test, feature = "lwt"))]
mod tests {
use super::*;
#[cfg(feature = "lwt")]
#[tokio::test]
async fn test_client_creation() {
let options = MqttierOptions {
connection: Connection::Tcp(TcpConnection {
hostname: "localhost".to_string(),
port: 1883,
}),
client_id: "test_client".to_string(),
ack_timeout_ms: 5000,
keepalive_secs: 60,
session_expiry_interval_secs: 1200,
availability_helper: Some(Arc::new(GenericAvailability::new("test_system"))),
publish_queue_size: 128,
max_incoming_packet_size: 10 * 1024,
max_inflight_messages: 100,
credentials: None,
};
let client = MqttierClient::new(options).unwrap();
assert_eq!(client.next_subscription_id.load(Ordering::SeqCst), 5);
}
#[cfg(feature = "lwt")]
#[tokio::test]
async fn test_client_creation_with_id() {
let client_id = "test_client".to_string();
let options = MqttierOptions {
connection: Connection::Tcp(TcpConnection {
hostname: "localhost".to_string(),
port: 1883,
}),
client_id,
ack_timeout_ms: 5000,
keepalive_secs: 60,
session_expiry_interval_secs: 1200,
availability_helper: Some(Arc::new(GenericAvailability::new("test_system"))),
publish_queue_size: 128,
max_incoming_packet_size: 10 * 1024,
max_inflight_messages: 100,
credentials: None,
};
let _client = MqttierClient::new(options).unwrap();
}
}
#[cfg(all(test, feature = "lwt"))]
mod validation_tests {
use super::*;
use stinger_mqtt_trait::Mqtt5PubSub;
#[cfg(feature = "lwt")]
#[tokio::test]
async fn test_mqtt_client_trait_implementation() {
let options = MqttierOptions {
connection: Connection::TcpLocalhost(1883),
client_id: "trait_test_client".to_string(),
ack_timeout_ms: 5000,
keepalive_secs: 60,
session_expiry_interval_secs: 1200,
availability_helper: Some(Arc::new(GenericAvailability::new("test_system"))),
publish_queue_size: 128,
max_incoming_packet_size: 10 * 1024,
max_inflight_messages: 100,
credentials: None,
};
let client = MqttierClient::new(options).expect("Failed to create client");
assert_eq!(client.get_client_id(), "trait_test_client");
let state_rx = client.get_state();
let current_state = *state_rx.borrow();
assert_eq!(current_state, MqttConnectionState::Disconnected);
}
}
#[cfg(test)]
mod builder_tests {
use super::*;
#[test]
fn test_mqttier_options_builder_defaults_and_override() {
let opts = MqttierOptionsBuilder::default()
.client_id("builder_test_client")
.ack_timeout_ms(1234u64)
.publish_queue_size(64u16)
.build()
.expect("Failed to build MqttierOptions");
assert_eq!(opts.client_id, "builder_test_client");
assert_eq!(opts.ack_timeout_ms, 1234);
assert_eq!(opts.publish_queue_size, 64);
match opts.connection {
Connection::TcpLocalhost(port) => assert_eq!(port, 1883),
_ => panic!("Unexpected default connection type"),
}
assert_eq!(opts.keepalive_secs, 60);
assert_eq!(opts.session_expiry_interval_secs, 1200);
}
}