use bevy_app::{App, Plugin, Update};
use bevy_ecs::prelude::*;
use bevy_ecs::{
event::{EntityEvent, Event},
message::Message,
observer::On,
};
use bevy_log::{debug, trace};
use bytes::Bytes;
use flume::{Receiver, bounded};
use regex::Regex;
pub use rumqttc;
use rumqttc::{ClientError, ConnectionError, QoS, SubscribeFilter};
use std::{
collections::VecDeque,
ops::{Deref, DerefMut},
panic, thread,
};
#[derive(Default)]
pub struct MqttPlugin;
impl Plugin for MqttPlugin {
fn build(&self, app: &mut App) {
app.add_message::<MqttEvent>()
.add_message::<MqttConnectError>()
.add_message::<MqttClientError>()
.add_message::<MqttPublishOutgoing>()
.add_message::<MqttPublishPacket>()
.add_message::<DisconnectMqttClient>()
.add_systems(Update, (connect_mqtt_clients, pending_subscribe_topic))
.add_systems(
Update,
(
handle_mqtt_events,
dispatch_publish_to_topic,
on_add_subscribe,
handle_outgoing_publish,
),
)
.add_observer(on_remove_subscribe);
}
}
#[derive(Component, Clone)]
pub struct MqttSetting {
pub mqtt_options: rumqttc::MqttOptions,
pub cap: usize,
}
#[derive(Component, Clone)]
pub struct MqttClient {
client: rumqttc::Client,
event_rx: Receiver<rumqttc::Event>,
error_rx: Receiver<ConnectionError>,
pending_subscribes: Vec<SubscribeFilter>,
}
impl Deref for MqttClient {
type Target = rumqttc::Client;
fn deref(&self) -> &Self::Target {
&self.client
}
}
impl DerefMut for MqttClient {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.client
}
}
#[derive(Component)]
pub struct MqttClientConnected;
#[derive(Debug, Clone, PartialEq, Eq, Event, Message)]
pub struct MqttEvent {
pub entity: Entity,
pub event: rumqttc::Event,
}
#[derive(Debug, Event, Message)]
pub struct MqttConnectError {
pub entity: Entity,
pub error: ConnectionError,
}
#[derive(Debug, Message)]
pub struct MqttClientError {
pub entity: Entity,
pub error: ClientError,
}
#[derive(Debug, Message)]
pub struct MqttPublishOutgoing {
pub entity: Entity,
pub topic: String,
pub qos: QoS,
pub retain: bool,
pub payload: Vec<u8>,
}
#[derive(Debug, Event, Message)]
pub struct MqttPublishPacket {
pub entity: Entity,
pub dup: bool,
pub qos: QoS,
pub retain: bool,
pub topic: String,
pub pkid: u16,
pub payload: Bytes,
}
#[derive(Event, Message)]
pub struct DisconnectMqttClient;
fn handle_mqtt_events(
clients: Query<(Entity, &MqttClient, &MqttSetting)>,
mut commands: Commands,
mut mqtt_events: MessageWriter<MqttEvent>,
mut error_events: MessageWriter<MqttConnectError>,
mut publish_incoming: MessageWriter<MqttPublishPacket>,
) {
for (entity, client, setting) in clients.iter() {
while let Ok(event) = client.event_rx.try_recv() {
match &event {
rumqttc::Event::Incoming(rumqttc::Incoming::ConnAck(_)) => {
debug!(
"Mqtt client connected to {:?}",
setting.mqtt_options.broker_address()
);
commands.entity(entity).insert(MqttClientConnected);
}
rumqttc::Event::Incoming(rumqttc::Incoming::Disconnect) => {
commands.entity(entity).remove::<MqttClientConnected>();
}
rumqttc::Event::Incoming(rumqttc::Incoming::Publish(publish)) => {
publish_incoming.write(MqttPublishPacket {
entity,
dup: publish.dup,
qos: publish.qos,
retain: publish.retain,
topic: publish.topic.clone(),
pkid: publish.pkid,
payload: publish.payload.clone(),
});
}
rumqttc::Event::Incoming(_) | rumqttc::Event::Outgoing(_) => {}
}
mqtt_events.write(MqttEvent {
entity,
event: event.clone(),
});
}
while let Ok(error) = client.error_rx.try_recv() {
commands
.entity(entity)
.remove::<(MqttClient, MqttClientConnected)>();
error_events.write(MqttConnectError { entity, error });
}
}
}
fn connect_mqtt_clients(
setting_query: Query<(Entity, &MqttSetting), Without<MqttClient>>,
mut commands: Commands,
) {
for (entity, setting) in setting_query.iter() {
debug!(
"Creating MQTT client for {:?}",
setting.mqtt_options.broker_address()
);
let (to_async_event, from_async_event) = bounded::<rumqttc::Event>(setting.cap);
let (to_async_error, from_async_error) = bounded::<ConnectionError>(setting.cap);
let (client, mut connection) =
rumqttc::Client::new(setting.mqtt_options.clone(), setting.cap);
let event_sender = to_async_event.clone();
let error_sender = to_async_error.clone();
thread::spawn(move || {
let thread_result = panic::catch_unwind(panic::AssertUnwindSafe(|| {
for notification in connection.iter() {
match notification {
Ok(event) => {
if event_sender.send(event).is_err() {
trace!("MQTT event channel closed, exiting thread");
return;
}
}
Err(connection_err) => {
if error_sender.send(connection_err).is_err() {
trace!("MQTT error channel closed, exiting thread.");
}
trace!("MQTT connection error, exiting thread for reconnection");
return;
}
}
}
trace!("MQTT connection iterator ended, exiting thread");
}));
if let Err(panic_info) = thread_result {
let panic_message = if let Some(s) = panic_info.downcast_ref::<&str>() {
s.to_string()
} else if let Some(s) = panic_info.downcast_ref::<String>() {
s.clone()
} else {
"Unknown panic occurred".to_string()
};
let synthetic_error = ConnectionError::Io(std::io::Error::other(format!(
"MQTT thread panicked: {}",
panic_message
)));
let _ = error_sender.send(synthetic_error);
debug!("MQTT thread panicked: {}", panic_message);
}
});
commands.entity(entity).insert(MqttClient {
client,
event_rx: from_async_event,
error_rx: from_async_error,
pending_subscribes: vec![],
});
}
}
#[derive(Debug, Clone, Component)]
pub struct SubscribeTopic {
topic: String,
qos: QoS,
re: Regex,
}
#[derive(Debug, Component)]
pub struct PacketCache {
pub packets: VecDeque<Bytes>,
pub capacity: usize,
}
impl Default for PacketCache {
fn default() -> Self {
Self::new(100) }
}
impl PacketCache {
pub fn new(capacity: usize) -> Self {
let safe_capacity = capacity.clamp(1, 1000); Self {
packets: VecDeque::with_capacity(safe_capacity),
capacity: safe_capacity,
}
}
pub fn push(&mut self, packet: Bytes) {
if self.packets.len() >= self.capacity {
self.packets.pop_front(); trace!(
"PacketCache at capacity {}, removing oldest packet",
self.capacity
);
}
self.packets.push_back(packet);
}
pub fn len(&self) -> usize {
self.packets.len()
}
pub fn is_empty(&self) -> bool {
self.packets.is_empty()
}
pub fn iter(&self) -> impl Iterator<Item = &Bytes> {
self.packets.iter()
}
pub fn clear(&mut self) {
self.packets.clear();
}
pub fn latest(&self) -> Option<&Bytes> {
self.packets.back()
}
pub fn oldest(&self) -> Option<&Bytes> {
self.packets.front()
}
}
impl SubscribeTopic {
pub fn new(topic: impl ToString, qos: QoS) -> Result<Self, regex::Error> {
let topic = topic.to_string();
let escaped_topic = topic
.chars()
.map(|c| match c {
'+' | '#' => c.to_string(),
'.' | '^' | '$' | '*' | '?' | '(' | ')' | '[' | ']' | '{' | '}' | '\\' | '|' => {
format!("\\{}", c)
}
_ => c.to_string(),
})
.collect::<String>();
let regex_pattern = escaped_topic.replace("+", "[^/]+").replace("#", ".+");
let re = Regex::new(&format!("^{}$", regex_pattern))?;
Ok(Self { topic, re, qos })
}
pub fn matches(&self, topic: &str) -> bool {
self.re.is_match(topic)
}
pub fn topic(&self) -> &str {
&self.topic
}
pub fn qos(&self) -> QoS {
self.qos
}
}
#[derive(Debug, EntityEvent)]
pub struct TopicMessage {
#[event_target]
pub target: Entity,
pub topic: String,
pub payload: Bytes,
}
fn dispatch_publish_to_topic(
mut publish_incoming: MessageReader<MqttPublishPacket>,
mut topic_query: Query<(Entity, &SubscribeTopic, Option<&mut PacketCache>)>,
parent_query: Query<&ChildOf>,
mut commands: Commands,
mut match_entities: Local<Vec<Entity>>,
) {
for packet in publish_incoming.read() {
for (e, subscribed_topic, opt_packet_cache) in topic_query.iter_mut() {
if subscribed_topic.matches(&packet.topic) {
trace!(
"{:?} {} Received matched packet",
e,
subscribed_topic.topic(),
);
match_entities.push(e);
if let Some(mut message_cache) = opt_packet_cache {
message_cache.push(packet.payload.clone());
}
for ancestor in parent_query.iter_ancestors(e) {
match_entities.push(ancestor);
}
}
}
if !match_entities.is_empty() {
let entities = std::mem::take(&mut *match_entities);
let topic = packet.topic.clone();
let payload = packet.payload.clone();
for entity in entities {
let topic_clone = topic.clone();
let payload_clone = payload.clone();
commands
.entity(entity)
.trigger(move |target: Entity| TopicMessage {
target,
topic: topic_clone,
payload: payload_clone,
});
}
}
}
}
fn pending_subscribe_topic(
mut clients: Query<(Entity, &mut MqttClient)>,
mut client_error: MessageWriter<MqttClientError>,
) {
for (entity, mut client) in clients.iter_mut() {
if client.pending_subscribes.is_empty() {
continue;
}
let sub_lists = client.pending_subscribes.drain(..).collect::<Vec<_>>();
if let Err(e) = client.subscribe_many(sub_lists) {
client_error.write(MqttClientError { entity, error: e });
}
}
}
fn on_add_subscribe(
mut clients: Query<&mut MqttClient>,
parent_query: Query<&ChildOf>,
query: Query<(Entity, &SubscribeTopic), Added<SubscribeTopic>>,
) {
for (entity, subscribe) in query.iter() {
let mut found_client = false;
for ancestor in parent_query.iter_ancestors(entity) {
if let Ok(mut client) = clients.get_mut(ancestor) {
client
.pending_subscribes
.push(SubscribeFilter::new(subscribe.topic.clone(), subscribe.qos));
found_client = true;
break; }
}
if !found_client {
debug!(
"No MQTT client found for SubscribeTopic entity {:?} with topic '{}'",
entity, subscribe.topic
);
}
}
}
fn on_remove_subscribe(
trigger: On<Remove, SubscribeTopic>,
parent_query: Query<&ChildOf>,
clients: Query<(Entity, &MqttClient)>,
subscribe_query: Query<&SubscribeTopic>,
mut client_error: MessageWriter<MqttClientError>,
) {
let target_entity = trigger.event().entity;
let subscribe = if let Ok(s) = subscribe_query.get(target_entity) {
s
} else {
trace!(
"SubscribeTopic component not found for entity {:?}",
target_entity
);
return;
};
for entity_to_check in
std::iter::once(target_entity).chain(parent_query.iter_ancestors(target_entity))
{
if let Ok((client_entity, client)) = clients.get(entity_to_check)
&& let Err(e) = client.try_unsubscribe(subscribe.topic.clone())
{
client_error.write(MqttClientError {
entity: client_entity,
error: e,
});
}
}
}
fn handle_outgoing_publish(
mut events: MessageReader<MqttPublishOutgoing>,
clients: Query<&MqttClient>,
mut client_error: MessageWriter<MqttClientError>,
) {
for event in events.read() {
if let Ok(client) = clients.get(event.entity) {
trace!(
"Publishing message to topic '{}' via client {:?}",
event.topic, event.entity
);
if let Err(e) = client.publish(
event.topic.clone(),
event.qos,
event.retain,
event.payload.clone(),
) {
client_error.write(MqttClientError {
entity: event.entity,
error: e,
});
}
} else {
debug!(
"Cannot publish to topic '{}': MqttClient not found for entity {:?}",
event.topic, event.entity
);
}
}
}
#[test]
fn test_topic_matches() {
let subscribe = SubscribeTopic::new("hello/+/world".to_string(), QoS::AtMostOnce).unwrap();
assert!(subscribe.matches("hello/1/world"));
}
#[test]
fn test_invalid_topic_pattern() {
let result = SubscribeTopic::new("hello/[invalid", QoS::AtMostOnce);
assert!(result.is_ok());
let subscribe = result.unwrap();
assert!(subscribe.matches("hello/[invalid"));
assert!(!subscribe.matches("hello/invalid"));
}
#[test]
fn test_topic_regex_escaping() {
let subscribe = SubscribeTopic::new("test/topic.with*special[chars]", QoS::AtMostOnce).unwrap();
assert!(subscribe.matches("test/topic.with*special[chars]"));
assert!(!subscribe.matches("test/topicXwithXspecialXchars"));
let subscribe_wildcard = SubscribeTopic::new("test/+/special.*", QoS::AtMostOnce).unwrap();
assert!(subscribe_wildcard.matches("test/anything/special.*"));
assert!(!subscribe_wildcard.matches("test/anything/specialXX"));
}
#[test]
fn test_packet_cache_capacity_limit() {
let mut cache = PacketCache::new(2);
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
cache.push(Bytes::from("packet1"));
assert_eq!(cache.len(), 1);
assert!(!cache.is_empty());
assert_eq!(cache.latest(), Some(&Bytes::from("packet1")));
assert_eq!(cache.oldest(), Some(&Bytes::from("packet1")));
cache.push(Bytes::from("packet2"));
assert_eq!(cache.len(), 2);
assert_eq!(cache.latest(), Some(&Bytes::from("packet2")));
assert_eq!(cache.oldest(), Some(&Bytes::from("packet1")));
cache.push(Bytes::from("packet3"));
assert_eq!(cache.len(), 2); assert_eq!(cache.latest(), Some(&Bytes::from("packet3")));
assert_eq!(cache.oldest(), Some(&Bytes::from("packet2")));
cache.clear();
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[test]
fn test_packet_cache_minimum_capacity() {
let cache = PacketCache::new(0);
assert_eq!(cache.capacity, 1);
}
#[test]
fn test_packet_cache_maximum_capacity() {
let cache = PacketCache::new(2000);
assert_eq!(cache.capacity, 1000); assert!(cache.packets.capacity() <= 1000); }
#[test]
fn test_packet_cache_capacity_consistency() {
let mut cache = PacketCache::new(5000); assert_eq!(cache.capacity, 1000);
for i in 0..1500 {
cache.push(Bytes::from(format!("packet{}", i)));
}
assert_eq!(cache.len(), 1000);
assert_eq!(cache.capacity, 1000);
assert_eq!(cache.oldest(), Some(&Bytes::from("packet500")));
assert_eq!(cache.latest(), Some(&Bytes::from("packet1499")));
}