diff --git a/fan-controller/src/main.rs b/fan-controller/src/main.rs index 29cb7ee..f276cae 100644 --- a/fan-controller/src/main.rs +++ b/fan-controller/src/main.rs @@ -463,23 +463,28 @@ async fn mqtt_task( match parts.r#type { Publish::TYPE => { info!("Received publish"); - let publish = - match Publish::read(parts.flags, &parts.variable_header_and_payload) { - Ok(publish) => publish, - Err(error) => { - warn!("Error reading publish: {:?}", error); - continue; - } - }; + let publish = match Publish::try_decode( + parts.flags, + &parts.variable_header_and_payload, + ) { + Ok(publish) => publish, + Err(error) => { + error!("Error reading publish: {:?}", error); + continue; + } + }; - handle_publish(&publish).await; + // debug!("Read publish"); + + // handle_publish(&publish).await; + // debug!("Handled publish"); } SubscribeAcknowledgement::TYPE => { let subscribe_acknowledgement = match SubscribeAcknowledgement::read(&parts.variable_header_and_payload) { Ok(acknowledgement) => acknowledgement, Err(error) => { - warn!("Error reading subscribe acknowledgement: {:?}", error); + error!("Error reading subscribe acknowledgement: {:?}", error); continue; } }; @@ -489,6 +494,7 @@ async fn mqtt_task( PingResponse::TYPE => { info!("Received ping response"); let ping_response = match PingResponse::try_decode( + parts.flags, &parts.variable_header_and_payload, ) { Ok(response) => response, @@ -503,7 +509,8 @@ async fn mqtt_task( Disconnect::TYPE => { info!("Received disconnect"); - let disconnect = Disconnect::try_decode(&parts.variable_header_and_payload); + let disconnect = + Disconnect::try_decode(parts.flags, &parts.variable_header_and_payload); info!("Disconnect {:?}", disconnect); //TODO disconnect TCP connection } @@ -532,11 +539,17 @@ async fn mqtt_task( match message { Message::Subscribe(subscribe) => { info!("Sending subscribe"); - send(&mut *writer, subscribe).await.unwrap(); + if let Err(error) = send(&mut *writer, subscribe).await { + error!("Error sending subscribe: {:?}", error); + continue; + } } Message::Publish(publish) => { info!("Sending publish"); - send(&mut *writer, publish).await.unwrap(); + if let Err(error) = send(&mut *writer, publish).await { + error!("Error sending publish: {:?}", error); + continue; + } } } @@ -588,8 +601,47 @@ async fn mqtt_task( info!("Set up subscriptions complete") } + async fn set_up_discovery() { + // Send discovery packet + // Configuration is like the YAML configuration that would be added in Home Assistant but as JSON + // Command topic: The MQTT topic to publish commands to change the state of the fan + //TODO set firmware version from Cargo.toml package version + //TODO think about setting hardware version, support url, and manufacturer + //TODO create single home assistant device with multiple entities for sensors in fan and the bypass + //TODO add diagnostic entity like IP address + //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 the effect though...) + // name -> name + // uniq_id -> unique_id + // stat_t -> state_topic + // cmd_t -> command_topic + // pct_stat_t -> percentage_state_topic + // pct_cmd_t -> percentage_command_topic + // spd_rng_max -> speed_range_max + // Don't need to set speed_range_min because it is 1 by default + const DISCOVERY_PAYLOAD: &[u8] = br#"{ + "name": "Fan", + "uniq_id": "testfan", + "stat_t": "testfan/on/state", + "cmd_t": "testfan/on/set", + "pct_stat_t": "testfan/speed/percentage_state", + "pct_cmd_t": "testfan/speed/percentage", + "spd_rng_max": 64000 + }"#; + + const DISCOVERY_PUBLISH: Message = Message::Publish(Publish { + topic_name: configuration::DISCOVERY_TOPIC, + payload: DISCOVERY_PAYLOAD, + }); + + OUTGOING.send(DISCOVERY_PUBLISH).await; + //TODO wait for packet acknowledgement + } + // Future 3 - let set_up = set_up_subscriptions(non_zero_u16!(1)); + let set_up = join(set_up_subscriptions(non_zero_u16!(1)), set_up_discovery()); // Keep alive task async fn keep_alive(writer: &Mutex>) { @@ -691,7 +743,7 @@ async fn send_discovery_and_keep_alive( let mut offset = 0; DISCOVERY_PUBLISH - .write(&mut send_buffer, &mut offset) + .try_encode(&mut send_buffer, &mut offset) .map_err(MqttError::WritePublishError)?; writer diff --git a/fan-controller/src/mqtt/mod.rs b/fan-controller/src/mqtt/mod.rs index 5aecff1..0ed9172 100644 --- a/fan-controller/src/mqtt/mod.rs +++ b/fan-controller/src/mqtt/mod.rs @@ -80,27 +80,27 @@ impl TryEncode for T { } } -pub(crate) trait Decode { - fn decode(variable_header_and_payload: &[u8]) -> Self +pub(crate) trait Decode<'a> { + fn decode(flags: u8, variable_header_and_payload: &'a [u8]) -> Self where Self: Sized; } -pub(crate) trait TryDecode { +pub(crate) trait TryDecode<'a> { type Error; - fn try_decode(variable_header_and_payload: &[u8]) -> Result + fn try_decode(flags: u8, variable_header_and_payload: &'a [u8]) -> Result where Self: Sized; } -impl TryDecode for T { +impl<'a, T: Decode<'a>> TryDecode<'a> for T { type Error = Infallible; - fn try_decode(variable_header_and_payload: &[u8]) -> Result + fn try_decode(flags: u8, variable_header_and_payload: &'a [u8]) -> Result where Self: Sized, { - let value = T::decode(variable_header_and_payload); + let value = T::decode(flags, variable_header_and_payload); Ok(value) } } diff --git a/fan-controller/src/mqtt/packet/connect.rs b/fan-controller/src/mqtt/packet/connect.rs index 86520f7..61ca10f 100644 --- a/fan-controller/src/mqtt/packet/connect.rs +++ b/fan-controller/src/mqtt/packet/connect.rs @@ -114,6 +114,14 @@ impl TryEncode for Connect<'_> { type Error = EncodeError; fn try_encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { + if buffer.is_empty() { + return Err(EncodeError::EmptyBuffer); + } + + // Fixed header + buffer[*offset] = Self::TYPE << 4; + *offset += 1; + let remaining_length = 11 + size_of::() + self.client_identifier.len() @@ -122,7 +130,11 @@ impl TryEncode for Connect<'_> { + size_of::() + self.password.len(); - let required_length = size_of_val(&Self::TYPE) + remaining_length; + //TODO check if we can even write fixed header + let length_length = variable_byte_integer::encode(remaining_length, buffer, offset) + .map_err(EncodeError::WriteRemainingLengthError)?; + let required_length = size_of_val(&Self::TYPE) + length_length + remaining_length; + if required_length > buffer.len() - *offset { return Err(EncodeError::BufferTooSmall { required: required_length, @@ -130,13 +142,6 @@ impl TryEncode for Connect<'_> { }); } - // Fixed header - buffer[*offset] = Self::TYPE << 4; - *offset += 1; - - variable_byte_integer::encode(remaining_length, buffer, offset) - .map_err(EncodeError::WriteRemainingLengthError)?; - // Variable header // Protocol name length buffer[*offset] = 0x00; @@ -205,12 +210,14 @@ impl TryEncode for Connect<'_> { *offset += 1; } + assert_eq!(required_length, buffer[..*offset].len()); Ok(()) } } #[derive(Debug, Format)] pub(crate) enum EncodeError { + EmptyBuffer, /// 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/packet/connect_acknowledgement.rs b/fan-controller/src/mqtt/packet/connect_acknowledgement.rs index 7acf3bb..553e74e 100644 --- a/fan-controller/src/mqtt/packet/connect_acknowledgement.rs +++ b/fan-controller/src/mqtt/packet/connect_acknowledgement.rs @@ -38,11 +38,11 @@ pub(crate) enum DecodeError { }, UnknownPropertyIdentifier(u8), } -impl TryDecode for ConnectAcknowledgement { +impl TryDecode<'_> for ConnectAcknowledgement { type Error = DecodeError; /// Reads the variable header and payload of a connect acknowledgement packet - fn decode(buffer: &[u8]) -> Result + fn try_decode(flags: u8, buffer: &[u8]) -> Result where Self: Sized, { diff --git a/fan-controller/src/mqtt/packet/disconnect.rs b/fan-controller/src/mqtt/packet/disconnect.rs index e0d3f80..1d92328 100644 --- a/fan-controller/src/mqtt/packet/disconnect.rs +++ b/fan-controller/src/mqtt/packet/disconnect.rs @@ -90,10 +90,10 @@ pub(crate) enum DecodeDisconnectError { UnknownReasonCode(UnknownReasonCode), } -impl TryDecode for Disconnect { +impl TryDecode<'_> for Disconnect { type Error = DecodeDisconnectError; - fn try_decode(variable_header_and_payload: &[u8]) -> Result { + fn try_decode(flags: u8, variable_header_and_payload: &[u8]) -> Result { // Variable header // Disconnect reason code let reason_code = variable_header_and_payload[0]; diff --git a/fan-controller/src/mqtt/packet/mod.rs b/fan-controller/src/mqtt/packet/mod.rs index c4efb00..981e905 100644 --- a/fan-controller/src/mqtt/packet/mod.rs +++ b/fan-controller/src/mqtt/packet/mod.rs @@ -97,14 +97,16 @@ where match parts.r#type { Connect::TYPE => Err(ReadError::UnsupportedPacketType(parts.r#type)), ConnectAcknowledgement::TYPE => { - let connect_acknowledgement = - ConnectAcknowledgement::try_decode(parts.variable_header_and_payload) - .map_err(ReadError::ConnectAcknowledgementError)?; + let connect_acknowledgement = ConnectAcknowledgement::try_decode( + parts.flags, + parts.variable_header_and_payload, + ) + .map_err(ReadError::ConnectAcknowledgementError)?; Ok(Packet::ConnectAcknowledgement(connect_acknowledgement)) } Publish::TYPE => { - let publish = Publish::read(parts.flags, parts.variable_header_and_payload) + let publish = Publish::try_decode(parts.flags, parts.variable_header_and_payload) .map_err(ReadError::PublishError)?; let packet = T::from_publish(publish); diff --git a/fan-controller/src/mqtt/packet/ping_response.rs b/fan-controller/src/mqtt/packet/ping_response.rs index 641fc54..2fc6a64 100644 --- a/fan-controller/src/mqtt/packet/ping_response.rs +++ b/fan-controller/src/mqtt/packet/ping_response.rs @@ -8,9 +8,9 @@ impl PingResponse { pub(crate) const TYPE: u8 = 13; } -impl Decode for PingResponse { +impl Decode<'_> for PingResponse { /// Returns a ping response as ping response don't have a variable header or payload - fn decode(_variable_header_and_payload: &[u8]) -> Self + fn decode(_flags: u8, _variable_header_and_payload: &[u8]) -> Self where Self: Sized, { diff --git a/fan-controller/src/mqtt/packet/publish.rs b/fan-controller/src/mqtt/packet/publish.rs index 2b7b64d..4a45bc9 100644 --- a/fan-controller/src/mqtt/packet/publish.rs +++ b/fan-controller/src/mqtt/packet/publish.rs @@ -1,9 +1,9 @@ use core::str::Utf8Error; -use crate::mqtt::variable_byte_integer; use crate::mqtt::variable_byte_integer::VariableByteIntegerEncodeError; use crate::mqtt::TryEncode; -use defmt::{write, Debug2Format, Format, Formatter}; +use crate::mqtt::{variable_byte_integer, TryDecode}; +use defmt::{debug, info, write, Debug2Format, Format, Formatter}; #[derive(Format, Clone)] pub(crate) struct Publish<'a> { @@ -11,8 +11,14 @@ pub(crate) struct Publish<'a> { pub(crate) payload: &'a [u8], } -#[derive(Debug)] +#[derive(Debug, Format)] pub(crate) enum EncodeError { + EmptyBuffer, + /// The buffer does not contain enough empty space to write the packet + BufferTooSmall { + required: usize, + available: usize, + }, VariableByteIntegerError(VariableByteIntegerEncodeError), } @@ -46,9 +52,16 @@ impl Format for ReadError { impl<'a> Publish<'a> { pub(crate) const TYPE: u8 = 3; +} + +impl TryEncode for Publish<'_> { + type Error = EncodeError; + + fn try_encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { + if buffer.is_empty() { + return Err(EncodeError::EmptyBuffer); + } - #[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; @@ -58,9 +71,19 @@ impl<'a> Publish<'a> { 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) + + let length_length = variable_byte_integer::encode(remaining_length, buffer, offset) .map_err(EncodeError::VariableByteIntegerError)?; + let required_length = size_of_val(&Self::TYPE) + length_length + remaining_length; + + if required_length > buffer.len() - *offset { + return Err(EncodeError::BufferTooSmall { + required: required_length, + available: buffer.len() - *offset, + }); + } + // Variable header // Topic name length buffer[*offset] = (topic_name_length >> 8) as u8; @@ -78,6 +101,7 @@ impl<'a> Publish<'a> { buffer[*offset] = 0; *offset += 1; + info!("Encode 2"); // Payload // No need to set length as it will be calculated for byte in self.payload { @@ -85,10 +109,16 @@ impl<'a> Publish<'a> { *offset += 1; } + assert_eq!(required_length, buffer[..*offset].len()); + Ok(()) } +} + +impl<'a> TryDecode<'a> for Publish<'a> { + type Error = ReadError; - pub(crate) fn read(flags: u8, buffer: &'a [u8]) -> Result { + fn try_decode(flags: u8, variable_header_and_payload: &'a [u8]) -> Result { // let is_re_delivery = (flags & 0b0000_1000) != 0; let quality_of_service_level = (flags & 0b0000_0110) >> 1; if quality_of_service_level > 0 { @@ -103,82 +133,38 @@ impl<'a> Publish<'a> { // let is_retain = (flags & 0b0000_0001) != 0; let mut offset = 0; - let remaining_length = variable_byte_integer::decode(buffer, &mut offset) - .map_err(ReadError::VariableByteIntegerError)?; - //TODO check lengths - // Variable header - let topic_length = ((buffer[offset] as u16) << 8) | buffer[offset + 1] as u16; + let topic_length = ((variable_header_and_payload[offset] as u16) << 8) + | variable_header_and_payload[offset + 1] as u16; + offset += 2; if topic_length == 0 { return Err(ReadError::ZeroLengthTopicName); } - let topic_name = core::str::from_utf8(&buffer[offset..offset + topic_length as usize]) - .map_err(ReadError::InvalidTopicName)?; + let topic_name = core::str::from_utf8( + &variable_header_and_payload[offset..offset + topic_length as usize], + ) + .map_err(ReadError::InvalidTopicName)?; //TODO validate topic name does not contain MQTT wildcard characters offset += topic_length as usize; // Properties - let properties_length = variable_byte_integer::decode(buffer, &mut offset) - .map_err(ReadError::VariableByteIntegerError)?; + let properties_length = + variable_byte_integer::decode(variable_header_and_payload, &mut offset) + .map_err(ReadError::VariableByteIntegerError)?; // Ignore properties for now offset += properties_length; // Payload - let payload_length = remaining_length - offset; - //TODO validate there is enough space left in the buffer - let payload = &buffer[offset..offset + payload_length]; + let payload = &variable_header_and_payload[offset..]; + debug!("8"); Ok(Publish { topic_name, payload, }) } } - -impl TryEncode for Publish<'_> { - type Error = EncodeError; - - fn try_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/packet/subscribe.rs b/fan-controller/src/mqtt/packet/subscribe.rs index f9b9c25..2088f72 100644 --- a/fan-controller/src/mqtt/packet/subscribe.rs +++ b/fan-controller/src/mqtt/packet/subscribe.rs @@ -69,6 +69,7 @@ impl<'a> Subscription<'a> { #[derive(Debug, Format)] pub(crate) enum EncodeError { + EmptyBuffer, /// The buffer does not contain enough empty space to write the packet BufferTooSmall { required: usize, @@ -93,16 +94,14 @@ impl TryEncode for Subscribe<'_> { // https://www.emqx.com/en/blog/mqtt-5-0-control-packets-03-subscribe-unsubscribe fn try_encode(&self, buffer: &mut [u8], offset: &mut usize) -> Result<(), Self::Error> { - // 82 0a 05 be 00 00 04 64 65 6d 6f 02 - // let test_packet = &[ - // 0x82, 0x0a, 0x05, 0xbe, 0x00, 0x00, 0x04, 0x64, 0x65, 0x6d, 0x6f, 0x02, - // ]; - // for byte in test_packet { - // buffer[*offset] = *byte; - // *offset += 1; - // } + if buffer.is_empty() { + return Err(EncodeError::EmptyBuffer); + } - // return Ok(()); + // Fixed header + // Need to set type and fixed/reserved bit in first byte + buffer[*offset] = Self::TYPE << 4 | 0b0000_0010; + *offset += 1; // Calculate lengths to bail early if buffer is too small let variable_header_length = size_of::() + size_of::(); @@ -115,7 +114,10 @@ impl TryEncode for Subscribe<'_> { .sum::(); let remaining_length = variable_header_length + payload_length; - let required_length = size_of_val(&Self::TYPE) + remaining_length; + let length_length = variable_byte_integer::encode(remaining_length, buffer, offset) + .map_err(EncodeError::RemainingLengthError)?; + + let required_length = size_of_val(&Self::TYPE) + length_length + remaining_length; if required_length > buffer.len() - *offset { return Err(EncodeError::BufferTooSmall { required: required_length, @@ -123,13 +125,6 @@ impl TryEncode for Subscribe<'_> { }); } - // Need to set type and fixed/reserved bit in first byte - buffer[*offset] = Self::TYPE << 4 | 0b0000_0010; - *offset += 1; - - variable_byte_integer::encode(remaining_length, buffer, offset) - .map_err(EncodeError::RemainingLengthError)?; - // Variable header // Packet Identifier buffer[*offset] = (Into::::into(self.packet_identifier) >> 8) as u8; @@ -165,6 +160,7 @@ impl TryEncode for Subscribe<'_> { *offset += 1; } + assert_eq!(required_length, buffer[..*offset].len()); Ok(()) } } diff --git a/fan-controller/src/mqtt/task.rs b/fan-controller/src/mqtt/task.rs index 571745e..764fd66 100644 --- a/fan-controller/src/mqtt/task.rs +++ b/fan-controller/src/mqtt/task.rs @@ -3,6 +3,7 @@ use crate::mqtt::packet::connect_acknowledgement::{ConnectAcknowledgement, Conne use crate::mqtt::packet::GetPartsError; use crate::mqtt::{packet, ConnectErrorReasonCode, DecodeError}; use crate::mqtt::{TryDecode, TryEncode}; +use core::fmt::Debug; use defmt::{info, warn, Format}; use embassy_net::tcp; use embassy_net::tcp::TcpSocket; @@ -14,7 +15,7 @@ use super::packet::connect_acknowledgement; ///! - Keep alive #[derive(Debug, Format)] -pub(crate) enum SendError { +pub(crate) enum SendError { EncodeError(T), SendError(tcp::Error), FlushError(tcp::Error), @@ -25,10 +26,11 @@ pub(crate) async fn send<'a, T>( packet: T, ) -> Result<(), SendError<::Error>> where - T: TryEncode, + T: TryEncode, { + info!("Sending packet"); let mut offset = 0; - let mut send_buffer = [0; 256]; + let mut send_buffer = [0; 512]; packet .try_encode(&mut send_buffer, &mut offset) .map_err(SendError::EncodeError)?; @@ -81,8 +83,9 @@ pub(crate) async fn connect<'a, 'b>( info!("Connect acknowledgement packet received"); - let acknowledgement = ConnectAcknowledgement::try_decode(parts.variable_header_and_payload) - .map_err(ConnectError::DecodeAcknowledgementError)?; + let acknowledgement = + ConnectAcknowledgement::try_decode(parts.flags, parts.variable_header_and_payload) + .map_err(ConnectError::DecodeAcknowledgementError)?; info!("Connect acknowledgement read"); if let ConnectReasonCode::ErrorCode(error_code) = acknowledgement.connect_reason_code { diff --git a/fan-controller/src/mqtt/variable_byte_integer.rs b/fan-controller/src/mqtt/variable_byte_integer.rs index 0371c3e..dcda6e2 100644 --- a/fan-controller/src/mqtt/variable_byte_integer.rs +++ b/fan-controller/src/mqtt/variable_byte_integer.rs @@ -8,11 +8,12 @@ pub(super) enum VariableByteIntegerEncodeError { EndOfBuffer, } const MAX: usize = 268_435_455; +/// Returns bytes written on success pub(super) fn encode( value: usize, buffer: &mut [u8], offset: &mut usize, -) -> Result<(), VariableByteIntegerEncodeError> { +) -> Result { // This checks if the length is also too large for u32 let length = match value { 0..=127 => 1, @@ -30,7 +31,7 @@ pub(super) fn encode( if value < 128 { buffer[*offset] = value as u8; *offset += 1; - return Ok(()); + return Ok(1); } let mut value = value as u32; @@ -47,7 +48,7 @@ pub(super) fn encode( } } - Ok(()) + Ok(length) } #[derive(Debug, Format, Clone)]