Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 47 additions & 17 deletions rumqttc/src/client.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
//! This module offers a high level synchronous and asynchronous abstraction to
//! async eventloop.
use std::sync::atomic::{AtomicU16, Ordering};
use std::sync::Arc;
use std::time::Duration;

use crate::mqttbytes::{v4::*, QoS};
Expand Down Expand Up @@ -42,17 +44,25 @@ impl From<TrySendError<Request>> for ClientError {
#[derive(Clone, Debug)]
pub struct AsyncClient {
request_tx: Sender<Request>,
pkid_counter: Arc<AtomicU16>,
max_inflight: u16,
}

impl AsyncClient {
/// Create a new `AsyncClient`.
///
/// `cap` specifies the capacity of the bounded async channel.
pub fn new(options: MqttOptions, cap: usize) -> (AsyncClient, EventLoop) {
let max_inflight = options.inflight;
let eventloop = EventLoop::new(options, cap);
let request_tx = eventloop.requests_tx.clone();
let pkid_counter = eventloop.pkid_counter();

let client = AsyncClient { request_tx };
let client = AsyncClient {
request_tx,
pkid_counter,
max_inflight,
};

(client, eventloop)
}
Expand All @@ -61,8 +71,24 @@ impl AsyncClient {
///
/// This is mostly useful for creating a test instance where you can
/// listen on the corresponding receiver.
pub fn from_senders(request_tx: Sender<Request>) -> AsyncClient {
AsyncClient { request_tx }
pub fn from_senders(request_tx: Sender<Request>, pkid_counter: Arc<AtomicU16>, max_inflight: u16) -> AsyncClient {
AsyncClient {
request_tx,
pkid_counter,
max_inflight,
}
}

/// Assigns a packet id for QoS 1/2 publishes using the shared atomic counter.
fn assign_pkid(&self, publish: &mut Publish) -> u16 {
if publish.qos != QoS::AtMostOnce {
let raw = self.pkid_counter.fetch_add(1, Ordering::Relaxed);
let pkid = (raw % self.max_inflight) + 1;
publish.pkid = pkid;
pkid
} else {
0
}
}

/// Sends a MQTT Publish to the `EventLoop`.
Expand All @@ -72,20 +98,21 @@ impl AsyncClient {
qos: QoS,
retain: bool,
payload: V,
) -> Result<(), ClientError>
) -> Result<u16, ClientError>
where
S: Into<String>,
V: Into<Vec<u8>>,
{
let topic = topic.into();
let mut publish = Publish::new(&topic, qos, payload);
publish.retain = retain;
let pkid = self.assign_pkid(&mut publish);
let publish = Request::Publish(publish);
if !valid_topic(&topic) {
return Err(ClientError::Request(publish));
}
self.request_tx.send_async(publish).await?;
Ok(())
Ok(pkid)
}

/// Attempts to send a MQTT Publish to the `EventLoop`.
Expand All @@ -95,20 +122,21 @@ impl AsyncClient {
qos: QoS,
retain: bool,
payload: V,
) -> Result<(), ClientError>
) -> Result<u16, ClientError>
where
S: Into<String>,
V: Into<Vec<u8>>,
{
let topic = topic.into();
let mut publish = Publish::new(&topic, qos, payload);
publish.retain = retain;
let pkid = self.assign_pkid(&mut publish);
let publish = Request::Publish(publish);
if !valid_topic(&topic) {
return Err(ClientError::TryRequest(publish));
}
self.request_tx.try_send(publish)?;
Ok(())
Ok(pkid)
}

/// Sends a MQTT PubAck to the `EventLoop`. Only needed in if `manual_acks` flag is set.
Expand Down Expand Up @@ -137,15 +165,16 @@ impl AsyncClient {
qos: QoS,
retain: bool,
payload: Bytes,
) -> Result<(), ClientError>
) -> Result<u16, ClientError>
where
S: Into<String>,
{
let mut publish = Publish::from_bytes(topic, qos, payload);
publish.retain = retain;
let pkid = self.assign_pkid(&mut publish);
let publish = Request::Publish(publish);
self.request_tx.send_async(publish).await?;
Ok(())
Ok(pkid)
}

/// Sends a MQTT Subscribe to the `EventLoop`
Expand Down Expand Up @@ -272,9 +301,9 @@ impl Client {
///
/// This is mostly useful for creating a test instance where you can
/// listen on the corresponding receiver.
pub fn from_sender(request_tx: Sender<Request>) -> Client {
pub fn from_sender(request_tx: Sender<Request>, pkid_counter: Arc<AtomicU16>, max_inflight: u16) -> Client {
Client {
client: AsyncClient::from_senders(request_tx),
client: AsyncClient::from_senders(request_tx, pkid_counter, max_inflight),
}
}

Expand All @@ -285,20 +314,21 @@ impl Client {
qos: QoS,
retain: bool,
payload: V,
) -> Result<(), ClientError>
) -> Result<u16, ClientError>
where
S: Into<String>,
V: Into<Vec<u8>>,
{
let topic = topic.into();
let mut publish = Publish::new(&topic, qos, payload);
publish.retain = retain;
let pkid = self.client.assign_pkid(&mut publish);
let publish = Request::Publish(publish);
if !valid_topic(&topic) {
return Err(ClientError::Request(publish));
}
self.client.request_tx.send(publish)?;
Ok(())
Ok(pkid)
}

pub fn try_publish<S, V>(
Expand All @@ -307,13 +337,12 @@ impl Client {
qos: QoS,
retain: bool,
payload: V,
) -> Result<(), ClientError>
) -> Result<u16, ClientError>
where
S: Into<String>,
V: Into<Vec<u8>>,
{
self.client.try_publish(topic, qos, retain, payload)?;
Ok(())
self.client.try_publish(topic, qos, retain, payload)
}

/// Sends a MQTT PubAck to the `EventLoop`. Only needed in if `manual_acks` flag is set.
Expand Down Expand Up @@ -540,7 +569,8 @@ mod test {
#[test]
fn should_be_able_to_build_test_client_from_channel() {
let (tx, rx) = flume::bounded(1);
let client = Client::from_sender(tx);
let pkid_counter = Arc::new(AtomicU16::new(0));
let client = Client::from_sender(tx, pkid_counter, 100);
client
.publish("hello/world", QoS::ExactlyOnce, false, "good bye")
.expect("Should be able to publish");
Expand Down
13 changes: 12 additions & 1 deletion rumqttc/src/eventloop.rs
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@ use std::collections::VecDeque;
use std::io;
use std::net::SocketAddr;
use std::pin::Pin;
use std::sync::atomic::AtomicU16;
use std::sync::Arc;
use std::time::Duration;

#[cfg(unix)]
Expand Down Expand Up @@ -85,6 +87,8 @@ pub struct EventLoop {
/// Keep alive time
keepalive_timeout: Option<Pin<Box<Sleep>>>,
pub network_options: NetworkOptions,
/// Shared atomic counter for packet id assignment
pub(crate) pkid_counter: Arc<AtomicU16>,
}

/// Events which can be yielded by the event loop
Expand All @@ -104,19 +108,26 @@ impl EventLoop {
let pending = VecDeque::new();
let max_inflight = mqtt_options.inflight;
let manual_acks = mqtt_options.manual_acks;
let pkid_counter = Arc::new(AtomicU16::new(0));

EventLoop {
mqtt_options,
state: MqttState::new(max_inflight, manual_acks),
state: MqttState::new(max_inflight, manual_acks, pkid_counter.clone()),
requests_tx,
requests_rx,
pending,
network: None,
keepalive_timeout: None,
network_options: NetworkOptions::new(),
pkid_counter,
}
}

/// Returns a clone of the shared packet id counter for use by the client
pub fn pkid_counter(&self) -> Arc<AtomicU16> {
self.pkid_counter.clone()
}

/// Last session might contain packets which aren't acked. MQTT says these packets should be
/// republished in the next session. Move pending messages from state to eventloop, drops the
/// underlying network connection and clears the keepalive timeout if any.
Expand Down
54 changes: 37 additions & 17 deletions rumqttc/src/state.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,9 @@ use crate::mqttbytes::v4::*;
use crate::mqttbytes::{self, *};
use fixedbitset::FixedBitSet;
use std::collections::VecDeque;
use std::fmt::{self, Debug, Formatter};
use std::sync::atomic::{AtomicU16, Ordering};
use std::sync::Arc;
use std::{io, time::Instant};

/// Errors during state handling
Expand Down Expand Up @@ -40,7 +43,6 @@ pub enum StateError {
// This is done for 2 reasons
// Bad acks or out of order acks aren't O(n) causing cpu spikes
// Any missing acks from the broker are detected during the next recycled use of packet ids
#[derive(Debug, Clone)]
pub struct MqttState {
/// Status of last ping
pub await_pingresp: bool,
Expand Down Expand Up @@ -72,19 +74,40 @@ pub struct MqttState {
pub events: VecDeque<Event>,
/// Indicates if acknowledgements should be send immediately
pub manual_acks: bool,
/// Shared atomic counter for packet id assignment
pub(crate) pkid_counter: Arc<AtomicU16>,
}

impl Debug for MqttState {
fn fmt(&self, f: &mut Formatter) -> fmt::Result {
f.debug_struct("MqttState")
.field("await_pingresp", &self.await_pingresp)
.field("collision_ping_count", &self.collision_ping_count)
.field("last_pkid", &self.last_pkid)
.field("last_puback", &self.last_puback)
.field("inflight", &self.inflight)
.field("max_inflight", &self.max_inflight)
.field("outgoing_pub", &self.outgoing_pub)
.field("outgoing_rel", &self.outgoing_rel)
.field("collision", &self.collision)
.field("events", &self.events)
.field("manual_acks", &self.manual_acks)
.finish()
}
}

impl MqttState {
/// Creates new mqtt state. Same state should be used during a
/// connection for persistent sessions while new state should
/// instantiated for clean sessions
pub fn new(max_inflight: u16, manual_acks: bool) -> Self {
pub fn new(max_inflight: u16, manual_acks: bool, pkid_counter: Arc<AtomicU16>) -> Self {
let last_pkid = pkid_counter.load(Ordering::Relaxed);
MqttState {
await_pingresp: false,
collision_ping_count: 0,
last_incoming: Instant::now(),
last_outgoing: Instant::now(),
last_pkid: 0,
last_pkid,
last_puback: 0,
inflight: 0,
max_inflight,
Expand All @@ -96,6 +119,7 @@ impl MqttState {
// TODO: Optimize these sizes later
events: VecDeque::with_capacity(100),
manual_acks,
pkid_counter,
}
}

Expand Down Expand Up @@ -312,6 +336,9 @@ impl MqttState {
if publish.qos != QoS::AtMostOnce {
if publish.pkid == 0 {
publish.pkid = self.next_pkid();
} else {
// Client pre-assigned pkid via shared atomic counter
self.last_pkid = publish.pkid;
}

let pkid = publish.pkid;
Expand Down Expand Up @@ -484,19 +511,10 @@ impl MqttState {
/// Packet ids are incremented till maximum set inflight messages and reset to 1 after that.
///
fn next_pkid(&mut self) -> u16 {
let next_pkid = self.last_pkid + 1;

// When next packet id is at the edge of inflight queue,
// set await flag. This instructs eventloop to stop
// processing requests until all the inflight publishes
// are acked
if next_pkid == self.max_inflight {
self.last_pkid = 0;
return next_pkid;
}

self.last_pkid = next_pkid;
next_pkid
let raw = self.pkid_counter.fetch_add(1, Ordering::Relaxed);
let pkid = (raw % self.max_inflight) + 1;
self.last_pkid = pkid;
pkid
}
}

Expand All @@ -506,6 +524,8 @@ mod test {
use crate::mqttbytes::v4::*;
use crate::mqttbytes::*;
use crate::{Event, Incoming, Outgoing, Request};
use std::sync::atomic::AtomicU16;
use std::sync::Arc;

fn build_outgoing_publish(qos: QoS) -> Publish {
let topic = "hello/world".to_owned();
Expand All @@ -527,7 +547,7 @@ mod test {
}

fn build_mqttstate() -> MqttState {
MqttState::new(100, false)
MqttState::new(100, false, Arc::new(AtomicU16::new(0)))
}

#[test]
Expand Down
Loading