#![allow(clippy::arc_with_non_send_sync)]
use super::{InferenceMessage, InferenceResponse, NetworkAdapter};
use crate::error::{InferenceError, InferenceResult};
use crate::streaming::StreamingEngine;
use rumqttc::{AsyncClient, Broker, Event, EventLoop, MqttOptions, Packet, QoS};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::RwLock;
use tracing::{debug, error, info};
pub struct MqttAdapter {
client: AsyncClient,
eventloop: Arc<RwLock<EventLoop>>,
engine: Arc<RwLock<StreamingEngine>>,
input_topic: String,
output_topic: String,
running: Arc<RwLock<bool>>,
}
impl MqttAdapter {
pub fn new(
broker_url: &str,
client_id: &str,
input_topic: impl Into<String>,
output_topic: impl Into<String>,
engine: StreamingEngine,
) -> InferenceResult<Self> {
let url = broker_url
.strip_prefix("mqtt://")
.or_else(|| broker_url.strip_prefix("tcp://"))
.unwrap_or(broker_url);
let (host, port) = if let Some(idx) = url.find(':') {
let (h, p) = url.split_at(idx);
let port: u16 = p[1..]
.parse()
.map_err(|_| InferenceError::NetworkError("Invalid port".to_string()))?;
(h.to_string(), port)
} else {
(url.to_string(), 1883)
};
let broker = Broker::tcp(host, port);
let mut mqttoptions = MqttOptions::new(client_id, broker);
mqttoptions.set_keep_alive(30u16);
let (client, eventloop) = AsyncClient::new(mqttoptions, 10);
Ok(Self {
client,
eventloop: Arc::new(RwLock::new(eventloop)),
engine: Arc::new(RwLock::new(engine)),
input_topic: input_topic.into(),
output_topic: output_topic.into(),
running: Arc::new(RwLock::new(false)),
})
}
pub async fn run(&self) -> InferenceResult<()> {
self.client
.subscribe(&self.input_topic, QoS::AtLeastOnce)
.await
.map_err(|e| InferenceError::NetworkError(e.to_string()))?;
info!("MQTT adapter subscribed to topic: {}", self.input_topic);
*self.running.write().await = true;
while *self.running.read().await {
let event = {
let mut eventloop = self.eventloop.write().await;
eventloop.poll().await
};
match event {
Ok(Event::Incoming(Packet::Publish(publish))) => {
debug!(
"Received message on topic: {} ({} bytes)",
String::from_utf8_lossy(&publish.topic),
publish.payload.len()
);
let request: InferenceMessage = match serde_json::from_slice(&publish.payload) {
Ok(req) => req,
Err(e) => {
error!("Failed to parse message: {}", e);
continue;
}
};
let start = std::time::Instant::now();
let input = request.to_array();
let output = {
let engine = self.engine.write().await;
match engine.step_async(input).await {
Ok(out) => out,
Err(e) => {
error!("Inference error: {}", e);
continue;
}
}
};
let latency_ms = start.elapsed().as_secs_f64() * 1000.0;
let response = InferenceResponse::new(
request.request_id,
output.to_vec(),
latency_ms,
output.len(),
);
let payload = match serde_json::to_vec(&response) {
Ok(p) => p,
Err(e) => {
error!("Failed to serialize response: {}", e);
continue;
}
};
if let Err(e) = self
.client
.publish(&self.output_topic, QoS::AtLeastOnce, false, payload)
.await
{
error!("Failed to publish response: {}", e);
}
}
Ok(Event::Incoming(_)) => {
}
Ok(Event::Outgoing(_)) => {
}
Err(e) => {
error!("MQTT connection error: {}", e);
tokio::time::sleep(Duration::from_secs(5)).await;
}
}
}
Ok(())
}
pub async fn request(
&self,
input: Vec<f32>,
request_id: &str,
) -> InferenceResult<InferenceResponse> {
let msg = InferenceMessage::new(request_id, input);
let payload = serde_json::to_vec(&msg)
.map_err(|e| InferenceError::SerializationError(e.to_string()))?;
self.client
.publish(&self.input_topic, QoS::AtLeastOnce, false, payload)
.await
.map_err(|e| InferenceError::NetworkError(e.to_string()))?;
Err(InferenceError::NotImplemented(
"Response handling not implemented in this example".to_string(),
))
}
}
impl NetworkAdapter for MqttAdapter {
async fn start(&mut self) -> InferenceResult<()> {
self.run().await
}
async fn stop(&mut self) -> InferenceResult<()> {
*self.running.write().await = false;
info!("MQTT adapter stopped");
Ok(())
}
fn is_running(&self) -> bool {
tokio::task::block_in_place(|| {
tokio::runtime::Handle::current().block_on(async { *self.running.read().await })
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test(flavor = "multi_thread")]
async fn test_mqtt_adapter_creation() {
use crate::streaming::StreamConfig;
let stream_config = StreamConfig::default();
let engine = StreamingEngine::new(stream_config).unwrap();
let adapter = MqttAdapter::new(
"mqtt://test.mosquitto.org:1883",
"test-client",
"agsp/input",
"agsp/output",
engine,
);
assert!(adapter.is_ok());
let adapter = adapter.unwrap();
assert!(!adapter.is_running());
}
#[tokio::test]
async fn test_mqtt_url_parsing() {
use crate::streaming::StreamConfig;
let stream_config1 = StreamConfig::default();
let engine1 = StreamingEngine::new(stream_config1).unwrap();
let adapter1 = MqttAdapter::new("mqtt://broker:1883", "client1", "in", "out", engine1);
assert!(adapter1.is_ok());
let stream_config2 = StreamConfig::default();
let engine2 = StreamingEngine::new(stream_config2).unwrap();
let adapter2 = MqttAdapter::new("tcp://broker:1883", "client2", "in", "out", engine2);
assert!(adapter2.is_ok());
let stream_config3 = StreamConfig::default();
let engine3 = StreamingEngine::new(stream_config3).unwrap();
let adapter3 = MqttAdapter::new("broker:1883", "client3", "in", "out", engine3);
assert!(adapter3.is_ok());
}
}