From c80a68a99ce20c8b463fd385963dd60d176e2f9d Mon Sep 17 00:00:00 2001 From: Claas Date: Sun, 20 Oct 2024 13:08:51 +0200 Subject: [PATCH] Save game 2 --- exponential_distribution/src/main.rs | 7 +- fan-controller/src/async_callback.rs | 14 +- fan-controller/src/configuration.rs | 7 + fan-controller/src/main.rs | 292 +++++++++++++++--- fan-controller/src/mqtt/client.rs | 136 ++++---- fan-controller/src/mqtt/connect.rs | 18 +- fan-controller/src/mqtt/mod.rs | 3 +- fan-controller/src/mqtt/packet.rs | 16 +- fan-controller/src/mqtt/ping_request.rs | 13 +- fan-controller/src/mqtt/ping_response.rs | 5 + fan-controller/src/mqtt/publish.rs | 54 +++- fan-controller/src/mqtt/subscribe.rs | 78 ++++- .../src/mqtt/subscribe_acknowledgement.rs | 109 +++---- fan-controller/src/mqtt/task.rs | 23 +- 14 files changed, 578 insertions(+), 197 deletions(-) create mode 100644 fan-controller/src/mqtt/ping_response.rs diff --git a/exponential_distribution/src/main.rs b/exponential_distribution/src/main.rs index b31bc88..db97995 100644 --- a/exponential_distribution/src/main.rs +++ b/exponential_distribution/src/main.rs @@ -5,7 +5,12 @@ use plotlib::view::ContinuousView; use rand::Rng; use rand::rngs::ThreadRng; use rand_distr::num_traits::Pow; - +///! Just a silly idea to randomize the time between keep alive ping request packets for MQTT as the +///! MQTT specification states that the client can send them at any time but should not exceed +///! the keep alive interval between packets. This might detect communication issues earlier in some +///! cases, but it is also not the intention to spam the broker which is why the random distribution +///! is exponential and should lean to the max keep alive value + // Written by AI. Is this exponential distribution? fn exponential_random(random: &mut ThreadRng, max_value: f64, rate: f64) -> usize { let exp_sample = random.gen::(); diff --git a/fan-controller/src/async_callback.rs b/fan-controller/src/async_callback.rs index d286013..014a734 100644 --- a/fan-controller/src/async_callback.rs +++ b/fan-controller/src/async_callback.rs @@ -2,18 +2,18 @@ use core::future::Future; /// Workaround for calling async function callback with lifetime parameter /// Source:https://www.reddit.com/r/rust/comments/hey4oa/comment/fvv1zql/ -pub trait AsyncCallback<'a> { +pub trait AsyncCallback<'a, T> { type Output: 'a + Future; - fn call(&self, argument1: &'a str, argument2: &'a [u8]) -> Self::Output; + fn call(&self, argument: &'a T) -> Self::Output; } -impl<'a, R: 'a, F> AsyncCallback<'a> for F +impl<'a, R: 'a, F, T: 'a> AsyncCallback<'a, T> for F where - F: Fn(&'a str, &'a [u8]) -> R, + F: Fn(&'a T) -> R, R: Future + 'a, { type Output = R; - fn call(&self, argument1: &'a str, argument2: &'a [u8]) -> Self::Output { - self(argument1, argument2) + fn call(&self, argument: &'a T) -> Self::Output { + self(argument) } -} \ No newline at end of file +} diff --git a/fan-controller/src/configuration.rs b/fan-controller/src/configuration.rs index fc5f622..f86cc6a 100644 --- a/fan-controller/src/configuration.rs +++ b/fan-controller/src/configuration.rs @@ -30,6 +30,8 @@ pub(crate) const MQTT_BROKER_PASSWORD: &[u8] = b""; /// Prefix is "homeassistant", but it can be changed in home assistant configuration pub(crate) const DISCOVERY_TOPIC: &str = "homeassistant/fan/testfan/config"; +/// The keep alive interval defines the maximum time between messages sent to the broker. +/// The broker will disconnect the client if no message is received within 1.5 times of the keep alive interval. pub(crate) const KEEP_ALIVE: Duration = Duration::from_secs(60); const _: () = { // Check if the representation as u16 is correct @@ -37,3 +39,8 @@ const _: () = { core::assert!(seconds == 60); }; + +/// The timeout not to be confused with the keep alive interval is used for packets that require a +/// response packet from the broker. If the client does not receive a response within the timeout +/// the client will stop waiting for a response which can lead to a disconnect or retry in some cases. +pub(crate) const TIMEOUT: Duration = Duration::from_secs(5); \ No newline at end of file diff --git a/fan-controller/src/main.rs b/fan-controller/src/main.rs index a65fff4..de65a9a 100644 --- a/fan-controller/src/main.rs +++ b/fan-controller/src/main.rs @@ -2,14 +2,16 @@ #![no_main] #![allow(warnings)] -use core::future::Future; -use core::ops::{Deref, DerefMut}; +use core::future::{poll_fn, Future}; +use core::ops::{Deref, DerefMut, Sub}; +use core::pin::pin; +use core::task::Poll; use crc::{Crc, CRC_16_MODBUS}; use cyw43::{Control, NetDriver}; use cyw43_pio::PioSpi; use defmt::*; use embassy_executor::Spawner; -use embassy_futures::join::join; +use embassy_futures::join::{join, join3}; use embassy_net::dns::{DnsQueryType, DnsSocket}; use embassy_net::tcp::client::{TcpClient, TcpClientState}; use embassy_net::tcp::{TcpReader, TcpSocket, TcpWriter}; @@ -25,10 +27,11 @@ use embassy_rp::{bind_interrupts, dma, pio, uart, Peripheral, Peripherals}; use embassy_sync::blocking_mutex::raw::{CriticalSectionRawMutex, NoopRawMutex}; use embassy_sync::channel; use embassy_sync::channel::Channel; -use embassy_sync::mutex::Mutex; +use embassy_sync::mutex::{Mutex, MutexGuard, TryLockError}; use embassy_sync::pubsub::PubSubChannel; use embassy_sync::signal::Signal; -use embassy_time::{Duration, Instant, Ticker, Timer}; +use embassy_sync::waitqueue::AtomicWaker; +use embassy_time::{with_deadline, with_timeout, Duration, Instant, Ticker, TimeoutError, Timer}; use embedded_io_async::{Read, Write}; use embedded_nal_async::{AddrType, Dns, SocketAddr, TcpConnect}; use mqtt::client::ConnectError; @@ -41,14 +44,16 @@ use {defmt_rtt as _, panic_probe as _}; use self::mqtt::packet; use self::mqtt::packet::Packet; use crate::async_callback::AsyncCallback; -use crate::configuration::*; use crate::mqtt::client::runner::State; use crate::mqtt::connect::Connect; use crate::mqtt::connect_acknowledgement::ConnectReasonCode; -use crate::mqtt::packet::{get_parts, FromPublish}; +use crate::mqtt::packet::{get_parts, FromPublish, FromSubscribeAcknowledgement}; use crate::mqtt::ping_request::PingRequest; +use crate::mqtt::ping_response::PingResponse; use crate::mqtt::publish::Publish; use crate::mqtt::subscribe::{Subscribe, Subscription}; +use crate::mqtt::subscribe_acknowledgement::SubscribeAcknowledgement; +use crate::mqtt::task::{send, Encode}; use crate::mqtt::QualityOfService; use crate::mqtt::{connect, publish, subscribe}; use fan::FanClient; @@ -75,9 +80,45 @@ async fn wifi_task( runner.run().await } +#[embassy_executor::task] +async fn network_task(stack: &'static Stack>) -> ! { + stack.run().await +} + #[derive(Clone, Copy)] enum Event {} static EVENT: PubSubChannel = PubSubChannel::new(); +async fn trigger_event(mutex: &Mutex>) { + let mut round = 0; + loop { + Timer::after_secs(3).await; + let mut mutex = mutex.lock().await; + *mutex = Some(round); + round += 1; + } +} + +async fn wait_for_event(mutex: &Mutex>) { + loop { + let value = poll_fn(|context| match mutex.try_lock() { + Ok(guard) => match *guard { + None => Poll::Pending, + Some(value) => Poll::Ready(value), + }, + Err(_error) => Poll::Pending, + }) + .await; + + info!("Got value {}", value); + } +} +async fn how_to() { + let mutex = Mutex::>::new(None); + + let f1 = trigger_event(&mutex); + let f2 = wait_for_event(&mutex); + join(f1, f2).await; +} async fn gain_control( spawner: Spawner, @@ -116,21 +157,16 @@ async fn gain_control( (net_device, control) } -#[embassy_executor::task] -async fn network_task(stack: &'static Stack>) -> ! { - stack.run().await -} - enum MqttError { - WriteConnectError(connect::WriteError), + WriteConnectError(connect::EncodeError), WriteError(tcp::Error), FlushError(tcp::Error), ReadError(tcp::Error), ReadPacketError(packet::ReadError), ConnectError(mqtt::ConnectErrorReasonCode), - WriteSubscribeError(subscribe::WriteError), + WriteSubscribeError(subscribe::EncodeError), UnexpectedPacketType(u8), - WritePublishError(publish::WriteError), + WritePublishError(publish::EncodeError), } #[embassy_executor::task] @@ -163,7 +199,10 @@ async fn mqtt_task( // Join Wi-Fi network loop { - match control.join_wpa2(WIFI_NETWORK, WIFI_PASSWORD).await { + match control + .join_wpa2(configuration::WIFI_NETWORK, configuration::WIFI_PASSWORD) + .await + { Ok(_) => break, Err(error) => info!("Error joining Wi-Fi network with status: {}", error.status), } @@ -199,7 +238,9 @@ async fn mqtt_task( // Get home assistant MQTT broker IP address let address = loop { //TODO support IPv6 - let result = dns_client.query(MQTT_BROKER_ADDRESS, DnsQueryType::A).await; + let result = dns_client + .query(configuration::MQTT_BROKER_ADDRESS, DnsQueryType::A) + .await; let mut addresses = match result { Ok(addresses) => addresses, @@ -227,7 +268,7 @@ async fn mqtt_task( info!("MQTT broker IP address resolved"); info!("Connecting to MQTT broker through TCP"); - let endpoint = IpEndpoint::new(address, MQTT_BROKER_PORT); + let endpoint = IpEndpoint::new(address, configuration::MQTT_BROKER_PORT); // Connect while let Err(error) = socket.connect(endpoint).await { info!( @@ -243,7 +284,7 @@ async fn mqtt_task( client_identifier: "testfan", username: configuration::MQTT_BROKER_USERNAME, password: configuration::MQTT_BROKER_PASSWORD, - keep_alive_seconds: 60, + keep_alive_seconds: configuration::KEEP_ALIVE.as_secs() as u16, }; info!("Establishing MQTT connection"); @@ -255,13 +296,58 @@ async fn mqtt_task( info!("Subscribing to MQTT topics"); + //TODO yes static "global" state is bad, but I am still learning how to use wakers and polling + // with futures so this will be refactored when I made it work + /// Contains the status of the subscribe packets send out. The packet identifier represents the + /// index in the array + static ACKNOWLEDGEMENTS: Mutex = Mutex::new([false, false]); + /// The waker needs to be woken to complete the subscribe acknowledgement future. + /// The embassy documentation does not explain when to use [`AtomicWaker`] but I am assuming + /// it is useful for cases like this where I need to mutate a static. + static WAKER: AtomicWaker = AtomicWaker::new(); + async fn handle_subscribe_acknowledgement<'f>( + acknowledgement: &'f SubscribeAcknowledgement<'f>, + ) { + info!("Received subscribe acknowledgement"); + let mut acknowledgements = ACKNOWLEDGEMENTS.lock().await; + // Validate server sends a valid packet identifier or we get bamboozled and panic + let Some(value) = acknowledgements.get_mut(acknowledgement.packet_identifier as usize) + else { + warn!("Received subscribe acknowledgement for out of bounds packet identifier"); + return; + }; + } + + async fn wait_for_acknowledgement() { + //TODO if this function gets called multiple times it might never be woken up because there + // is only one waker and the other call of this function will lock the mutex. To solve this + // we could use structs from [embassy-sync::waitqueue] and/or a blocking mutex to remove the + // try_lock which is used because the lock function is async and we can not easily await here + poll_fn(|context| match ACKNOWLEDGEMENTS.try_lock() { + Ok(mut guard) => { + let packet = guard.get_mut(PACKET_IDENTIFIER as usize).unwrap(); + if *packet { + Poll::Ready(()) + } else { + // Waker needs to be overwritten on each poll. Read the Rust async book on wakers + // for more details + let waker = context.waker(); + WAKER.register(waker); + Poll::Pending + } + } + Err(_error) => Poll::Pending, + }) + .await; + } + /// A handler that takes MQTT publishes and sets the fan settings accordingly - async fn handle_publish<'f>(topic_name: &'f str, payload: &'f [u8]) { + async fn handle_publish<'f>(publish: &'f Publish<'f>) { info!("Received publish"); // This part is not MQTT and application specific - match topic_name { + match publish.topic_name { "testfan/speed/percentage" => { - let payload = match core::str::from_utf8(&payload) { + let payload = match core::str::from_utf8(publish.payload) { Ok(payload) => payload, Err(error) => { warn!("Expected percentage_command_topic payload (speed percentage) to be a valid UTF-8 string with a number"); @@ -293,30 +379,95 @@ async fn mqtt_task( let fans = FANS.lock().await; } - other => info!("Unexpected topic: {} with payload: {}", other, payload), + other => info!( + "Unexpected topic: {} with payload: {}", + other, publish.payload + ), } } - async fn listen_for_publish<'reader, F>(reader: &mut TcpReader<'reader>, on_publish: F) - where - F: for<'a> AsyncCallback<'a>, + static PING_RESPONSE: Signal = Signal::new(); + /// Callback handler for [listen](crate::listen) + async fn handle_ping_response(ping_response: PingResponse) { + info!("Received ping response"); + PING_RESPONSE.signal(ping_response); + } + + async fn listen<'reader, FS, FP, FR, F>( + reader: &mut TcpReader<'reader>, + on_subscribe_acknowledgement: FS, + on_publish: FP, + on_ping_response: FR, + ) where + FS: for<'a> AsyncCallback<'a, SubscribeAcknowledgement<'a>>, + FP: for<'a> AsyncCallback<'a, Publish<'a>>, + FR: Fn(PingResponse) -> F, + F: Future, { let mut buffer = [0; 1024]; loop { let bytes_read = reader.read(&mut buffer).await.unwrap(); let parts = get_parts(&buffer[..bytes_read]).unwrap(); - if parts.r#type != Publish::TYPE { - continue; + match parts.r#type { + Publish::TYPE => { + let publish = + Publish::read(parts.flags, &parts.variable_header_and_payload).unwrap(); + + on_publish.call(&publish).await; + } + SubscribeAcknowledgement::TYPE => { + let subscribe_acknowledgement = + SubscribeAcknowledgement::read(&parts.variable_header_and_payload).unwrap(); + + on_subscribe_acknowledgement + .call(&subscribe_acknowledgement) + .await; + } + + other => info!("Unsupported packet type {}", other), } + } + } + + let (mut reader, mut writer) = socket.split(); + // Future 1 + let listen = listen( + &mut reader, + handle_subscribe_acknowledgement, + handle_publish, + handle_ping_response, + ); + + enum Message<'a> { + Subscribe(Subscribe<'a>), + Publish(Publish<'a>), + } - let publish = Publish::read(parts.flags, &parts.variable_header_and_payload).unwrap(); + static OUTGOING: Channel = Channel::new(); + /// The instant when the last packet was sent to determine when the next keep alive has to be sent + static LAST_PACKET: Signal = Signal::new(); + async fn talk(writer: &Mutex>) { + loop { + let message = OUTGOING.receive().await; + let mut writer = writer.lock().await; + match message { + Message::Subscribe(subscribe) => { + send(&mut *writer, subscribe).await.unwrap(); + } + Message::Publish(publish) => { + send(&mut *writer, publish).await.unwrap(); + } + } - on_publish.call(publish.topic_name, publish.payload).await; + LAST_PACKET.signal(Instant::now()); } } - let (mut reader, writer) = socket.split(); - listen_for_publish(&mut reader, handle_publish).await; + // Using a mutex for the writer, so it can be shared between the task that sends messages (for + // subscribing and publishing fan speed updates) and the task that sends the keep alive ping + let writer = Mutex::>::new(writer); + // Future 2 + let talk = talk(&writer); // Subscribe to home assistant topics const SUBSCRIPTIONS: [Subscription; 2] = [ @@ -339,6 +490,72 @@ async fn mqtt_task( ), }, ]; + + async fn set_up_subscriptions() { + let message = Message::Subscribe(Subscribe { + subscriptions: &SUBSCRIPTIONS, + //TODO free identifier management + packet_identifier: PACKET_IDENTIFIER, + }); + + OUTGOING.send(message).await; + + wait_for_acknowledgement::().await; + } + + // Future 3 + let set_up = set_up_subscriptions::<0>(); + + enum ClientState { + Disconnected, + Connected, + ConnectionLost, + } + static CLIENT_STATE: Signal = Signal::new(); + + // Keep alive task + async fn keep_alive(writer: &Mutex>) { + // Send keep alive packets after connection with connect packet is established. + // Wait for a packet to be sent or the keep alive interval to time the wait out. + // If the keep alive timed the wait out, send a ping request packet. + // Else if a packet was sent, wait for the next packet to be sent with the keep alive + // interval as timeout. + + let start = LAST_PACKET.try_take().unwrap_or(Instant::now()); + let mut last_send = start; + loop { + // The server waits for 1.5 times the keep alive interval, so being off by a bit due to + // network, async overhead or the clock not being exactly precise is fine + let deadline = last_send + configuration::KEEP_ALIVE; + let result = with_deadline(deadline, LAST_PACKET.wait()).await; + match result { + Err(TimeoutError) => { + let mut writer = writer.lock().await; + // Send keep alive ping request + send(&mut *writer, PingRequest).await.unwrap(); + // Keep alive is time from when the last packet was sent and not when the ping response + // was received. Therefore, we need to reset it here + last_send = Instant::now(); + + // Wait for ping response + if let Err(TimeoutError) = + with_timeout(configuration::TIMEOUT, PING_RESPONSE.wait()).await + { + // Assume disconnect from server + error!("Timeout waiting for ping response. Disconnecting"); + CLIENT_STATE.signal(ClientState::ConnectionLost); + return; + } + } + Ok(last_packet) => { + last_send = last_packet; + } + } + } + } + + //TODO cancel all tasks when client loses connection + // join3(listen, talk, set_up).await; } enum PublishReceiveError { @@ -470,12 +687,13 @@ async fn mqtt_routine<'a>(mut socket: TcpSocket<'a>) { // Ok(()) } -async fn listen_for_publishes( +async fn listen_for_publishes( mut reader: TcpReader<'_>, mut buffer: [u8; 1024], ) -> Result<(), PublishReceiveError> where T: FromPublish, + S: FromSubscribeAcknowledgement, { loop { // Should read at least one byte @@ -484,7 +702,7 @@ where .await .map_err(PublishReceiveError::ReadError)?; - let packet = Packet::::read(&buffer[..bytes_read]) + let packet = Packet::::read(&buffer[..bytes_read]) .map_err(PublishReceiveError::ReadPacketError)?; match packet { //TODO handle disconnect packet @@ -582,7 +800,7 @@ async fn send_discovery_and_keep_alive( //TODO availability topic //TODO remove whitespace at compile time through macro, build script or const fn - // Using abbreviations to save space of binary and on the wire (haven't measured effect though...) + // Using abbreviations to save space of binary and on the wire (haven't measured the effect though...) // name -> name // uniq_id -> unique_id // stat_t -> state_topic @@ -602,7 +820,7 @@ async fn send_discovery_and_keep_alive( }"#; const DISCOVERY_PUBLISH: Publish = Publish { - topic_name: DISCOVERY_TOPIC, + topic_name: configuration::DISCOVERY_TOPIC, payload: DISCOVERY_PAYLOAD, }; @@ -623,7 +841,7 @@ async fn send_discovery_and_keep_alive( loop { keep_alive.next().await; // Send keep alive ping request - PingRequest.write(&mut send_buffer, &mut offset); + let _ = PingRequest.encode(&mut send_buffer, &mut offset); writer .write_all(&send_buffer[..offset]) .await diff --git a/fan-controller/src/mqtt/client.rs b/fan-controller/src/mqtt/client.rs index 87846f8..9659daf 100644 --- a/fan-controller/src/mqtt/client.rs +++ b/fan-controller/src/mqtt/client.rs @@ -1,4 +1,4 @@ -use crate::mqtt::packet::{FromPublish, Packet}; +use crate::mqtt::packet::{FromPublish, FromSubscribeAcknowledgement, Packet}; use super::{ connect, connect_acknowledgement::ConnectReasonCode, packet, Connect, ConnectErrorReasonCode, @@ -12,7 +12,7 @@ pub(crate) struct Connected; #[derive(Debug)] pub(crate) enum ConnectError { - WriteConnectError(connect::WriteError), + WriteConnectError(connect::EncodeError), TcpWriteError(tcp::Error), TcpFlushError(tcp::Error), ReadError(ReadError), @@ -61,7 +61,7 @@ impl<'a> MqttClient<'a, NotConnected> { } /// The server must not send any data before send - pub(crate) async fn connect( + pub(crate) async fn connect( mut self, username: &str, password: &[u8], @@ -69,6 +69,7 @@ impl<'a> MqttClient<'a, NotConnected> { ) -> Result, (Self, ConnectError)> where T: FromPublish, + S: FromSubscribeAcknowledgement, { // Send MQTT connect packet let packet = Connect { @@ -100,7 +101,7 @@ impl<'a> MqttClient<'a, NotConnected> { Err(error) => return Err((self, ConnectError::ReadError(error))), }; - let packet = match Packet::::read(bytes) { + let packet = match Packet::::read(bytes) { Ok(packet) => packet, Err(error) => return Err((self, ConnectError::DecodePacketError(error))), }; @@ -169,9 +170,10 @@ impl<'a> MqttReceiver<'a> { Ok(&self.receive_buffer[..bytes_read]) } - pub(crate) async fn receive(&mut self) -> Result, ReceiveError> + pub(crate) async fn receive(&mut self) -> Result, ReceiveError> where T: FromPublish, + S: FromSubscribeAcknowledgement, { let result = self.read().await; let bytes = match result { @@ -186,9 +188,9 @@ impl<'a> MqttReceiver<'a> { pub(crate) mod runner { use crate::mqtt::connect::Connect; use crate::mqtt::connect_acknowledgement::{ConnectAcknowledgement, ConnectReasonCode}; - use crate::mqtt::packet::{FromPublish, Packet}; + use crate::mqtt::packet::{FromPublish, FromSubscribeAcknowledgement, Packet}; use crate::mqtt::subscribe::{Subscribe, Subscription}; - use crate::mqtt::subscribe_acknowledgement::{SubscribeAcknowledgement, SubscribeErrorReasonCode, SubscribeReasonCode}; + use crate::mqtt::subscribe_acknowledgement::SubscribeErrorReasonCode; use crate::mqtt::{connect, packet, subscribe, ConnectErrorReasonCode}; use defmt::{warn, Format}; use embassy_futures::select::{select, Either}; @@ -200,9 +202,8 @@ pub(crate) mod runner { mod message { use crate::mqtt::connect_acknowledgement::ConnectReasonCode; - use crate::mqtt::packet::FromPublish; + use crate::mqtt::packet::{FromPublish, FromSubscribeAcknowledgement}; use crate::mqtt::subscribe::Subscription; - use crate::mqtt::subscribe_acknowledgement::SubscribeAcknowledgement; pub(super) enum Outgoing<'a> { Connect { @@ -216,12 +217,13 @@ pub(crate) mod runner { }, } - pub(super) enum Incoming + pub(super) enum Incoming where T: FromPublish, + S: FromSubscribeAcknowledgement, { ConnectAcknowledgement(ConnectReasonCode), - SubscribeAcknowledgement(SubscribeAcknowledgement), + SubscribeAcknowledgement(S), Publish(T), } } @@ -232,17 +234,19 @@ pub(crate) mod runner { }, } - pub(crate) struct State<'ch, T> + pub(crate) struct State<'ch, T, S> where T: FromPublish, + S: FromSubscribeAcknowledgement, { - incoming: Channel, 8>, + incoming: Channel, 8>, outgoing: Channel, 8>, } - - impl<'ch, T> State<'ch, T> + + impl<'ch, T, S> State<'ch, T, S> where T: FromPublish, + S: FromSubscribeAcknowledgement, { pub(crate) const fn new() -> Self { Self { @@ -258,23 +262,25 @@ pub(crate) mod runner { DecodePacketError(packet::ReadError), } - struct MqttRunner<'socket,'state, T> + struct MqttRunner<'socket, 'state, T, S> where T: FromPublish, + S: FromSubscribeAcknowledgement, { tcp_receiver: TcpReader<'socket>, /// Subscribe to messages to be sent from the client (like an actor handle) outgoing: Receiver<'state, CriticalSectionRawMutex, message::Outgoing<'socket>, 8>, - incoming: Sender<'state, CriticalSectionRawMutex, message::Incoming, 8>, + incoming: Sender<'state, CriticalSectionRawMutex, message::Incoming, 8>, receive_buffer: [u8; 1024], send_buffer: [u8; 256], timeout: Duration, tcp_writer: TcpWriter<'socket>, } - impl<'socket,'state, T> MqttRunner<'socket,'state, T> + impl<'socket, 'state, T, S> MqttRunner<'socket, 'state, T, S> where T: FromPublish, + S: FromSubscribeAcknowledgement, { pub(self) async fn read(&mut self) -> Result<&[u8], ReadError> { let read = self.tcp_receiver.read(&mut self.receive_buffer); @@ -310,7 +316,8 @@ pub(crate) mod runner { }; let mut offset = 0; - if let Err(error) = packet.encode(&mut self.send_buffer, &mut offset) + if let Err(error) = + packet.encode(&mut self.send_buffer, &mut offset) { //TODO handle error warn!("Error encoding connect packet: {:?}", error); @@ -375,7 +382,7 @@ pub(crate) mod runner { } }; - let result = Packet::::read(&self.receive_buffer[..bytes_read]); + let result = Packet::::read(&self.receive_buffer[..bytes_read]); let packet = match result { Ok(packet) => packet, Err(error) => { @@ -419,7 +426,7 @@ pub(crate) mod runner { #[derive(Debug)] pub(crate) enum ConnectError { - WriteConnectError(connect::WriteError), + WriteConnectError(connect::EncodeError), TcpWriteError(tcp::Error), TcpFlushError(tcp::Error), ReceiveError(ReceiveError), @@ -429,31 +436,33 @@ pub(crate) mod runner { } pub(crate) enum SubscribeError { - WriteSubscribeError(subscribe::WriteError), + WriteSubscribeError(subscribe::EncodeError), TcpWriteError(tcp::Error), TcpFlushError(tcp::Error), Timeout(TimeoutError), } - pub(crate) struct MqttClient<'a, T> + pub(crate) struct MqttClient<'a, T, S> where T: FromPublish, + S: FromSubscribeAcknowledgement, { timeout: Duration, /// Publisher for sending messages to the runner to execute outgoing: Sender<'a, CriticalSectionRawMutex, message::Outgoing<'a>, 8>, - incoming: Receiver<'a, CriticalSectionRawMutex, message::Incoming, 8>, + incoming: Receiver<'a, CriticalSectionRawMutex, message::Incoming, 8>, } - impl<'state,'socket, T> MqttClient<'state,T> + impl<'state, 'socket, T, S> MqttClient<'state, T, S> where T: FromPublish, + S: FromSubscribeAcknowledgement, 'state: 'socket, { pub(crate) fn new( - state: &'state State<'state, T>, + state: &'state State<'state, T, S>, mut socket: &'socket mut TcpSocket<'socket>, - ) -> (MqttRunner<'socket, 'state, T>, Self) { + ) -> (MqttRunner<'socket, 'state, T, S>, Self) { let (reader, writer) = socket.split(); //TODO handle out of subscribers/publishers error let send_incoming = state.incoming.sender(); @@ -539,44 +548,45 @@ pub(crate) mod runner { let deadline = Instant::now() + self.timeout; loop { - let result = with_deadline(deadline, self.incoming.receive()).await; - let message = match result { - Ok(message) => message, - Err(error) => { - //TODO handle error - warn!("Timed out receiving subscribe acknowledgement: {:?}", error); - return Err(SubscribeError::Timeout(error)); - } - }; - let (identifier, reason_codes) = match message { - message::Incoming::SubscribeAcknowledgement(SubscribeAcknowledgement { - packet_identifier, - reason_codes, - }) => (packet_identifier, reason_codes), - _other => continue, - }; - - if packet_identifier != identifier { - continue; - } + todo!(); + // let result = with_deadline(deadline, self.incoming.receive()).await; + // let message = match result { + // Ok(message) => message, + // Err(error) => { + // //TODO handle error + // warn!("Timed out receiving subscribe acknowledgement: {:?}", error); + // return Err(SubscribeError::Timeout(error)); + // } + // }; + // let (identifier, reason_codes) = match message { + // message::Incoming::SubscribeAcknowledgement(SubscribeAcknowledgement { + // packet_identifier, + // reason_codes, + // }) => (packet_identifier, reason_codes), + // _other => continue, + // }; + // + // if packet_identifier != identifier { + // continue; + // } //TODO don't assume we have 2 results - let mut results = [Ok(()), Ok(())]; - - // Write codes to results - for (index, code) in reason_codes.into_iter().enumerate() { - if index >= results.len() { - break; - } - - let Some(SubscribeReasonCode::ErrorCode(code)) = code else { - continue; - }; - - results[index] = Err(code); - } - - return Ok(results); + // let mut results = [Ok(()), Ok(())]; + // + // // Write codes to results + // for (index, code) in reason_codes.into_iter().enumerate() { + // if index >= results.len() { + // break; + // } + // + // let Some(SubscribeReasonCode::ErrorCode(code)) = code else { + // continue; + // }; + // + // results[index] = Err(code); + // } + // + // return Ok(results); } } } diff --git a/fan-controller/src/mqtt/connect.rs b/fan-controller/src/mqtt/connect.rs index d9e5e3f..15f8a25 100644 --- a/fan-controller/src/mqtt/connect.rs +++ b/fan-controller/src/mqtt/connect.rs @@ -1,7 +1,7 @@ -use defmt::Format; use crate::mqtt::task::Encode; use crate::mqtt::variable_byte_integer; use crate::mqtt::variable_byte_integer::VariableByteIntegerEncodeError; +use defmt::Format; pub(crate) struct Connect<'a> { pub(crate) client_identifier: &'a str, @@ -12,7 +12,9 @@ pub(crate) struct Connect<'a> { impl<'a> Connect<'a> { pub(crate) const TYPE: u8 = 1; - pub(crate) fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), WriteError> { + + #[deprecated(note = "Use Encode trait")] + pub(crate) fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), EncodeError> { let remaining_length = 11 + size_of::() + self.client_identifier.len() @@ -23,7 +25,7 @@ impl<'a> Connect<'a> { let required_length = size_of_val(&Self::TYPE) + remaining_length; if required_length > buffer.len() - *offset { - return Err(WriteError::BufferTooSmall { + return Err(EncodeError::BufferTooSmall { required: required_length, available: buffer.len() - *offset, }); @@ -34,7 +36,7 @@ impl<'a> Connect<'a> { *offset += 1; variable_byte_integer::encode(remaining_length, buffer, offset) - .map_err(WriteError::WriteRemainingLengthError)?; + .map_err(EncodeError::WriteRemainingLengthError)?; // Variable header // Protocol name length @@ -109,7 +111,7 @@ impl<'a> Connect<'a> { } impl Encode for Connect<'_> { - type Error = WriteError; + type Error = EncodeError; fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { let remaining_length = 11 @@ -122,7 +124,7 @@ impl Encode for Connect<'_> { let required_length = size_of_val(&Self::TYPE) + remaining_length; if required_length > buffer.len() - *offset { - return Err(WriteError::BufferTooSmall { + return Err(EncodeError::BufferTooSmall { required: required_length, available: buffer.len() - *offset, }); @@ -133,7 +135,7 @@ impl Encode for Connect<'_> { *offset += 1; variable_byte_integer::encode(remaining_length, buffer, offset) - .map_err(WriteError::WriteRemainingLengthError)?; + .map_err(EncodeError::WriteRemainingLengthError)?; // Variable header // Protocol name length @@ -208,7 +210,7 @@ impl Encode for Connect<'_> { } #[derive(Debug, Format)] -pub(crate) enum WriteError { +pub(crate) enum EncodeError { /// Client identifier + user name + password together are larger than [VariableByteInteger::MAX] DataTooLarge, /// The buffer does not contain enough empty space to write the packet diff --git a/fan-controller/src/mqtt/mod.rs b/fan-controller/src/mqtt/mod.rs index 51fded4..83c5c11 100644 --- a/fan-controller/src/mqtt/mod.rs +++ b/fan-controller/src/mqtt/mod.rs @@ -14,9 +14,10 @@ pub(crate) mod packet; pub(crate) mod ping_request; pub(crate) mod publish; pub(crate) mod subscribe; -mod subscribe_acknowledgement; +pub(crate) mod subscribe_acknowledgement; mod variable_byte_integer; pub(crate) mod task; +pub(crate) mod ping_response; #[derive(Debug, Format, Clone)] pub(super) enum ConnectErrorReasonCode { diff --git a/fan-controller/src/mqtt/packet.rs b/fan-controller/src/mqtt/packet.rs index 57a4cc0..f73a592 100644 --- a/fan-controller/src/mqtt/packet.rs +++ b/fan-controller/src/mqtt/packet.rs @@ -37,12 +37,13 @@ pub(crate) struct PacketParts<'a> { /// contains the topic name and payload. This is to get around the problem that publish topic names /// and payloads can have a variable unknown length and are difficult to pass around with lifetimes. #[derive(Format)] -pub(crate) enum Packet +pub(crate) enum Packet where T: FromPublish, + S: FromSubscribeAcknowledgement, { ConnectAcknowledgement(ConnectAcknowledgement), - SubscribeAcknowledgement(SubscribeAcknowledgement), + SubscribeAcknowledgement(S), Publish(T), } @@ -67,11 +68,12 @@ pub(crate) fn get_parts(buffer: &[u8]) -> Result { }) } -impl Packet +impl Packet where T: FromPublish, + S: FromSubscribeAcknowledgement, { - pub(crate) fn read(buffer: &[u8]) -> Result, ReadError> { + pub(crate) fn read(buffer: &[u8]) -> Result, ReadError> { // Fixed header let parts = get_parts(buffer).map_err(ReadError::PartsError)?; @@ -98,7 +100,7 @@ where SubscribeAcknowledgement::read(parts.variable_header_and_payload) .map_err(ReadError::SubscribeAcknowledgementError)?; - Ok(Packet::SubscribeAcknowledgement(subscribe_acknowledgement)) + Ok(Packet::SubscribeAcknowledgement(S::from_subscribe_acknowledgement(subscribe_acknowledgement))) } unexpected => Err(ReadError::UnexpectedPacketType(unexpected)), @@ -124,3 +126,7 @@ impl FromPublish for () { () } } + +pub(crate) trait FromSubscribeAcknowledgement { + fn from_subscribe_acknowledgement(subscribe_acknowledgement: SubscribeAcknowledgement) -> Self; +} diff --git a/fan-controller/src/mqtt/ping_request.rs b/fan-controller/src/mqtt/ping_request.rs index 5beab24..27e5276 100644 --- a/fan-controller/src/mqtt/ping_request.rs +++ b/fan-controller/src/mqtt/ping_request.rs @@ -1,12 +1,21 @@ +use core::convert::Infallible; +use core::pin::Pin; +use crate::mqtt::task::Encode; + pub(crate) struct PingRequest; impl PingRequest { pub(crate) const TYPE: u8 = 12; +} + +impl Encode for PingRequest { + type Error = Infallible; - pub(crate) fn write(self, buffer: &mut [u8], offset: &mut usize) { + fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { buffer[*offset] = Self::TYPE << 4; *offset += 1; - // This might not be 0 if we reuse the buffer + // Setting it to 0, because this might not be 0 if we reuse the buffer buffer[*offset] = 0; + Ok(()) } } diff --git a/fan-controller/src/mqtt/ping_response.rs b/fan-controller/src/mqtt/ping_response.rs new file mode 100644 index 0000000..750bd48 --- /dev/null +++ b/fan-controller/src/mqtt/ping_response.rs @@ -0,0 +1,5 @@ +pub(crate) struct PingResponse; + +impl PingResponse { + pub(crate) const TYPE: u8 = 13; +} diff --git a/fan-controller/src/mqtt/publish.rs b/fan-controller/src/mqtt/publish.rs index fe6f890..d2bd867 100644 --- a/fan-controller/src/mqtt/publish.rs +++ b/fan-controller/src/mqtt/publish.rs @@ -1,7 +1,7 @@ use core::str::Utf8Error; use defmt::{write, Debug2Format, Format, Formatter}; - +use crate::mqtt::task::Encode; use crate::mqtt::variable_byte_integer; use crate::mqtt::variable_byte_integer::{ VariableByteIntegerDecodeError, VariableByteIntegerEncodeError, @@ -13,7 +13,8 @@ pub(crate) struct Publish<'a> { pub(crate) payload: &'a [u8], } -pub(crate) enum WriteError { +#[derive(Debug)] +pub(crate) enum EncodeError { VariableByteIntegerError(VariableByteIntegerEncodeError), } @@ -48,7 +49,8 @@ impl Format for ReadError { impl<'a> Publish<'a> { pub(crate) const TYPE: u8 = 3; - pub(crate) fn write(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), WriteError> { + #[deprecated(note = "Use Encode trait")] + pub(crate) fn write(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), EncodeError> { // Fixed header //TODO set flags buffer[*offset] = Self::TYPE << 4; @@ -59,7 +61,7 @@ impl<'a> Publish<'a> { let variable_header_length = size_of::() + topic_name_length + size_of::(); let remaining_length = variable_header_length + self.payload.len(); variable_byte_integer::encode(remaining_length, buffer, offset) - .map_err(WriteError::VariableByteIntegerError)?; + .map_err(EncodeError::VariableByteIntegerError)?; // Variable header // Topic name length @@ -138,3 +140,47 @@ impl<'a> Publish<'a> { }) } } + +impl Encode for Publish<'_> { + type Error = EncodeError; + + fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { + // Fixed header + //TODO set flags + buffer[*offset] = Self::TYPE << 4; + *offset += 1; + + // Remaining length + let topic_name_length = self.topic_name.len(); + let variable_header_length = size_of::() + topic_name_length + size_of::(); + let remaining_length = variable_header_length + self.payload.len(); + variable_byte_integer::encode(remaining_length, buffer, offset) + .map_err(EncodeError::VariableByteIntegerError)?; + + // Variable header + // Topic name length + buffer[*offset] = (topic_name_length >> 8) as u8; + *offset += 1; + buffer[*offset] = topic_name_length as u8; + *offset += 1; + // Topic name + for byte in self.topic_name.as_bytes() { + buffer[*offset] = *byte; + *offset += 1; + } + + // Property length + // No properties supported for now so set to 0 + buffer[*offset] = 0; + *offset += 1; + + // Payload + // No need to set length as it will be calculated + for byte in self.payload { + buffer[*offset] = *byte; + *offset += 1; + } + + Ok(()) + } +} diff --git a/fan-controller/src/mqtt/subscribe.rs b/fan-controller/src/mqtt/subscribe.rs index 48d4edb..dad9ac1 100644 --- a/fan-controller/src/mqtt/subscribe.rs +++ b/fan-controller/src/mqtt/subscribe.rs @@ -1,7 +1,9 @@ use defmt::Format; use crate::mqtt::{QualityOfService, variable_byte_integer}; +use crate::mqtt::task::Encode; +#[derive(Debug)] pub(crate) struct Options(u8); /// Retain handling option of the subscription options @@ -51,6 +53,7 @@ impl Options { } } +#[derive(Debug)] pub(crate) struct Subscription<'a> { pub(crate) topic_filter: &'a str, pub(crate) options: Options, @@ -63,7 +66,7 @@ impl<'a> Subscription<'a> { } #[derive(Debug, Format)] -pub(crate) enum WriteError { +pub(crate) enum EncodeError { /// The buffer does not contain enough empty space to write the packet BufferTooSmall { required: usize, @@ -72,6 +75,7 @@ pub(crate) enum WriteError { RemainingLengthError(variable_byte_integer::VariableByteIntegerEncodeError), } +#[derive(Debug)] pub(crate) struct Subscribe<'a> { pub(crate) subscriptions: &'a [Subscription<'a>], pub(crate) packet_identifier: u16, @@ -79,7 +83,9 @@ pub(crate) struct Subscribe<'a> { impl<'a> Subscribe<'a> { pub(crate) const TYPE: u8 = 8; - pub(crate) fn write(self, buffer: &mut [u8], offset: &mut usize) -> Result<(), WriteError> { + + #[deprecated(note = "Use Encode trait")] + pub(crate) fn write(self, buffer: &mut [u8], offset: &mut usize) -> Result<(), EncodeError> { let variable_header_length = size_of::() + size_of::(); let payload_length = self.subscriptions.len() * size_of::() @@ -92,7 +98,7 @@ impl<'a> Subscribe<'a> { let remaining_length = variable_header_length + payload_length; let required_length = size_of_val(&Self::TYPE) + remaining_length; if required_length > buffer.len() - *offset { - return Err(WriteError::BufferTooSmall { + return Err(EncodeError::BufferTooSmall { required: required_length, available: buffer.len() - *offset, }); @@ -102,7 +108,71 @@ impl<'a> Subscribe<'a> { *offset += 1; variable_byte_integer::encode(remaining_length, buffer, offset) - .map_err(WriteError::RemainingLengthError)?; + .map_err(EncodeError::RemainingLengthError)?; + + // Variable header + // Packet Identifier + buffer[*offset] = (self.packet_identifier >> 8) as u8; + *offset += 1; + + buffer[*offset] = self.packet_identifier as u8; + *offset += 1; + + // Property length + // No properties supported for now so set to 0 + buffer[*offset] = 0; + *offset += 1; + + for subscription in self.subscriptions { + // Topic name length + let topic_name_length = subscription.topic_filter.len() as u16; + buffer[*offset] = (topic_name_length >> 8) as u8; + *offset += 1; + buffer[*offset] = topic_name_length as u8; + *offset += 1; + + // Topic name + for byte in subscription.topic_filter.as_bytes() { + buffer[*offset] = *byte; + *offset += 1; + } + + // Options + buffer[*offset] = subscription.options.0; + *offset += 1; + } + + Ok(()) + } +} + +impl Encode for Subscribe<'_> { + type Error = EncodeError; + + fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { + let variable_header_length = size_of::() + size_of::(); + + let payload_length = self.subscriptions.len() * size_of::() + + self + .subscriptions + .iter() + .map(|subscription| subscription.length()) + .sum::(); + + let remaining_length = variable_header_length + payload_length; + let required_length = size_of_val(&Self::TYPE) + remaining_length; + if required_length > buffer.len() - *offset { + return Err(EncodeError::BufferTooSmall { + required: required_length, + available: buffer.len() - *offset, + }); + } + + buffer[*offset] = Self::TYPE << 4; + *offset += 1; + + variable_byte_integer::encode(remaining_length, buffer, offset) + .map_err(EncodeError::RemainingLengthError)?; // Variable header // Packet Identifier diff --git a/fan-controller/src/mqtt/subscribe_acknowledgement.rs b/fan-controller/src/mqtt/subscribe_acknowledgement.rs index 6d05119..608f95e 100644 --- a/fan-controller/src/mqtt/subscribe_acknowledgement.rs +++ b/fan-controller/src/mqtt/subscribe_acknowledgement.rs @@ -1,6 +1,6 @@ use crate::mqtt::variable_byte_integer::VariableByteIntegerDecodeError; use crate::mqtt::{variable_byte_integer, QualityOfService}; -use defmt::{warn, Format}; +use defmt::Format; #[derive(Debug, Clone, Format)] pub(crate) enum SubscribeAcknowledgementError { @@ -27,13 +27,14 @@ pub(crate) enum SubscribeReasonCode { } #[derive(Format)] -pub(crate) struct SubscribeAcknowledgement { +pub(crate) struct SubscribeAcknowledgement<'a> { pub(crate) packet_identifier: u16, /// Reason code for each subscribed topic in the same order - pub(crate) reason_codes: [Option; 2], + pub(crate) reason_codes: &'a [u8], } -impl<'a> SubscribeAcknowledgement { - pub(super) const TYPE: u8 = 9; + +impl<'a> SubscribeAcknowledgement<'a> { + pub(crate) const TYPE: u8 = 9; pub(crate) fn read(buffer: &'a [u8]) -> Result { // Variable header let packet_identifier: u16 = ((buffer[0] as u16) << 8) | buffer[1] as u16; @@ -47,57 +48,57 @@ impl<'a> SubscribeAcknowledgement { offset += properties_length; - const DEFAULT: Option = None; - let mut reason_codes = [DEFAULT; 2]; + // const DEFAULT: Option = None; + // let mut reason_codes = [DEFAULT; 2]; // Payload // Reason code for each subscribed topic in the same order - // let reason_codes = &buffer[offset..]; - let mut index = 0; - while index < reason_codes.len() { - let Some(code) = buffer.get(offset + index) else { - break; - }; - - reason_codes[index] = Some(match code { - 0x00 => SubscribeReasonCode::GrantedQualityOfService( - QualityOfService::AtMostOnceDelivery, - ), - 0x01 => SubscribeReasonCode::GrantedQualityOfService( - QualityOfService::AtLeastOnceDelivery, - ), - 0x02 => SubscribeReasonCode::GrantedQualityOfService( - QualityOfService::ExactlyOnceDelivery, - ), - 0x80 => SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::UnspecifiedError), - 0x83 => SubscribeReasonCode::ErrorCode( - SubscribeErrorReasonCode::ImplementationSpecificError, - ), - 0x87 => SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::NotAuthorized), - 0x8F => { - SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::TopicFilterInvalid) - } - 0x91 => { - SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::PacketIdentifierInUse) - } - 0x97 => SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::QuotaExceeded), - 0x9E => SubscribeReasonCode::ErrorCode( - SubscribeErrorReasonCode::SharedSubscriptionsNotSupported, - ), - 0xA1 => SubscribeReasonCode::ErrorCode( - SubscribeErrorReasonCode::SubscriptionIdentifiersNotSupported, - ), - 0xA2 => SubscribeReasonCode::ErrorCode( - SubscribeErrorReasonCode::WildcardSubscriptionsNotSupported, - ), - other => { - //TODO handle invalid reason code - warn!("Invalid reason code: {:?}", other); - break; - } - }); - - index += 1; - } + let reason_codes = &buffer[offset..]; + // let mut index = 0; + // while index < reason_codes.len() { + // let Some(code) = buffer.get(offset + index) else { + // break; + // }; + // + // reason_codes[index] = Some(match code { + // 0x00 => SubscribeReasonCode::GrantedQualityOfService( + // QualityOfService::AtMostOnceDelivery, + // ), + // 0x01 => SubscribeReasonCode::GrantedQualityOfService( + // QualityOfService::AtLeastOnceDelivery, + // ), + // 0x02 => SubscribeReasonCode::GrantedQualityOfService( + // QualityOfService::ExactlyOnceDelivery, + // ), + // 0x80 => SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::UnspecifiedError), + // 0x83 => SubscribeReasonCode::ErrorCode( + // SubscribeErrorReasonCode::ImplementationSpecificError, + // ), + // 0x87 => SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::NotAuthorized), + // 0x8F => { + // SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::TopicFilterInvalid) + // } + // 0x91 => { + // SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::PacketIdentifierInUse) + // } + // 0x97 => SubscribeReasonCode::ErrorCode(SubscribeErrorReasonCode::QuotaExceeded), + // 0x9E => SubscribeReasonCode::ErrorCode( + // SubscribeErrorReasonCode::SharedSubscriptionsNotSupported, + // ), + // 0xA1 => SubscribeReasonCode::ErrorCode( + // SubscribeErrorReasonCode::SubscriptionIdentifiersNotSupported, + // ), + // 0xA2 => SubscribeReasonCode::ErrorCode( + // SubscribeErrorReasonCode::WildcardSubscriptionsNotSupported, + // ), + // other => { + // //TODO handle invalid reason code + // warn!("Invalid reason code: {:?}", other); + // break; + // } + // }); + // + // index += 1; + // } Ok(SubscribeAcknowledgement { packet_identifier, diff --git a/fan-controller/src/mqtt/task.rs b/fan-controller/src/mqtt/task.rs index 440325b..2f1719c 100644 --- a/fan-controller/src/mqtt/task.rs +++ b/fan-controller/src/mqtt/task.rs @@ -1,10 +1,11 @@ -use crate::mqtt::connect::Connect; +use crate::mqtt::connect::{Connect, EncodeError}; use crate::mqtt::connect_acknowledgement::{ConnectAcknowledgement, ConnectReasonCode}; use crate::mqtt::packet::GetPartsError; use crate::mqtt::{packet, ConnectErrorReasonCode, ReadConnectAcknowledgementError}; use defmt::{warn, Format}; use embassy_net::tcp; use embassy_net::tcp::TcpSocket; +use embedded_io_async::Write; ///! Tasks that need to be done to run MQTT ///! - Keep alive @@ -14,17 +15,17 @@ pub(crate) trait Encode { fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error>; } -#[derive(Format)] -enum SendError -where - T: Encode, -{ - EncodeError(T::Error), +#[derive(Debug, Format)] +pub(crate) enum SendError { + EncodeError(T), SendError(tcp::Error), FlushError(tcp::Error), } -async fn send<'a, T>(socket: &mut TcpSocket<'a>, packet: T) -> Result<(), SendError> +pub(crate) async fn send<'a, T>( + socket: &mut impl Write, + packet: T, +) -> Result<(), SendError<::Error>> where T: Encode, { @@ -42,8 +43,8 @@ where } #[derive(Format)] -pub(crate) enum ConnectError<'a> { - SendError(SendError>), +pub(crate) enum ConnectError { + SendError(SendError), ReadError(tcp::Error), PartsError(GetPartsError), InvalidResponsePacketType(u8), @@ -54,7 +55,7 @@ pub(crate) enum ConnectError<'a> { pub(crate) async fn connect<'a, 'b>( socket: &mut TcpSocket<'a>, packet: Connect<'b>, -) -> Result<(), ConnectError<'b>> { +) -> Result<(), ConnectError> { send(socket, packet) .await .map_err(ConnectError::SendError)?; -- 2.51.2