Fix disconnect when there are pending writes

Signed-off-by: vm-vloz <73361520+vm-vloz@users.noreply.github.com>
This commit is contained in:
vm-vloz 2021-07-26 18:51:32 -07:00 committed by vm-vloz
parent de1bc380f5
commit fadf34d1e5
7 changed files with 130 additions and 17 deletions

View File

@ -137,6 +137,7 @@ typedef struct
List* inboundMsgs; /**< inbound in flight messages */
List* outboundMsgs; /**< outbound in flight messages */
List* messageQueue; /**< inbound complete but undelivered messages */
List* outboundQueue; /**< outbound queued messages */
unsigned int qentry_seqno;
void* phandle; /**< the persistence handle */
MQTTClient_persistence* persistence; /**< a persistence implementation */

View File

@ -356,6 +356,7 @@ int MQTTAsync_createWithOptions(MQTTAsync* handle, const char* serverURI, const
bstate->clients = ListInitialize();
Socket_outInitialize();
Socket_setWriteCompleteCallback(MQTTAsync_writeComplete);
Socket_setWriteAvailableCallback(MQTTProtocol_writeAvailable);
MQTTAsync_handles = ListInitialize();
MQTTAsync_commands = ListInitialize();
#if defined(OPENSSL)
@ -408,9 +409,10 @@ int MQTTAsync_createWithOptions(MQTTAsync* handle, const char* serverURI, const
m->c->outboundMsgs = ListInitialize();
m->c->inboundMsgs = ListInitialize();
m->c->messageQueue = ListInitialize();
m->c->outboundQueue = ListInitialize();
m->c->clientID = MQTTStrdup(clientId);
if (m->c->context == NULL || m->c->outboundMsgs == NULL || m->c->inboundMsgs == NULL ||
m->c->messageQueue == NULL || m->c->clientID == NULL)
m->c->messageQueue == NULL || m->c->outboundQueue == NULL || m->c->clientID == NULL)
{
rc = PAHO_MEMORY_ERROR;
goto exit;

View File

@ -428,6 +428,7 @@ int MQTTClient_createWithOptions(MQTTClient* handle, const char* serverURI, cons
bstate->clients = ListInitialize();
Socket_outInitialize();
Socket_setWriteCompleteCallback(MQTTClient_writeComplete);
Socket_setWriteAvailableCallback(MQTTProtocol_writeAvailable);
handles = ListInitialize();
#if defined(OPENSSL)
SSLSocket_initialize();
@ -486,6 +487,7 @@ int MQTTClient_createWithOptions(MQTTClient* handle, const char* serverURI, cons
m->c->outboundMsgs = ListInitialize();
m->c->inboundMsgs = ListInitialize();
m->c->messageQueue = ListInitialize();
m->c->outboundQueue = ListInitialize();
m->c->clientID = MQTTStrdup(clientId);
m->connect_sem = Thread_create_sem(&rc);
m->connack_sem = Thread_create_sem(&rc);

View File

@ -36,6 +36,7 @@
#if !defined(NO_PERSISTENCE)
#include "MQTTPersistence.h"
#endif
#include "Socket.h"
#include "SocketBuffer.h"
#include "StackTrace.h"
#include "Heap.h"
@ -55,6 +56,13 @@ static int MQTTProtocol_startPublishCommon(
int retained);
static void MQTTProtocol_retries(START_TIME_TYPE now, Clients* client, int regardless);
static int MQTTProtocol_queueAck(Clients* client, int ackType, int msgId);
typedef struct {
int messageId;
int ackType;
} AckRequest;
/**
* List callback function for comparing Message structures by message id
@ -312,6 +320,7 @@ int MQTTProtocol_handlePublishes(void* pack, int sock)
Clients* client = NULL;
char* clientid = NULL;
int rc = TCPSOCKET_COMPLETE;
int socketHasPendingWrites = 0;
FUNC_ENTRY;
client = (Clients*)(ListFindItem(bstate->clients, &sock, clientSocketCompare)->content);
@ -320,15 +329,23 @@ int MQTTProtocol_handlePublishes(void* pack, int sock)
publish->header.bits.retain, publish->payloadlen, min(20, publish->payloadlen), publish->payload);
if (publish->header.bits.qos == 0)
Protocol_processPublication(publish, client, 1);
else if (!Socket_noPendingWrites(sock))
rc = SOCKET_ERROR; /* queue acks? */
else if (publish->header.bits.qos == 1)
{
/* send puback before processing the publications because a lot of return publications could fill up the socket buffer */
rc = MQTTPacket_send_puback(publish->MQTTVersion, publish->msgId, &client->net, client->clientID);
/* if we get a socket error from sending the puback, should we ignore the publication? */
Protocol_processPublication(publish, client, 1);
Protocol_processPublication(publish, client, 1);
goto exit;
}
socketHasPendingWrites = !Socket_noPendingWrites(sock);
if (publish->header.bits.qos == 1)
{
if (socketHasPendingWrites)
rc = MQTTProtocol_queueAck(client, PUBACK, publish->msgId);
else
/* send puback before processing the publications because a lot of return publications could fill up the socket buffer */
rc = MQTTPacket_send_puback(publish->MQTTVersion, publish->msgId, &client->net, client->clientID);
/* if we get a socket error from sending the puback, should we ignore the publication? */
Protocol_processPublication(publish, client, 1);
}
else if (publish->header.bits.qos == 2)
{
@ -364,7 +381,12 @@ int MQTTProtocol_handlePublishes(void* pack, int sock)
already_received = 1;
} else
ListAppend(client->inboundMsgs, m, sizeof(Messages) + len);
rc = MQTTPacket_send_pubrec(publish->MQTTVersion, publish->msgId, &client->net, client->clientID);
if (socketHasPendingWrites)
rc = MQTTProtocol_queueAck(client, PUBREC, publish->msgId);
else
rc = MQTTPacket_send_pubrec(publish->MQTTVersion, publish->msgId, &client->net, client->clientID);
if (m->MQTTVersion >= MQTTVERSION_5 && already_received == 0)
{
Publish publish1;
@ -539,7 +561,7 @@ int MQTTProtocol_handlePubrels(void* pack, int sock)
if (pubrel->header.bits.dup == 0)
Log(TRACE_MIN, 3, NULL, "PUBREL", client->clientID, pubrel->msgId);
else if (!Socket_noPendingWrites(sock))
rc = SOCKET_ERROR; /* queue acks? */
rc = MQTTProtocol_queueAck(client, PUBCOMP, pubrel->msgId);
else
/* Apparently this is "normal" behaviour, so we don't need to issue a warning */
rc = MQTTPacket_send_pubcomp(pubrel->MQTTVersion, pubrel->msgId, &client->net, client->clientID);
@ -551,15 +573,16 @@ int MQTTProtocol_handlePubrels(void* pack, int sock)
Log(TRACE_MIN, 4, NULL, "PUBREL", client->clientID, pubrel->msgId, m->qos);
else if (m->nextMessageType != PUBREL)
Log(TRACE_MIN, 5, NULL, "PUBREL", client->clientID, pubrel->msgId);
else if (!Socket_noPendingWrites(sock))
rc = SOCKET_ERROR; /* queue acks? */
else
{
Publish publish;
memset(&publish, '\0', sizeof(publish));
/* send pubcomp before processing the publications because a lot of return publications could fill up the socket buffer */
rc = MQTTPacket_send_pubcomp(pubrel->MQTTVersion, pubrel->msgId, &client->net, client->clientID);
if (!Socket_noPendingWrites(sock))
rc = MQTTProtocol_queueAck(client, PUBCOMP, pubrel->msgId);
else
rc = MQTTPacket_send_pubcomp(pubrel->MQTTVersion, pubrel->msgId, &client->net, client->clientID);
publish.header.bits.qos = m->qos;
publish.header.bits.retain = m->retain;
publish.msgId = m->msgid;
@ -576,9 +599,9 @@ int MQTTProtocol_handlePubrels(void* pack, int sock)
else
Protocol_processPublication(&publish, client, 0); /* only for 3.1.1 and lower */
#if !defined(NO_PERSISTENCE)
rc += MQTTPersistence_remove(client,
(m->MQTTVersion >= MQTTVERSION_5) ? PERSISTENCE_V5_PUBLISH_RECEIVED : PERSISTENCE_PUBLISH_RECEIVED,
m->qos, pubrel->msgId);
rc += MQTTPersistence_remove(client,
(m->MQTTVersion >= MQTTVERSION_5) ? PERSISTENCE_V5_PUBLISH_RECEIVED : PERSISTENCE_PUBLISH_RECEIVED,
m->qos, pubrel->msgId);
#endif
if (m->MQTTVersion >= MQTTVERSION_5)
MQTTProperties_free(&m->properties);
@ -791,6 +814,35 @@ exit:
}
/**
* Queue an ack message. This is used when the socket is full (e.g. SSL_ERROR_WANT_WRITE).
* To be completed/cleared when the socket is no longer full
* @param client the client that received the published message
* @param ackType the type of ack to send
* @param msgId the msg id of the message we are acknowledging
* @return the completion code
*/
int MQTTProtocol_queueAck(Clients* client, int ackType, int msgId)
{
int rc = 0;
AckRequest* ackReq = NULL;
FUNC_ENTRY;
ackReq = malloc(sizeof(AckRequest));
if (!ackReq)
rc = PAHO_MEMORY_ERROR;
else
{
ackReq->messageId = msgId;
ackReq->ackType = ackType;
ListAppend(client->outboundQueue, ackReq, sizeof(AckRequest));
}
FUNC_EXIT_RC(rc);
return rc;
}
/**
* MQTT retry protocol and socket pending writes processing.
* @param now current time
@ -835,6 +887,7 @@ void MQTTProtocol_freeClient(Clients* client)
MQTTProtocol_freeMessageList(client->outboundMsgs);
MQTTProtocol_freeMessageList(client->inboundMsgs);
ListFree(client->messageQueue);
ListFree(client->outboundQueue);
free(client->clientID);
client->clientID = NULL;
if (client->will)
@ -917,6 +970,48 @@ void MQTTProtocol_freeMessageList(List* msgList)
}
/**
* Callback that is invoked when the socket is available for writing.
* This is the last attempt made to acknowledge a message. Failures that
* occur here are ignored.
* @param socket the socket that is available for writing
*/
void MQTTProtocol_writeAvailable(int socket)
{
Clients* client = NULL;
ListElement* current = NULL;
int rc = 0;
FUNC_ENTRY;
client = (Clients*)(ListFindItem(bstate->clients, &socket, clientSocketCompare)->content);
current = NULL;
while (ListNextElement(client->outboundQueue, &current) && rc == 0)
{
AckRequest* ackReq = (AckRequest*)(current->content);
switch (ackReq->ackType)
{
case PUBACK:
rc = MQTTPacket_send_puback(client->MQTTVersion, ackReq->messageId, &client->net, client->clientID);
break;
case PUBREC:
rc = MQTTPacket_send_pubrec(client->MQTTVersion, ackReq->messageId, &client->net, client->clientID);
break;
case PUBCOMP:
rc = MQTTPacket_send_pubcomp(client->MQTTVersion, ackReq->messageId, &client->net, client->clientID);
break;
default:
Log(LOG_ERROR, -1, "unknown ACK type %d, dropping msg", ackReq->ackType);
break;
}
}
ListEmpty(client->outboundQueue);
FUNC_EXIT_RC(rc);
}
/**
* Copy no more than dest_size -1 characters from the string pointed to by src to the array pointed to by dest.
* The destination string will always be null-terminated.

View File

@ -55,6 +55,8 @@ void MQTTProtocol_freeMessageList(List* msgList);
char* MQTTStrncpy(char *dest, const char* src, size_t num);
char* MQTTStrdup(const char* src);
void MQTTProtocol_writeAvailable(int socket);
//#define MQTTStrdup(src) MQTTStrncpy(malloc(strlen(src)+1), src, strlen(src)+1)
#endif

View File

@ -861,7 +861,12 @@ void Socket_setWriteCompleteCallback(Socket_writeComplete* mywritecomplete)
writecomplete = mywritecomplete;
}
static Socket_writeAvailable* writeAvailable = NULL;
void Socket_setWriteAvailableCallback(Socket_writeAvailable* mywriteavailable)
{
writeAvailable = mywriteavailable;
}
/**
* Continue an outstanding write for a particular socket
@ -1010,6 +1015,9 @@ int Socket_continueWrites(fd_set* pwset, int* sock)
}
curpending = mod_s.write_pending->current;
if (writeAvailable && rc > 0)
(*writeAvailable)(socket);
if (writecomplete)
(*writecomplete)(socket, rc);
}

View File

@ -142,4 +142,7 @@ void Socket_clearPendingWrite(int socket);
typedef void Socket_writeComplete(int socket, int rc);
void Socket_setWriteCompleteCallback(Socket_writeComplete*);
typedef void Socket_writeAvailable(int socket);
void Socket_setWriteAvailableCallback(Socket_writeAvailable*);
#endif /* SOCKET_H */