use crate::Error;
use crate::codec::util::{
decode_byte, decode_string, decode_variable_integer, encode_string, encode_variable_integer,
};
use crate::codec::{Decode, Encode, RawPacket};
use crate::protocol::util::len_bytes;
use crate::protocol::v5::property::{
Property, PropertyFrame, property_decode, property_encode, property_len,
};
use crate::protocol::v5::util::id_header;
use crate::protocol::{FixedHeader, Flags, PacketType, QoS, traits, util};
use bytes::{Buf, BufMut, Bytes, BytesMut};
use std::borrow::Borrow;
use std::ops::{Index, IndexMut};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SubscribeProperties {
pub subscription_id: Option<u32>,
pub user_properties: Vec<(String, String)>,
}
impl PropertyFrame for SubscribeProperties {
fn encoded_len(&self) -> usize {
let mut len = 0usize;
if let Some(value) = self.subscription_id {
len += 1 + len_bytes(value as usize);
}
len += property_len!(&self.user_properties);
len
}
fn encode(&self, buf: &mut BytesMut) {
if let Some(value) = self.subscription_id {
buf.put_u8(Property::SubscriptionIdentifier.into());
encode_variable_integer(buf, value).expect("");
}
property_encode!(&self.user_properties, Property::UserProp, buf);
}
fn decode(buf: &mut Bytes) -> Result<Option<Self>, Error>
where
Self: Sized,
{
if buf.is_empty() {
return Ok(None);
}
let mut subscription_id: Option<u32> = None;
let mut user_properties: Vec<(String, String)> = Vec::new();
while buf.has_remaining() {
let property: Property = decode_byte(buf)?.try_into()?;
match property {
Property::SubscriptionIdentifier => {
if subscription_id.is_some() {
return Err(Error::ProtocolError);
}
let value = decode_variable_integer(buf)?;
buf.advance(len_bytes(value as usize));
subscription_id = Some(value);
}
Property::UserProp => {
property_decode!(&mut user_properties, buf);
}
_ => return Err(Error::PropertyMismatch),
}
}
Ok(Some(SubscribeProperties {
subscription_id,
user_properties,
}))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd)]
pub enum RetainHandling {
Send = 0,
SendForNewSub = 1,
DoNotSend = 2,
}
impl TryFrom<u8> for RetainHandling {
type Error = Error;
fn try_from(value: u8) -> Result<Self, Self::Error> {
match value {
0 => Ok(RetainHandling::Send),
1 => Ok(RetainHandling::SendForNewSub),
2 => Ok(RetainHandling::DoNotSend),
n => Err(Error::InvalidRetainHandling(n)),
}
}
}
impl From<RetainHandling> for u8 {
fn from(value: RetainHandling) -> Self {
value as u8
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TopicOptionFilter {
pub topic: String,
pub qos: QoS,
pub no_local: bool,
pub retain_as_published: bool,
pub retain_handling: RetainHandling,
}
impl TopicOptionFilter {
pub fn new<S: Into<String>>(
topic: S,
qos: QoS,
no_local: bool,
retain_as_published: bool,
retain_handling: RetainHandling,
) -> Self {
let topic = topic.into();
if !util::is_valid_topic_filter(&topic) {
panic!("Invalid topic filter: '{}'", topic);
}
TopicOptionFilter {
topic,
qos,
no_local,
retain_as_published,
retain_handling,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TopicOptionFilters(Vec<TopicOptionFilter>);
#[allow(clippy::len_without_is_empty)]
impl TopicOptionFilters {
pub fn new<T: IntoIterator<Item = TopicOptionFilter>>(filters: T) -> Self {
let values: Vec<TopicOptionFilter> = filters.into_iter().collect();
if values.is_empty() {
panic!("At least one topic filter is required");
}
TopicOptionFilters(values)
}
pub fn len(&self) -> usize {
self.0.len()
}
pub(crate) fn decode(payload: &mut Bytes) -> Result<Self, Error> {
let mut filters = Vec::with_capacity(1);
while payload.has_remaining() {
let topic = decode_string(payload)?;
if !util::is_valid_topic_filter(&topic) {
return Err(Error::InvalidTopicFilter(topic));
}
let flags = decode_byte(payload)?;
if flags & 0b1100_0000 > 0 {
return Err(Error::MalformedPacket);
}
let qos = (flags & 0x03).try_into()?;
let no_local = flags & 0x04 != 0;
let retain_as_published = flags & 0x08 != 0;
let retain_handling = ((flags >> 4) & 0x03).try_into()?;
filters.push(TopicOptionFilter::new(
topic,
qos,
no_local,
retain_as_published,
retain_handling,
));
}
if filters.is_empty() {
return Err(Error::NoTopic);
}
Ok(TopicOptionFilters(filters))
}
pub(crate) fn encode(&self, buf: &mut BytesMut) {
self.0.iter().for_each(|f| {
let qos: u8 = f.qos.into();
let retain_handling: u8 = f.retain_handling.into();
let options: u8 = retain_handling << 4
| (f.retain_as_published as u8) << 3
| (f.no_local as u8) << 2
| qos;
encode_string(buf, &f.topic);
buf.put_u8(options);
});
}
pub(crate) fn encoded_len(&self) -> usize {
self.0.iter().fold(0, |acc, f| acc + 2 + f.topic.len() + 1)
}
}
impl AsRef<Vec<TopicOptionFilter>> for TopicOptionFilters {
#[inline]
fn as_ref(&self) -> &Vec<TopicOptionFilter> {
&self.0
}
}
impl Borrow<Vec<TopicOptionFilter>> for TopicOptionFilters {
fn borrow(&self) -> &Vec<TopicOptionFilter> {
&self.0
}
}
impl IntoIterator for TopicOptionFilters {
type Item = TopicOptionFilter;
type IntoIter = std::vec::IntoIter<TopicOptionFilter>;
fn into_iter(self) -> Self::IntoIter {
self.0.into_iter()
}
}
impl FromIterator<TopicOptionFilter> for TopicOptionFilters {
fn from_iter<T: IntoIterator<Item = TopicOptionFilter>>(iter: T) -> Self {
TopicOptionFilters(Vec::from_iter(iter))
}
}
impl From<TopicOptionFilters> for Vec<TopicOptionFilter> {
#[inline]
fn from(value: TopicOptionFilters) -> Self {
value.0
}
}
impl From<Vec<TopicOptionFilter>> for TopicOptionFilters {
#[inline]
fn from(value: Vec<TopicOptionFilter>) -> Self {
TopicOptionFilters(value)
}
}
impl Index<usize> for TopicOptionFilters {
type Output = TopicOptionFilter;
fn index(&self, index: usize) -> &Self::Output {
self.0.index(index)
}
}
impl IndexMut<usize> for TopicOptionFilters {
fn index_mut(&mut self, index: usize) -> &mut Self::Output {
self.0.index_mut(index)
}
}
id_header!(SubscribeHeader, SubscribeProperties);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Subscribe {
header: SubscribeHeader,
filters: TopicOptionFilters,
}
impl Subscribe {
pub fn new<T: IntoIterator<Item = TopicOptionFilter>>(
packet_id: u16,
properties: Option<SubscribeProperties>,
filters: T,
) -> Self {
let header = SubscribeHeader::new(packet_id, properties);
let filters = TopicOptionFilters::new(filters);
Subscribe { header, filters }
}
pub fn packet_id(&self) -> u16 {
self.header.packet_id
}
pub fn properties(&self) -> Option<SubscribeProperties> {
self.header.properties.clone()
}
pub fn filters(&self) -> TopicOptionFilters {
self.filters.clone()
}
}
impl Encode for Subscribe {
fn encode(&self, buf: &mut BytesMut) -> Result<(), Error> {
let header = FixedHeader::with_flags(
PacketType::Subscribe,
Flags::new(QoS::AtLeastOnce),
self.payload_len(),
);
header.encode(buf)?;
self.header.encode(buf)?;
self.filters.encode(buf);
Ok(())
}
fn payload_len(&self) -> usize {
self.header.encoded_len() + self.filters.encoded_len()
}
}
impl Decode for Subscribe {
fn decode(mut packet: RawPacket) -> Result<Self, Error> {
if packet.header.packet_type() != PacketType::Subscribe
|| packet.header.flags() != Flags::new(QoS::AtLeastOnce)
{
return Err(Error::MalformedPacket);
}
let header = SubscribeHeader::decode(&mut packet.payload)?;
let filters = TopicOptionFilters::decode(&mut packet.payload)?;
Ok(Subscribe::new(header.packet_id, header.properties, filters))
}
}
impl traits::Subscribe for Subscribe {}
#[cfg(test)]
mod tests {
use super::*;
use crate::codec::PacketCodec;
use tokio_util::codec::Decoder;
#[test]
fn subscribe_properties_decode_advances_past_subscription_identifier() {
let mut buf = BytesMut::new();
buf.put_u8(Property::SubscriptionIdentifier.into());
encode_variable_integer(&mut buf, 42).unwrap();
buf.put_u8(Property::UserProp.into());
encode_string(&mut buf, "client");
encode_string(&mut buf, "rust");
let mut buf = buf.freeze();
let properties = SubscribeProperties::decode(&mut buf).unwrap().unwrap();
assert_eq!(properties.subscription_id, Some(42));
assert_eq!(
properties.user_properties,
vec![("client".to_string(), "rust".to_string())]
);
assert!(buf.is_empty(), "buffer should be fully consumed");
}
#[test]
fn subscribe_properties_decode_rejects_duplicate_subscription_identifier() {
let mut buf = BytesMut::new();
buf.put_u8(Property::SubscriptionIdentifier.into());
encode_variable_integer(&mut buf, 1).unwrap();
buf.put_u8(Property::SubscriptionIdentifier.into());
encode_variable_integer(&mut buf, 2).unwrap();
let mut buf = buf.freeze();
let result = SubscribeProperties::decode(&mut buf);
assert!(matches!(result, Err(Error::ProtocolError)));
}
#[test]
fn subscribe_decode_full_packet_with_subscription_identifier() {
let mut codec = PacketCodec::new(None, None);
let mut properties_buf = BytesMut::new();
properties_buf.put_u8(Property::SubscriptionIdentifier.into());
encode_variable_integer(&mut properties_buf, 7).unwrap();
let mut payload = BytesMut::new();
payload.put_u16(0x1234);
encode_variable_integer(&mut payload, properties_buf.len() as u32).unwrap();
payload.extend_from_slice(&properties_buf);
encode_string(&mut payload, "sensors/#");
payload.put_u8(0x00);
let mut stream = BytesMut::new();
stream.put_u8(((PacketType::Subscribe as u8) << 4) | 0x02);
encode_variable_integer(&mut stream, payload.len() as u32).unwrap();
stream.extend_from_slice(&payload);
let raw_packet = codec.decode(&mut stream).unwrap().unwrap();
let packet = Subscribe::decode(raw_packet).unwrap();
assert_eq!(packet.packet_id(), 0x1234);
assert_eq!(packet.properties().unwrap().subscription_id, Some(7));
assert_eq!(packet.filters().len(), 1);
assert_eq!(packet.filters()[0].topic, "sensors/#");
}
}