From ff5da1de69e2ba462d01f3c5955378f35930f825 Mon Sep 17 00:00:00 2001 From: Claas Date: Sun, 20 Oct 2024 13:04:39 +0200 Subject: [PATCH] Save game 1 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit idk what I did here 😅 --- fan-controller/src/async_callback.rs | 19 ++ fan-controller/src/main.rs | 149 ++++++++------- fan-controller/src/mqtt/client.rs | 171 +++++++++++++++--- fan-controller/src/mqtt/connect.rs | 103 ++++++++++- fan-controller/src/mqtt/mod.rs | 2 + fan-controller/src/mqtt/packet.rs | 63 ++++--- .../src/mqtt/subscribe_acknowledgement.rs | 96 ++++++++-- fan-controller/src/mqtt/task.rs | 90 +++++++++ 8 files changed, 562 insertions(+), 131 deletions(-) create mode 100644 fan-controller/src/async_callback.rs create mode 100644 fan-controller/src/mqtt/task.rs diff --git a/fan-controller/src/async_callback.rs b/fan-controller/src/async_callback.rs new file mode 100644 index 0000000..d286013 --- /dev/null +++ b/fan-controller/src/async_callback.rs @@ -0,0 +1,19 @@ +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> { + type Output: 'a + Future; + fn call(&self, argument1: &'a str, argument2: &'a [u8]) -> Self::Output; +} + +impl<'a, R: 'a, F> AsyncCallback<'a> for F +where + F: Fn(&'a str, &'a [u8]) -> R, + R: Future + 'a, +{ + type Output = R; + fn call(&self, argument1: &'a str, argument2: &'a [u8]) -> Self::Output { + self(argument1, argument2) + } +} \ No newline at end of file diff --git a/fan-controller/src/main.rs b/fan-controller/src/main.rs index 7d85613..a65fff4 100644 --- a/fan-controller/src/main.rs +++ b/fan-controller/src/main.rs @@ -2,6 +2,7 @@ #![no_main] #![allow(warnings)] +use core::future::Future; use core::ops::{Deref, DerefMut}; use crc::{Crc, CRC_16_MODBUS}; use cyw43::{Control, NetDriver}; @@ -39,10 +40,12 @@ 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::FromPublish; +use crate::mqtt::packet::{get_parts, FromPublish}; use crate::mqtt::ping_request::PingRequest; use crate::mqtt::publish::Publish; use crate::mqtt::subscribe::{Subscribe, Subscription}; @@ -50,6 +53,7 @@ use crate::mqtt::QualityOfService; use crate::mqtt::{connect, publish, subscribe}; use fan::FanClient; +mod async_callback; mod configuration; mod fan; mod modbus; @@ -233,71 +237,86 @@ async fn mqtt_task( Timer::after_millis(500).await; } info!("Connected to MQTT broker through TCP"); - use mqtt::client::MqttClient; - let mut client = MqttClient::new(socket); - //TODO refactor this into a retry policy or of the sorts - let mut client = loop { - let result = client - .connect::<()>(MQTT_BROKER_USERNAME, MQTT_BROKER_PASSWORD, KEEP_ALIVE) - .await; - - match result { - Ok(client) => break client, - // Retry - Err((old_client, error)) => { - match error { - error @ ConnectError::TcpFlushError(_) - | error @ ConnectError::TcpWriteError(_) - | error @ ConnectError::ReadError(_) => { - info!( - "Error connecting. Trying again ({:?})", - Debug2Format(&error) - ); - client = old_client; - Timer::after_millis(500).await; - continue; - } - // Disconnect on these errors. (Returning does drop the tcp socket) - // This error is most likely on us and not recoverable. - //TODO retry when configuration changed - error @ ConnectError::WriteConnectError(_) - | error @ ConnectError::ErrorCode(_) - | error @ ConnectError::InvalidPacket(_) - | error @ ConnectError::DecodePacketError(_) => { - error!( - "Unrecoverable error. Closing connection. {:?}", - Debug2Format(&error) + use mqtt::task; + let packet = Connect { + client_identifier: "testfan", + username: configuration::MQTT_BROKER_USERNAME, + password: configuration::MQTT_BROKER_PASSWORD, + keep_alive_seconds: 60, + }; + + info!("Establishing MQTT connection"); + if let Err(error) = task::connect(&mut socket, packet).await { + warn!("Error connecting to MQTT broker: {:?}", error); + return; + }; + info!("MQTT connection established"); + + info!("Subscribing to MQTT topics"); + + /// A handler that takes MQTT publishes and sets the fan settings accordingly + async fn handle_publish<'f>(topic_name: &'f str, payload: &'f [u8]) { + info!("Received publish"); + // This part is not MQTT and application specific + match topic_name { + "testfan/speed/percentage" => { + let payload = match core::str::from_utf8(&payload) { + Ok(payload) => payload, + Err(error) => { + warn!("Expected percentage_command_topic payload (speed percentage) to be a valid UTF-8 string with a number"); + return; + } + }; + + // And then to an integer... + let set_point = payload.parse::(); + let set_point = match set_point { + Ok(set_point) => set_point, + Err(error) => { + warn!( + "Expected speed percentage to be a number string. Payload is: {}", + payload ); return; } - } + }; + + let Ok(setting) = FanSetting::new(set_point) else { + warn!( + "Setting fan speed out of bounds. Not accepting new setting: {}", + set_point + ); + return; + }; + + let fans = FANS.lock().await; } + + other => info!("Unexpected topic: {} with payload: {}", other, payload), } - }; - const ARRAY_REPEAT_VALUE: Option> = None; - let pending_subscriptions: [Option; 2] = [ARRAY_REPEAT_VALUE; 2]; - let (mut receiver, writer) = client.split(); - join( - async { - loop { - match receiver.receive::<()>().await { - Ok(packet) => { - info!("Received packet: {}", packet); - - match packet { - Packet::Publish(publish) => {} - Packet::SubscribeAcknowledgement(subscribe_acknowledgement) => {} - other => continue, - } - } - Err(error) => error!("Error receiving packet: {:?}", Debug2Format(&error)), - } + } + + async fn listen_for_publish<'reader, F>(reader: &mut TcpReader<'reader>, on_publish: F) + where + F: for<'a> AsyncCallback<'a>, + { + 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; } - }, - async { Timer::after_millis(3).await }, - ) - .await; + + let publish = Publish::read(parts.flags, &parts.variable_header_and_payload).unwrap(); + + on_publish.call(publish.topic_name, publish.payload).await; + } + } + + let (mut reader, writer) = socket.split(); + listen_for_publish(&mut reader, handle_publish).await; // Subscribe to home assistant topics const SUBSCRIPTIONS: [Subscription; 2] = [ @@ -320,11 +339,6 @@ async fn mqtt_task( ), }, ]; - - const SUBSCRIBE_PACKET: Subscribe = Subscribe { - subscriptions: &SUBSCRIPTIONS, - packet_identifier: 42, - }; } enum PublishReceiveError { @@ -691,6 +705,7 @@ async fn input_task(pin_18: PIN_18) { } type Fans = Mutex>>; +/// Use this to make calls to the fans through modbus static FANS: Fans = Mutex::new(None); #[embassy_executor::main] async fn main(spawner: Spawner) { @@ -723,9 +738,9 @@ async fn main(spawner: Spawner) { unwrap!(spawner.spawn(input_task(pin_18))); // The MQTT task waits for publishes from MQTT and sends them to the modbus task. // It also sends updates from the modbus task that happen through button inputs to MQTT - unwrap!(spawner.spawn(mqtt_task( - spawner, pin_23, pin_25, pio0, dma_ch0, pin_24, pin_29 - ))); + // unwrap!(spawner.spawn(mqtt_task( + // spawner, pin_23, pin_25, pio0, dma_ch0, pin_24, pin_29 + // ))); } #[cfg(test)] diff --git a/fan-controller/src/mqtt/client.rs b/fan-controller/src/mqtt/client.rs index 780b87e..87846f8 100644 --- a/fan-controller/src/mqtt/client.rs +++ b/fan-controller/src/mqtt/client.rs @@ -79,7 +79,7 @@ impl<'a> MqttClient<'a, NotConnected> { }; let mut offset = 0; - if let Err(error) = packet.write(&mut self.send_buffer, &mut offset) { + if let Err(error) = packet.encode(&mut self.send_buffer, &mut offset) { return Err((self, ConnectError::WriteConnectError(error))); }; @@ -183,10 +183,12 @@ impl<'a> MqttReceiver<'a> { } } -mod runner { +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::subscribe::{Subscribe, Subscription}; + use crate::mqtt::subscribe_acknowledgement::{SubscribeAcknowledgement, SubscribeErrorReasonCode, SubscribeReasonCode}; use crate::mqtt::{connect, packet, subscribe, ConnectErrorReasonCode}; use defmt::{warn, Format}; use embassy_futures::select::{select, Either}; @@ -195,11 +197,12 @@ mod runner { use embassy_sync::blocking_mutex::raw::CriticalSectionRawMutex; use embassy_sync::channel::{Channel, Receiver, Sender}; use embassy_time::{with_deadline, with_timeout, Duration, Instant, TimeoutError}; - use crate::mqtt::client::Connected; mod message { use crate::mqtt::connect_acknowledgement::ConnectReasonCode; - use crate::mqtt::packet::{FromPublish, Packet}; + use crate::mqtt::packet::FromPublish; + use crate::mqtt::subscribe::Subscription; + use crate::mqtt::subscribe_acknowledgement::SubscribeAcknowledgement; pub(super) enum Outgoing<'a> { Connect { @@ -207,6 +210,10 @@ mod runner { username: &'a str, password: &'a [u8], }, + Subscribe { + subscriptions: &'a [Subscription<'a>], + packet_identifier: u16, + }, } pub(super) enum Incoming @@ -214,6 +221,7 @@ mod runner { T: FromPublish, { ConnectAcknowledgement(ConnectReasonCode), + SubscribeAcknowledgement(SubscribeAcknowledgement), Publish(T), } } @@ -224,13 +232,25 @@ mod runner { }, } - struct State<'ch, T> + pub(crate) struct State<'ch, T> where T: FromPublish, { incoming: Channel, 8>, outgoing: Channel, 8>, } + + impl<'ch, T> State<'ch, T> + where + T: FromPublish, + { + pub(crate) const fn new() -> Self { + Self { + incoming: Channel::new(), + outgoing: Channel::new(), + } + } + } #[derive(Clone, Debug, Format)] enum ReceiveError { @@ -238,21 +258,21 @@ mod runner { DecodePacketError(packet::ReadError), } - struct MqttRunner<'a, T> + struct MqttRunner<'socket,'state, T> where T: FromPublish, { - tcp_receiver: TcpReader<'a>, + tcp_receiver: TcpReader<'socket>, /// Subscribe to messages to be sent from the client (like an actor handle) - outgoing: Receiver<'a, CriticalSectionRawMutex, message::Outgoing<'a>, 8>, - incoming: Sender<'a, CriticalSectionRawMutex, message::Incoming, 8>, + outgoing: Receiver<'state, CriticalSectionRawMutex, message::Outgoing<'socket>, 8>, + incoming: Sender<'state, CriticalSectionRawMutex, message::Incoming, 8>, receive_buffer: [u8; 1024], send_buffer: [u8; 256], timeout: Duration, - tcp_writer: TcpWriter<'a>, + tcp_writer: TcpWriter<'socket>, } - impl<'a, T> MqttRunner<'a, T> + impl<'socket,'state, T> MqttRunner<'socket,'state, T> where T: FromPublish, { @@ -266,7 +286,7 @@ mod runner { Ok(&self.receive_buffer[..bytes_read]) } - pub(crate) async fn run(mut self) -> ! { + pub(crate) async fn run(&'socket mut self) -> ! { loop { let result = select( self.outgoing.receive(), @@ -290,7 +310,7 @@ mod runner { }; let mut offset = 0; - if let Err(error) = packet.write(&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); @@ -311,6 +331,38 @@ mod runner { continue; } } + message::Outgoing::Subscribe { + subscriptions, + packet_identifier, + } => { + //TODO packet identifier + let packet = Subscribe { + subscriptions, + packet_identifier, + }; + + let mut offset = 0; + if let Err(error) = packet.write(&mut self.send_buffer, &mut offset) + { + //TODO handle error + warn!("Error encoding subscribe packet: {:?}", error); + continue; + }; + + if let Err(error) = + self.tcp_writer.write(&self.send_buffer[..offset]).await + { + //TODO handle error + warn!("Error writing subscribe packet: {:?}", error); + continue; + } + + if let Err(error) = self.tcp_writer.flush().await { + //TODO handle error + warn!("Error flushing subscribe packet: {:?}", error); + continue; + } + } } } Either::Second(result) => { @@ -345,7 +397,12 @@ mod runner { } //TODO - Packet::SubscribeAcknowledgement(_) => {} + Packet::SubscribeAcknowledgement(acknowledgement) => { + let message = + message::Incoming::SubscribeAcknowledgement(acknowledgement); + self.incoming.send(message).await; + continue; + } Packet::Publish(_) => {} } } @@ -378,7 +435,7 @@ mod runner { Timeout(TimeoutError), } - struct MqttClient<'a, T> + pub(crate) struct MqttClient<'a, T> where T: FromPublish, { @@ -388,14 +445,15 @@ mod runner { incoming: Receiver<'a, CriticalSectionRawMutex, message::Incoming, 8>, } - impl<'a, T> MqttClient<'a, T> + impl<'state,'socket, T> MqttClient<'state,T> where T: FromPublish, + 'state: 'socket, { - pub fn new( - state: &'a State<'a, T>, - mut socket: &'a mut TcpSocket<'a>, - ) -> (MqttRunner<'a, T>, Self) { + pub(crate) fn new( + state: &'state State<'state, T>, + mut socket: &'socket mut TcpSocket<'socket>, + ) -> (MqttRunner<'socket, 'state, T>, Self) { let (reader, writer) = socket.split(); //TODO handle out of subscribers/publishers error let send_incoming = state.incoming.sender(); @@ -425,9 +483,9 @@ mod runner { pub async fn connect( &self, - client_identifier: &'a str, - username: &'a str, - password: &'a [u8], + client_identifier: &'state str, + username: &'state str, + password: &'state [u8], ) -> Result<(), ConnectError> { let message = message::Outgoing::Connect { client_identifier, @@ -454,12 +512,71 @@ mod runner { } }; match message { - message::Incoming::ConnectAcknowledgement(reason_code) => return match reason_code { - ConnectReasonCode::Success => Ok(()), - ConnectReasonCode::ErrorCode(error_code) => Err(ConnectError::ErrorCode(error_code)), - }, + message::Incoming::ConnectAcknowledgement(reason_code) => { + return match reason_code { + ConnectReasonCode::Success => Ok(()), + ConnectReasonCode::ErrorCode(error_code) => { + Err(ConnectError::ErrorCode(error_code)) + } + } + } + _other => continue, + }; + } + } + + pub async fn subscribe( + &self, + subscriptions: &'state [Subscription<'state>; 2], + ) -> Result<[Result<(), SubscribeErrorReasonCode>; 2], SubscribeError> { + //TODO manage packet identifier + let packet_identifier = 42; + let message = message::Outgoing::Subscribe { + subscriptions, + packet_identifier, + }; + self.outgoing.send(message).await; + + 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 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); } } } diff --git a/fan-controller/src/mqtt/connect.rs b/fan-controller/src/mqtt/connect.rs index 73d095a..d9e5e3f 100644 --- a/fan-controller/src/mqtt/connect.rs +++ b/fan-controller/src/mqtt/connect.rs @@ -1,5 +1,5 @@ use defmt::Format; - +use crate::mqtt::task::Encode; use crate::mqtt::variable_byte_integer; use crate::mqtt::variable_byte_integer::VariableByteIntegerEncodeError; @@ -12,7 +12,106 @@ pub(crate) struct Connect<'a> { impl<'a> Connect<'a> { pub(crate) const TYPE: u8 = 1; - pub(crate) fn write(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), WriteError> { + pub(crate) fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), WriteError> { + let remaining_length = 11 + + size_of::() + + self.client_identifier.len() + + size_of::() + + self.username.len() + + size_of::() + + self.password.len(); + + let required_length = size_of_val(&Self::TYPE) + remaining_length; + if required_length > buffer.len() - *offset { + return Err(WriteError::BufferTooSmall { + required: required_length, + available: buffer.len() - *offset, + }); + } + + // Fixed header + buffer[*offset] = Self::TYPE << 4; + *offset += 1; + + variable_byte_integer::encode(remaining_length, buffer, offset) + .map_err(WriteError::WriteRemainingLengthError)?; + + // Variable header + // Protocol name length + buffer[*offset] = 0x00; + *offset += 1; + buffer[*offset] = 4; + *offset += 1; + + // Protocol name + buffer[*offset] = b'M'; + *offset += 1; + buffer[*offset] = b'Q'; + *offset += 1; + buffer[*offset] = b'T'; + *offset += 1; + buffer[*offset] = b'T'; + *offset += 1; + // Protocol version + buffer[*offset] = 5; + *offset += 1; + + // Connect Flags + // USER_NAME_FLAG | PASSWORD_FLAG | CLEAN_START + buffer[*offset] = 0b1100_0010; + *offset += 1; + + // Keep alive + buffer[*offset] = (self.keep_alive_seconds >> 8) as u8; + *offset += 1; + buffer[*offset] = self.keep_alive_seconds as u8; + *offset += 1; + // Property length 0 (no properties). Has to be set to 0 if there are no properties + buffer[*offset] = 0; + *offset += 1; + + // Payload + // Client identifier + let length = self.client_identifier.len(); + buffer[*offset] = (length >> 8) as u8; + *offset += 1; + buffer[*offset] = length as u8; + *offset += 1; + + for byte in self.client_identifier.as_bytes() { + buffer[*offset] = *byte; + *offset += 1; + } + // Username + let length = self.username.len(); + buffer[*offset] = (length >> 8) as u8; + *offset += 1; + buffer[*offset] = length as u8; + *offset += 1; + for byte in self.username.as_bytes() { + buffer[*offset] = *byte; + *offset += 1; + } + + // Password + let length = self.password.len(); + buffer[*offset] = (length >> 8) as u8; + *offset += 1; + buffer[*offset] = length as u8; + *offset += 1; + for byte in self.password { + buffer[*offset] = *byte; + *offset += 1; + } + + Ok(()) + } +} + +impl Encode for Connect<'_> { + type Error = WriteError; + + fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { let remaining_length = 11 + size_of::() + self.client_identifier.len() diff --git a/fan-controller/src/mqtt/mod.rs b/fan-controller/src/mqtt/mod.rs index 7418702..51fded4 100644 --- a/fan-controller/src/mqtt/mod.rs +++ b/fan-controller/src/mqtt/mod.rs @@ -16,6 +16,7 @@ pub(crate) mod publish; pub(crate) mod subscribe; mod subscribe_acknowledgement; mod variable_byte_integer; +pub(crate) mod task; #[derive(Debug, Format, Clone)] pub(super) enum ConnectErrorReasonCode { @@ -45,6 +46,7 @@ pub(super) enum ConnectErrorReasonCode { #[derive(Debug, Clone, Format)] pub struct UnknownConnectErrorReasonCode(u8); +#[derive(Format)] pub(super) enum QualityOfService { /// At most once delivery or 0 AtMostOnceDelivery = 0x00, diff --git a/fan-controller/src/mqtt/packet.rs b/fan-controller/src/mqtt/packet.rs index 2573bdd..57a4cc0 100644 --- a/fan-controller/src/mqtt/packet.rs +++ b/fan-controller/src/mqtt/packet.rs @@ -9,24 +9,34 @@ use crate::mqtt::variable_byte_integer::VariableByteIntegerDecodeError; use crate::mqtt::{publish, variable_byte_integer, ReadConnectAcknowledgementError}; use defmt::Format; +#[derive(Clone, Format, Debug)] +pub(crate) enum GetPartsError { + InvalidRemainingLength(VariableByteIntegerDecodeError), + MissingBytes(usize), +} #[derive(Debug, Clone, Format)] pub(crate) enum ReadError { /// The packet type is not supported. This can happen if there is a packet received that is /// only intended for the broker and not the client. Or the packet type is not yet implemented. UnsupportedPacketType(u8), UnexpectedPacketType(u8), - MissingBytes(usize), - InvalidRemainingLength(VariableByteIntegerDecodeError), + PartsError(GetPartsError), ConnectAcknowledgementError(ReadConnectAcknowledgementError), PublishError(publish::ReadError), SubscribeAcknowledgementError(SubscribeAcknowledgementError), } +pub(crate) struct PacketParts<'a> { + pub(crate) r#type: u8, + pub(crate) flags: u8, + pub(crate) variable_header_and_payload: &'a [u8], +} + /// T is for users of this MQTT implementation to define as the publish packets they expect vary by /// application. The only requirement is that they can be created from a publish packet which /// 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, Clone)] +#[derive(Format)] pub(crate) enum Packet where T: FromPublish, @@ -36,45 +46,56 @@ where Publish(T), } +pub(crate) fn get_parts(buffer: &[u8]) -> Result { + let packet_type = buffer[0] >> 4; + let flags = buffer[0] & 0b0000_1111; + let mut offset = 2; + let remaining_length = variable_byte_integer::decode(buffer, &mut offset) + .map_err(GetPartsError::InvalidRemainingLength)?; + + if buffer.len() < offset + remaining_length { + let missing_bytes = buffer.len() - offset - remaining_length; + return Err(GetPartsError::MissingBytes(missing_bytes)); + } + + let variable_header_and_payload = &buffer[offset..offset + remaining_length]; + + Ok(PacketParts { + r#type: packet_type, + flags, + variable_header_and_payload, + }) +} + impl Packet where T: FromPublish, { pub(crate) fn read(buffer: &[u8]) -> Result, ReadError> { // Fixed header - let packet_type = buffer[0] >> 4; - let flags = buffer[0] & 0b0000_1111; - let mut offset = 2; - let remaining_length = variable_byte_integer::decode(buffer, &mut offset) - .map_err(ReadError::InvalidRemainingLength)?; - - if buffer.len() < offset + remaining_length { - return Err(ReadError::MissingBytes( - buffer.len() - offset - remaining_length, - )); - } - let variable_header_and_payload = &buffer[offset..offset + remaining_length]; - match packet_type { - Connect::TYPE => Err(ReadError::UnsupportedPacketType(packet_type)), + let parts = get_parts(buffer).map_err(ReadError::PartsError)?; + + match parts.r#type { + Connect::TYPE => Err(ReadError::UnsupportedPacketType(parts.r#type)), ConnectAcknowledgement::TYPE => { let connect_acknowledgement = - ConnectAcknowledgement::read(variable_header_and_payload) + ConnectAcknowledgement::read(parts.variable_header_and_payload) .map_err(ReadError::ConnectAcknowledgementError)?; Ok(Packet::ConnectAcknowledgement(connect_acknowledgement)) } Publish::TYPE => { - let publish = Publish::read(flags, variable_header_and_payload) + let publish = Publish::read(parts.flags, parts.variable_header_and_payload) .map_err(ReadError::PublishError)?; let packet = T::from_publish(publish); Ok(Packet::Publish(packet)) } - Subscribe::TYPE => Err(ReadError::UnsupportedPacketType(packet_type)), + Subscribe::TYPE => Err(ReadError::UnsupportedPacketType(parts.r#type)), SubscribeAcknowledgement::TYPE => { let subscribe_acknowledgement = - SubscribeAcknowledgement::read(variable_header_and_payload) + SubscribeAcknowledgement::read(parts.variable_header_and_payload) .map_err(ReadError::SubscribeAcknowledgementError)?; Ok(Packet::SubscribeAcknowledgement(subscribe_acknowledgement)) diff --git a/fan-controller/src/mqtt/subscribe_acknowledgement.rs b/fan-controller/src/mqtt/subscribe_acknowledgement.rs index 3f0cf05..6d05119 100644 --- a/fan-controller/src/mqtt/subscribe_acknowledgement.rs +++ b/fan-controller/src/mqtt/subscribe_acknowledgement.rs @@ -1,17 +1,40 @@ -use crate::mqtt::variable_byte_integer; use crate::mqtt::variable_byte_integer::VariableByteIntegerDecodeError; -use defmt::Format; +use crate::mqtt::{variable_byte_integer, QualityOfService}; +use defmt::{warn, Format}; #[derive(Debug, Clone, Format)] pub(crate) enum SubscribeAcknowledgementError { InvalidPropertiesLength(VariableByteIntegerDecodeError), } -#[derive(Format, Clone)] -pub(crate) struct SubscribeAcknowledgement; -impl SubscribeAcknowledgement { +#[derive(Format)] +pub(crate) enum SubscribeErrorReasonCode { + UnspecifiedError = 0x80, + ImplementationSpecificError = 0x83, + NotAuthorized = 0x87, + TopicFilterInvalid = 0x8F, + PacketIdentifierInUse = 0x91, + QuotaExceeded = 0x97, + SharedSubscriptionsNotSupported = 0x9E, + SubscriptionIdentifiersNotSupported = 0xA1, + WildcardSubscriptionsNotSupported = 0xA2, +} + +#[derive(Format)] +pub(crate) enum SubscribeReasonCode { + GrantedQualityOfService(QualityOfService), + ErrorCode(SubscribeErrorReasonCode), +} + +#[derive(Format)] +pub(crate) struct SubscribeAcknowledgement { + pub(crate) packet_identifier: u16, + /// Reason code for each subscribed topic in the same order + pub(crate) reason_codes: [Option; 2], +} +impl<'a> SubscribeAcknowledgement { pub(super) const TYPE: u8 = 9; - pub(crate) fn read(buffer: &[u8]) -> Result { + pub(crate) fn read(buffer: &'a [u8]) -> Result { // Variable header let packet_identifier: u16 = ((buffer[0] as u16) << 8) | buffer[1] as u16; @@ -23,17 +46,62 @@ impl SubscribeAcknowledgement { //TODO check if topics are acknowledged offset += properties_length; - + + const DEFAULT: Option = None; + let mut reason_codes = [DEFAULT; 2]; // Payload // Reason code for each subscribed topic in the same order - let mut topic_index = 0; - - loop { - let reason_code = buffer[offset]; - offset += 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) + Ok(SubscribeAcknowledgement { + packet_identifier, + reason_codes, + }) } } diff --git a/fan-controller/src/mqtt/task.rs b/fan-controller/src/mqtt/task.rs new file mode 100644 index 0000000..440325b --- /dev/null +++ b/fan-controller/src/mqtt/task.rs @@ -0,0 +1,90 @@ +use crate::mqtt::connect::Connect; +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; + +///! Tasks that need to be done to run MQTT +///! - Keep alive + +pub(crate) trait Encode { + type Error; + fn encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error>; +} + +#[derive(Format)] +enum SendError +where + T: Encode, +{ + EncodeError(T::Error), + SendError(tcp::Error), + FlushError(tcp::Error), +} + +async fn send<'a, T>(socket: &mut TcpSocket<'a>, packet: T) -> Result<(), SendError> +where + T: Encode, +{ + let mut offset = 0; + let mut send_buffer = [0; 256]; + packet + .encode(&mut send_buffer, &mut offset) + .map_err(SendError::EncodeError)?; + socket + .write(&send_buffer[..offset]) + .await + .map_err(SendError::SendError)?; + socket.flush().await.map_err(SendError::FlushError)?; + Ok(()) +} + +#[derive(Format)] +pub(crate) enum ConnectError<'a> { + SendError(SendError>), + ReadError(tcp::Error), + PartsError(GetPartsError), + InvalidResponsePacketType(u8), + DecodeAcknowledgementError(ReadConnectAcknowledgementError), + ErrorReasonCode(ConnectErrorReasonCode), +} + +pub(crate) async fn connect<'a, 'b>( + socket: &mut TcpSocket<'a>, + packet: Connect<'b>, +) -> Result<(), ConnectError<'b>> { + send(socket, packet) + .await + .map_err(ConnectError::SendError)?; + + // Wait for connect acknowledgement + // Discard all messages before the connect acknowledgement + // The server has to send a connect acknowledgement before sending any other packet + let mut receive_buffer = [0; 1024]; + let bytes_read = socket + .read(&mut receive_buffer) + .await + .map_err(ConnectError::ReadError)?; + + let parts = + packet::get_parts(&receive_buffer[..bytes_read]).map_err(ConnectError::PartsError)?; + if parts.r#type != ConnectAcknowledgement::TYPE { + warn!( + "Expected connect acknowledgement packet, got: {:?}", + parts.r#type + ); + return Err(ConnectError::InvalidResponsePacketType(parts.r#type)); + } + + let acknowledgement = ConnectAcknowledgement::read(parts.variable_header_and_payload) + .map_err(ConnectError::DecodeAcknowledgementError)?; + + if let ConnectReasonCode::ErrorCode(error_code) = acknowledgement.connect_reason_code { + warn!("Connect error: {:?}", error_code); + return Err(ConnectError::ErrorReasonCode(error_code)); + } + + Ok(()) +} -- 2.51.2