[message] add range-based for loop for message/priority queue (#7481)

This commit adds support for range-based `for` loop iteration over
`MessageQueue` and `PriorityQueue`. It adds `ConstIterator` and
`Iterator` in `Message` class. The non-const `for` loop iteration
(which uses `Message::Iterator`) works properly and is safe to use
even when the entry is removed from the queue during iteration. This
commit updates other core modules to use the new `for` loop iteration
model which helps simplify the code. It also updates the unit tests
to validate the behavior of the newly added iteration mechanisms.
This commit is contained in:
Abtin Keshavarzian
2022-03-21 09:28:12 -07:00
committed by GitHub
parent c24e5dd6c4
commit 237c91b939
16 changed files with 503 additions and 250 deletions
+32 -40
View File
@@ -79,18 +79,15 @@ void CoapBase::ClearRequests(const Ip6::Address &aAddress)
void CoapBase::ClearRequests(const Ip6::Address *aAddress)
{
Message *nextMessage;
for (Message *message = mPendingRequests.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mPendingRequests)
{
Metadata metadata;
nextMessage = message->GetNextCoapMessage();
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if ((aAddress == nullptr) || (metadata.mSourceAddress == *aAddress))
{
FinalizeCoapTransaction(*message, metadata, nullptr, nullptr, kErrorAbort);
FinalizeCoapTransaction(message, metadata, nullptr, nullptr, kErrorAbort);
}
}
}
@@ -421,19 +418,16 @@ void CoapBase::HandleRetransmissionTimer(void)
TimeMilli now = TimerMilli::GetNow();
TimeMilli nextTime = now.GetDistantFuture();
Metadata metadata;
Message * nextMessage;
Ip6::MessageInfo messageInfo;
for (Message *message = mPendingRequests.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mPendingRequests)
{
nextMessage = message->GetNextCoapMessage();
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if (now >= metadata.mNextTimerShot)
{
#if OPENTHREAD_CONFIG_COAP_OBSERVE_API_ENABLE
if (message->IsRequest() && metadata.mObserve && metadata.mAcknowledged)
if (message.IsRequest() && metadata.mObserve && metadata.mAcknowledged)
{
// This is a RFC7641 subscription. Do not time out.
continue;
@@ -443,7 +437,7 @@ void CoapBase::HandleRetransmissionTimer(void)
if (!metadata.mConfirmable || (metadata.mRetransmissionsRemaining == 0))
{
// No expected response or acknowledgment.
FinalizeCoapTransaction(*message, metadata, nullptr, nullptr, kErrorResponseTimeout);
FinalizeCoapTransaction(message, metadata, nullptr, nullptr, kErrorResponseTimeout);
continue;
}
@@ -451,7 +445,7 @@ void CoapBase::HandleRetransmissionTimer(void)
metadata.mRetransmissionsRemaining--;
metadata.mRetransmissionTimeout *= 2;
metadata.mNextTimerShot = now + metadata.mRetransmissionTimeout;
metadata.UpdateIn(*message);
metadata.UpdateIn(message);
// Retransmit
if (!metadata.mAcknowledged)
@@ -465,7 +459,7 @@ void CoapBase::HandleRetransmissionTimer(void)
#endif
messageInfo.SetMulticastLoop(metadata.mMulticastLoop);
SendCopy(*message, messageInfo);
SendCopy(message, messageInfo);
}
}
@@ -498,17 +492,15 @@ void CoapBase::FinalizeCoapTransaction(Message & aRequest,
Error CoapBase::AbortTransaction(ResponseHandler aHandler, void *aContext)
{
Error error = kErrorNotFound;
Message *nextMessage;
Metadata metadata;
for (Message *message = mPendingRequests.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mPendingRequests)
{
nextMessage = message->GetNextCoapMessage();
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if (metadata.mResponseHandler == aHandler && metadata.mResponseContext == aContext)
{
FinalizeCoapTransaction(*message, metadata, nullptr, nullptr, kErrorAbort);
FinalizeCoapTransaction(message, metadata, nullptr, nullptr, kErrorAbort);
error = kErrorNone;
}
}
@@ -965,11 +957,11 @@ Message *CoapBase::FindRelatedRequest(const Message & aResponse,
const Ip6::MessageInfo &aMessageInfo,
Metadata & aMetadata)
{
Message *message;
Message *request = nullptr;
for (message = mPendingRequests.GetHead(); message != nullptr; message = message->GetNextCoapMessage())
for (Message &message : mPendingRequests)
{
aMetadata.ReadFrom(*message);
aMetadata.ReadFrom(message);
if (((aMetadata.mDestinationAddress == aMessageInfo.GetPeerAddr()) ||
aMetadata.mDestinationAddress.IsMulticast() ||
@@ -980,8 +972,9 @@ Message *CoapBase::FindRelatedRequest(const Message & aResponse,
{
case kTypeReset:
case kTypeAck:
if (aResponse.GetMessageId() == message->GetMessageId())
if (aResponse.GetMessageId() == message.GetMessageId())
{
request = &message;
ExitNow();
}
@@ -989,8 +982,9 @@ Message *CoapBase::FindRelatedRequest(const Message & aResponse,
case kTypeConfirmable:
case kTypeNonConfirmable:
if (aResponse.IsTokenEqual(*message))
if (aResponse.IsTokenEqual(message))
{
request = &message;
ExitNow();
}
@@ -1000,7 +994,7 @@ Message *CoapBase::FindRelatedRequest(const Message & aResponse,
}
exit:
return message;
return request;
}
void CoapBase::Receive(ot::Message &aMessage, const Ip6::MessageInfo &aMessageInfo)
@@ -1449,25 +1443,26 @@ exit:
const Message *ResponsesQueue::FindMatchedResponse(const Message &aRequest, const Ip6::MessageInfo &aMessageInfo) const
{
Message *message;
const Message *response = nullptr;
for (message = mQueue.GetHead(); message != nullptr; message = message->GetNextCoapMessage())
for (const Message &message : mQueue)
{
if (message->GetMessageId() == aRequest.GetMessageId())
if (message.GetMessageId() == aRequest.GetMessageId())
{
ResponseMetadata metadata;
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if ((metadata.mMessageInfo.GetPeerPort() == aMessageInfo.GetPeerPort()) &&
(metadata.mMessageInfo.GetPeerAddr() == aMessageInfo.GetPeerAddr()))
{
response = &message;
break;
}
}
}
return message;
return response;
}
void ResponsesQueue::EnqueueResponse(Message & aMessage,
@@ -1506,15 +1501,15 @@ void ResponsesQueue::UpdateQueue(void)
// `kMaxCachedResponses` remove the one with earliest dequeue
// time.
for (Message *message = mQueue.GetHead(); message != nullptr; message = message->GetNextCoapMessage())
for (Message &message : mQueue)
{
ResponseMetadata metadata;
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if ((earliestMsg == nullptr) || (metadata.mDequeueTime < earliestDequeueTime))
{
earliestMsg = message;
earliestMsg = &message;
earliestDequeueTime = metadata.mDequeueTime;
}
@@ -1546,19 +1541,16 @@ void ResponsesQueue::HandleTimer(void)
{
TimeMilli now = TimerMilli::GetNow();
TimeMilli nextDequeueTime = now.GetDistantFuture();
Message * nextMessage;
for (Message *message = mQueue.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mQueue)
{
ResponseMetadata metadata;
nextMessage = message->GetNextCoapMessage();
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if (now >= metadata.mDequeueTime)
{
DequeueResponse(*message);
DequeueResponse(message);
continue;
}
+10
View File
@@ -476,6 +476,16 @@ const char *Message::CodeToString(void) const
}
#endif // OPENTHREAD_CONFIG_COAP_API_ENABLE
Message::Iterator MessageQueue::begin(void)
{
return Message::Iterator(GetHead());
}
Message::ConstIterator MessageQueue::begin(void) const
{
return Message::ConstIterator(GetHead());
}
Error Option::Iterator::Init(const Message &aMessage)
{
Error error = kErrorParse;
+42 -1
View File
@@ -166,6 +166,7 @@ enum OptionNumber : uint16_t
class Message : public ot::Message
{
friend class Option;
friend class MessageQueue;
public:
static constexpr uint8_t kDefaultTokenLength = OT_COAP_DEFAULT_TOKEN_LENGTH; ///< Default token length.
@@ -956,6 +957,27 @@ private:
#endif
};
class ConstIterator : public ot::Message::ConstIterator
{
public:
using ot::Message::ConstIterator::ConstIterator;
const Message &operator*(void) { return static_cast<const Message &>(ot::Message::ConstIterator::operator*()); }
const Message *operator->(void)
{
return static_cast<const Message *>(ot::Message::ConstIterator::operator->());
}
};
class Iterator : public ot::Message::Iterator
{
public:
using ot::Message::Iterator::Iterator;
Message &operator*(void) { return static_cast<Message &>(ot::Message::Iterator::operator*()); }
Message *operator->(void) { return static_cast<Message *>(ot::Message::Iterator::operator->()); }
};
static_assert(sizeof(HelpData) <= sizeof(Ip6::Header) + sizeof(Ip6::HopByHopHeader) + sizeof(Ip6::OptionMpl) +
sizeof(Ip6::Udp::Header),
"HelpData size exceeds the size of the reserved region in the message");
@@ -1000,7 +1022,15 @@ public:
* @returns A pointer to the first message.
*
*/
Message *GetHead(void) const { return static_cast<Message *>(ot::MessageQueue::GetHead()); }
Message *GetHead(void) { return static_cast<Message *>(ot::MessageQueue::GetHead()); }
/**
* This method returns a pointer to the first message.
*
* @returns A pointer to the first message.
*
*/
const Message *GetHead(void) const { return static_cast<const Message *>(ot::MessageQueue::GetHead()); }
/**
* This method adds a message to the end of the queue.
@@ -1034,6 +1064,17 @@ public:
*
*/
void DequeueAndFree(Message &aMessage) { ot::MessageQueue::DequeueAndFree(aMessage); }
// The following methods are intended to support range-based `for`
// loop iteration over the queue entries and should not be used
// directly. The range-based `for` works correctly even if the
// current entry is removed from the queue during iteration.
Message::Iterator begin(void);
Message::Iterator end(void) { return Message::Iterator(); }
Message::ConstIterator begin(void) const;
Message::ConstIterator end(void) const { return Message::ConstIterator(); }
};
/**
+42 -35
View File
@@ -198,6 +198,15 @@ const Message::Settings &Message::Settings::From(const otMessageSettings *aSetti
return (aSettings == nullptr) ? GetDefault() : AsCoreType(aSettings);
}
//---------------------------------------------------------------------------------------------------------------------
// Message::Iterator
void Message::Iterator::Advance(void)
{
mItem = mNext;
mNext = NextMessage(mNext);
}
//---------------------------------------------------------------------------------------------------------------------
// Message
@@ -740,16 +749,6 @@ void Message::SetPriorityQueue(PriorityQueue *aPriorityQueue)
//---------------------------------------------------------------------------------------------------------------------
// MessageQueue
MessageQueue::MessageQueue(void)
{
SetTail(nullptr);
}
Message *MessageQueue::GetHead(void) const
{
return (GetTail() == nullptr) ? nullptr : GetTail()->Next();
}
void MessageQueue::Enqueue(Message &aMessage, QueuePosition aPosition)
{
OT_ASSERT(!aMessage.IsInAQueue());
@@ -821,37 +820,39 @@ void MessageQueue::DequeueAndFreeAll(void)
}
}
Message::Iterator MessageQueue::begin(void)
{
return Message::Iterator(GetHead());
}
Message::ConstIterator MessageQueue::begin(void) const
{
return Message::ConstIterator(GetHead());
}
void MessageQueue::GetInfo(uint16_t &aMessageCount, uint16_t &aBufferCount) const
{
aMessageCount = 0;
aBufferCount = 0;
for (const Message *message = GetHead(); message != nullptr; message = message->GetNext())
for (const Message &message : *this)
{
aMessageCount++;
aBufferCount += message->GetBufferCount();
aBufferCount += message.GetBufferCount();
}
}
//---------------------------------------------------------------------------------------------------------------------
// PriorityQueue
PriorityQueue::PriorityQueue(void)
{
for (Message *&tail : mTails)
{
tail = nullptr;
}
}
Message *PriorityQueue::FindFirstNonNullTail(Message::Priority aStartPriorityLevel) const
const Message *PriorityQueue::FindFirstNonNullTail(Message::Priority aStartPriorityLevel) const
{
// Find the first non-`nullptr` tail starting from the given priority
// level and moving forward (wrapping from priority value
// `kNumPriorities` -1 back to 0).
Message *tail = nullptr;
uint8_t priority;
const Message *tail = nullptr;
uint8_t priority;
priority = static_cast<uint8_t>(aStartPriorityLevel);
@@ -869,19 +870,15 @@ Message *PriorityQueue::FindFirstNonNullTail(Message::Priority aStartPriorityLev
return tail;
}
Message *PriorityQueue::GetHead(void) const
const Message *PriorityQueue::GetHead(void) const
{
Message *tail;
tail = FindFirstNonNullTail(Message::kPriorityLow);
return (tail == nullptr) ? nullptr : tail->Next();
return Message::NextOf(FindFirstNonNullTail(Message::kPriorityLow));
}
Message *PriorityQueue::GetHeadForPriority(Message::Priority aPriority) const
const Message *PriorityQueue::GetHeadForPriority(Message::Priority aPriority) const
{
Message *head;
Message *previousTail;
const Message *head;
const Message *previousTail;
if (mTails[aPriority] != nullptr)
{
@@ -899,7 +896,7 @@ Message *PriorityQueue::GetHeadForPriority(Message::Priority aPriority) const
return head;
}
Message *PriorityQueue::GetTail(void) const
const Message *PriorityQueue::GetTail(void) const
{
return FindFirstNonNullTail(Message::kPriorityLow);
}
@@ -983,15 +980,25 @@ void PriorityQueue::DequeueAndFreeAll(void)
}
}
Message::Iterator PriorityQueue::begin(void)
{
return Message::Iterator(GetHead());
}
Message::ConstIterator PriorityQueue::begin(void) const
{
return Message::ConstIterator(GetHead());
}
void PriorityQueue::GetInfo(uint16_t &aMessageCount, uint16_t &aBufferCount) const
{
aMessageCount = 0;
aBufferCount = 0;
for (const Message *message = GetHead(); message != nullptr; message = message->GetNext())
for (const Message &message : *this)
{
aMessageCount++;
aBufferCount += message->GetBufferCount();
aBufferCount += message.GetBufferCount();
}
}
+120 -10
View File
@@ -42,10 +42,12 @@
#include <openthread/platform/messagepool.h>
#include "common/as_core_type.hpp"
#include "common/clearable.hpp"
#include "common/code_utils.hpp"
#include "common/const_cast.hpp"
#include "common/data.hpp"
#include "common/encoding.hpp"
#include "common/iterator_utils.hpp"
#include "common/linked_list.hpp"
#include "common/locator.hpp"
#include "common/non_copyable.hpp"
@@ -1225,6 +1227,45 @@ public:
#endif // #if OPENTHREAD_CONFIG_MULTI_RADIO
protected:
class ConstIterator : public ItemPtrIterator<const Message, ConstIterator>
{
friend class ItemPtrIterator<const Message, ConstIterator>;
public:
ConstIterator(void) = default;
explicit ConstIterator(const Message *aMessage)
: ItemPtrIterator(aMessage)
{
}
private:
void Advance(void) { mItem = mItem->GetNext(); }
};
class Iterator : public ItemPtrIterator<Message, Iterator>
{
friend class ItemPtrIterator<Message, Iterator>;
public:
Iterator(void)
: mNext(nullptr)
{
}
explicit Iterator(Message *aMessage)
: ItemPtrIterator(aMessage)
, mNext(NextMessage(aMessage))
{
}
private:
void Advance(void);
static Message *NextMessage(Message *aMessage) { return (aMessage != nullptr) ? aMessage->GetNext() : nullptr; }
Message *mNext;
};
uint16_t GetReserved(void) const { return GetMetadata().mReserved; }
void SetReserved(uint16_t aReservedHeader) { GetMetadata().mReserved = aReservedHeader; }
@@ -1269,6 +1310,9 @@ private:
Message *const &Next(void) const { return GetMetadata().mNext; }
Message *& Prev(void) { return GetMetadata().mPrev; }
static Message * NextOf(Message *aMessage) { return (aMessage != nullptr) ? aMessage->Next() : nullptr; }
static const Message *NextOf(const Message *aMessage) { return (aMessage != nullptr) ? aMessage->Next() : nullptr; }
Error ResizeMessage(uint16_t aLength);
};
@@ -1297,7 +1341,7 @@ public:
* This constructor initializes the message queue.
*
*/
MessageQueue(void);
MessageQueue(void) { SetTail(nullptr); }
/**
* This method returns a pointer to the first message.
@@ -1305,7 +1349,15 @@ public:
* @returns A pointer to the first message.
*
*/
Message *GetHead(void) const;
Message *GetHead(void) { return Message::NextOf(GetTail()); }
/**
* This method returns a pointer to the first message.
*
* @returns A pointer to the first message.
*
*/
const Message *GetHead(void) const { return Message::NextOf(GetTail()); }
/**
* This method adds a message to the end of the list.
@@ -1355,16 +1407,28 @@ public:
*/
void GetInfo(uint16_t &aMessageCount, uint16_t &aBufferCount) const;
// The following methods are intended to support range-based `for`
// loop iteration over the queue entries and should not be used
// directly. The range-based `for` works correctly even if the
// current entry is removed from the queue during iteration.
Message::Iterator begin(void);
Message::Iterator end(void) { return Message::Iterator(); }
Message::ConstIterator begin(void) const;
Message::ConstIterator end(void) const { return Message::ConstIterator(); }
private:
Message *GetTail(void) const { return static_cast<Message *>(mData); }
void SetTail(Message *aMessage) { mData = aMessage; }
Message * GetTail(void) { return static_cast<Message *>(mData); }
const Message *GetTail(void) const { return static_cast<const Message *>(mData); }
void SetTail(Message *aMessage) { mData = aMessage; }
};
/**
* This class implements a priority queue.
*
*/
class PriorityQueue
class PriorityQueue : private Clearable<PriorityQueue>
{
friend class Message;
friend class MessageQueue;
@@ -1375,7 +1439,7 @@ public:
* This constructor initializes the priority queue.
*
*/
PriorityQueue(void);
PriorityQueue(void) { Clear(); }
/**
* This method returns a pointer to the first message.
@@ -1383,7 +1447,15 @@ public:
* @returns A pointer to the first message.
*
*/
Message *GetHead(void) const;
Message *GetHead(void) { return AsNonConst(AsConst(this)->GetHead()); }
/**
* This method returns a pointer to the first message.
*
* @returns A pointer to the first message.
*
*/
const Message *GetHead(void) const;
/**
* This method returns a pointer to the first message for a given priority level.
@@ -1394,7 +1466,21 @@ public:
* this priority level.
*
*/
Message *GetHeadForPriority(Message::Priority aPriority) const;
Message *GetHeadForPriority(Message::Priority aPriority)
{
return AsNonConst(AsConst(this)->GetHeadForPriority(aPriority));
}
/**
* This method returns a pointer to the first message for a given priority level.
*
* @param[in] aPriority Priority level.
*
* @returns A pointer to the first message with given priority level or `nullptr` if there is no messages with
* this priority level.
*
*/
const Message *GetHeadForPriority(Message::Priority aPriority) const;
/**
* This method adds a message to the queue.
@@ -1441,7 +1527,26 @@ public:
* @returns A pointer to the tail of the list.
*
*/
Message *GetTail(void) const;
Message *GetTail(void) { return AsNonConst(AsConst(this)->GetTail()); }
/**
* This method returns the tail of the list (last message in the list)
*
* @returns A pointer to the tail of the list.
*
*/
const Message *GetTail(void) const;
// The following methods are intended to support range-based `for`
// loop iteration over the queue entries and should not be used
// directly. The range-based `for` works correctly even if the
// current entry is removed from the queue during iteration.
Message::Iterator begin(void);
Message::Iterator end(void) { return Message::Iterator(); }
Message::ConstIterator begin(void) const;
Message::ConstIterator end(void) const { return Message::ConstIterator(); }
private:
uint8_t PrevPriority(uint8_t aPriority) const
@@ -1449,7 +1554,12 @@ private:
return (aPriority == Message::kNumPriorities - 1) ? 0 : (aPriority + 1);
}
Message *FindFirstNonNullTail(Message::Priority aStartPriorityLevel) const;
const Message *FindFirstNonNullTail(Message::Priority aStartPriorityLevel) const;
Message *FindFirstNonNullTail(Message::Priority aStartPriorityLevel)
{
return AsNonConst(AsConst(this)->FindFirstNonNullTail(aStartPriorityLevel));
}
Message *mTails[Message::kNumPriorities]; // Tail pointers associated with different priority levels.
};
+9 -11
View File
@@ -905,20 +905,21 @@ void Client::GetCallback(const Query &aQuery, Callback &aCallback, void *&aConte
Client::Query *Client::FindQueryById(uint16_t aMessageId)
{
Query * query;
Query * matchedQuery = nullptr;
QueryInfo info;
for (query = mQueries.GetHead(); query != nullptr; query = query->GetNext())
for (Query &query : mQueries)
{
info.ReadFrom(*query);
info.ReadFrom(query);
if (info.mMessageId == aMessageId)
{
matchedQuery = &query;
break;
}
}
return query;
return matchedQuery;
}
void Client::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMsgInfo)
@@ -1070,24 +1071,21 @@ void Client::HandleTimer(void)
{
TimeMilli now = TimerMilli::GetNow();
TimeMilli nextTime = now.GetDistantFuture();
Query * nextQuery;
QueryInfo info;
for (Query *query = mQueries.GetHead(); query != nullptr; query = nextQuery)
for (Query &query : mQueries)
{
nextQuery = query->GetNext();
info.ReadFrom(*query);
info.ReadFrom(query);
if (now >= info.mRetransmissionTime)
{
if (info.mTransmissionCount >= info.mConfig.GetMaxTxAttempts())
{
FinalizeQuery(*query, kErrorResponseTimeout);
FinalizeQuery(query, kErrorResponseTimeout);
continue;
}
SendQuery(*query, info, /* aUpdateTimer */ false);
SendQuery(query, info, /* aUpdateTimer */ false);
}
if (nextTime > info.mRetransmissionTime)
+9 -12
View File
@@ -708,13 +708,14 @@ Error Ip6::HandleFragment(Message &aMessage, Netif *aNetif, MessageInfo &aMessag
ExitNow();
}
for (message = mReassemblyList.GetHead(); message; message = message->GetNext())
for (Message &msg : mReassemblyList)
{
SuccessOrExit(error = message->Read(0, headerBuffer));
SuccessOrExit(error = msg.Read(0, headerBuffer));
if (message->GetDatagramTag() == fragmentHeader.GetIdentification() &&
if (msg.GetDatagramTag() == fragmentHeader.GetIdentification() &&
headerBuffer.GetSource() == header.GetSource() && headerBuffer.GetDestination() == header.GetDestination())
{
message = &msg;
break;
}
}
@@ -816,22 +817,18 @@ void Ip6::HandleTimeTick(void)
void Ip6::UpdateReassemblyList(void)
{
Message *next;
for (Message *message = mReassemblyList.GetHead(); message; message = next)
for (Message &message : mReassemblyList)
{
next = message->GetNext();
if (message->GetTimeout() > 0)
if (message.GetTimeout() > 0)
{
message->DecrementTimeout();
message.DecrementTimeout();
}
else
{
LogNote("Reassembly timeout.");
SendIcmpError(*message, Icmp::Header::kTypeTimeExceeded, Icmp::Header::kCodeFragmReasTimeEx);
SendIcmpError(message, Icmp::Header::kTypeTimeExceeded, Icmp::Header::kCodeFragmReasTimeEx);
mReassemblyList.DequeueAndFree(*message);
mReassemblyList.DequeueAndFree(message);
}
}
}
+9 -13
View File
@@ -344,14 +344,10 @@ void Mpl::HandleRetransmissionTimer(void)
TimeMilli now = TimerMilli::GetNow();
TimeMilli nextTime = now.GetDistantFuture();
Metadata metadata;
Message * message;
Message * nextMessage;
for (message = mBufferedMessageSet.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mBufferedMessageSet)
{
nextMessage = message->GetNext();
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if (now < metadata.mTransmissionTime)
{
@@ -367,7 +363,7 @@ void Mpl::HandleRetransmissionTimer(void)
if (metadata.mTransmissionCount < GetTimerExpirations())
{
Message *messageCopy = message->Clone(message->GetLength() - sizeof(Metadata));
Message *messageCopy = message.Clone(message.GetLength() - sizeof(Metadata));
if (messageCopy != nullptr)
{
@@ -380,7 +376,7 @@ void Mpl::HandleRetransmissionTimer(void)
}
metadata.GenerateNextTransmissionTime(now, kDataMessageInterval);
metadata.UpdateIn(*message);
metadata.UpdateIn(message);
if (nextTime > metadata.mTransmissionTime)
{
@@ -389,22 +385,22 @@ void Mpl::HandleRetransmissionTimer(void)
}
else
{
mBufferedMessageSet.Dequeue(*message);
mBufferedMessageSet.Dequeue(message);
if (metadata.mTransmissionCount == GetTimerExpirations())
{
if (metadata.mTransmissionCount > 1)
{
message->SetSubType(Message::kSubTypeMplRetransmission);
message.SetSubType(Message::kSubTypeMplRetransmission);
}
metadata.RemoveFrom(*message);
Get<Ip6>().EnqueueDatagram(*message);
metadata.RemoveFrom(message);
Get<Ip6>().EnqueueDatagram(message);
}
else
{
// Stop retransmitting if the number of timer expirations is already exceeded.
message->Free();
message.Free();
}
}
}
+15 -28
View File
@@ -113,18 +113,12 @@ exit:
Error Client::Stop(void)
{
Message * message = mPendingQueries.GetHead();
Message * messageToRemove;
QueryMetadata queryMetadata;
// Remove all pending queries.
while (message != nullptr)
for (Message &message : mPendingQueries)
{
messageToRemove = message;
message = message->GetNext();
QueryMetadata queryMetadata;
queryMetadata.ReadFrom(*messageToRemove);
FinalizeSntpTransaction(*messageToRemove, queryMetadata, 0, kErrorAbort);
queryMetadata.ReadFrom(message);
FinalizeSntpTransaction(message, queryMetadata, 0, kErrorAbort);
}
return mSocket.Close();
@@ -245,24 +239,21 @@ exit:
Message *Client::FindRelatedQuery(const Header &aResponseHeader, QueryMetadata &aQueryMetadata)
{
Header header;
Message *message = mPendingQueries.GetHead();
Message *matchedMessage = nullptr;
while (message != nullptr)
for (Message &message : mPendingQueries)
{
// Read originate timestamp.
aQueryMetadata.ReadFrom(*message);
aQueryMetadata.ReadFrom(message);
if (aQueryMetadata.mTransmitTimestamp == aResponseHeader.GetOriginateTimestampSeconds())
{
ExitNow();
matchedMessage = &message;
break;
}
message = message->GetNext();
}
exit:
return message;
return matchedMessage;
}
void Client::FinalizeSntpTransaction(Message & aQuery,
@@ -288,36 +279,32 @@ void Client::HandleRetransmissionTimer(void)
TimeMilli now = TimerMilli::GetNow();
TimeMilli nextTime = now.GetDistantFuture();
QueryMetadata queryMetadata;
Message * message;
Message * nextMessage;
Ip6::MessageInfo messageInfo;
for (message = mPendingQueries.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mPendingQueries)
{
nextMessage = message->GetNext();
queryMetadata.ReadFrom(*message);
queryMetadata.ReadFrom(message);
if (now >= queryMetadata.mTransmissionTime)
{
if (queryMetadata.mRetransmissionCount >= kMaxRetransmit)
{
// No expected response.
FinalizeSntpTransaction(*message, queryMetadata, 0, kErrorResponseTimeout);
FinalizeSntpTransaction(message, queryMetadata, 0, kErrorResponseTimeout);
continue;
}
// Increment retransmission counter and timer.
queryMetadata.mRetransmissionCount++;
queryMetadata.mTransmissionTime = now + kResponseTimeout;
queryMetadata.UpdateIn(*message);
queryMetadata.UpdateIn(message);
// Retransmit
messageInfo.SetPeerAddr(queryMetadata.mDestinationAddress);
messageInfo.SetPeerPort(queryMetadata.mDestinationPort);
messageInfo.SetSockAddr(queryMetadata.mSourceAddress);
SendCopy(*message, messageInfo);
SendCopy(message, messageInfo);
}
if (nextTime > queryMetadata.mTransmissionTime)
+13 -17
View File
@@ -136,18 +136,13 @@ exit:
void IndirectSender::ClearAllMessagesForSleepyChild(Child &aChild)
{
Message *message;
Message *nextMessage;
VerifyOrExit(aChild.GetIndirectMessageCount() > 0);
for (message = Get<MeshForwarder>().mSendQueue.GetHead(); message; message = nextMessage)
for (Message &message : Get<MeshForwarder>().mSendQueue)
{
nextMessage = message->GetNext();
message.ClearChildMask(Get<ChildTable>().GetChildIndex(aChild));
message->ClearChildMask(Get<ChildTable>().GetChildIndex(aChild));
Get<MeshForwarder>().RemoveMessageIfNoPendingTx(*message);
Get<MeshForwarder>().RemoveMessageIfNoPendingTx(message);
}
aChild.SetIndirectMessage(nullptr);
@@ -186,12 +181,12 @@ void IndirectSender::HandleChildModeChange(Child &aChild, Mle::DeviceMode aOldMo
{
uint16_t childIndex = Get<ChildTable>().GetChildIndex(aChild);
for (Message *message = Get<MeshForwarder>().mSendQueue.GetHead(); message; message = message->GetNext())
for (Message &message : Get<MeshForwarder>().mSendQueue)
{
if (message->GetChildMask(childIndex))
if (message.GetChildMask(childIndex))
{
message->ClearChildMask(childIndex);
message->SetDirectTransmission();
message.ClearChildMask(childIndex);
message.SetDirectTransmission();
}
}
@@ -214,19 +209,20 @@ void IndirectSender::HandleChildModeChange(Child &aChild, Mle::DeviceMode aOldMo
Message *IndirectSender::FindIndirectMessage(Child &aChild, bool aSupervisionTypeOnly)
{
Message *message;
Message *msg = nullptr;
uint16_t childIndex = Get<ChildTable>().GetChildIndex(aChild);
for (message = Get<MeshForwarder>().mSendQueue.GetHead(); message; message = message->GetNext())
for (Message &message : Get<MeshForwarder>().mSendQueue)
{
if (message->GetChildMask(childIndex) &&
(!aSupervisionTypeOnly || (message->GetType() == Message::kTypeSupervision)))
if (message.GetChildMask(childIndex) &&
(!aSupervisionTypeOnly || (message.GetType() == Message::kTypeSupervision)))
{
msg = &message;
break;
}
}
return message;
return msg;
}
void IndirectSender::RequestMessageUpdate(Child &aChild)
+18 -21
View File
@@ -1291,15 +1291,16 @@ void MeshForwarder::HandleFragment(const uint8_t * aFrame,
}
else // Received frame is a "next fragment".
{
for (message = mReassemblyList.GetHead(); message; message = message->GetNext())
for (Message &msg : mReassemblyList)
{
// Security Check: only consider reassembly buffers that had the same Security Enabled setting.
if (message->GetLength() == fragmentHeader.GetDatagramSize() &&
message->GetDatagramTag() == fragmentHeader.GetDatagramTag() &&
message->GetOffset() == fragmentHeader.GetDatagramOffset() &&
message->GetOffset() + aFrameLength <= fragmentHeader.GetDatagramSize() &&
message->IsLinkSecurityEnabled() == aLinkInfo.IsLinkSecurityEnabled())
if (msg.GetLength() == fragmentHeader.GetDatagramSize() &&
msg.GetDatagramTag() == fragmentHeader.GetDatagramTag() &&
msg.GetOffset() == fragmentHeader.GetDatagramOffset() &&
msg.GetOffset() + aFrameLength <= fragmentHeader.GetDatagramSize() &&
msg.IsLinkSecurityEnabled() == aLinkInfo.IsLinkSecurityEnabled())
{
message = &msg;
break;
}
}
@@ -1346,17 +1347,17 @@ exit:
void MeshForwarder::ClearReassemblyList(void)
{
for (const Message *message = mReassemblyList.GetHead(); message != nullptr; message = message->GetNext())
for (Message &message : mReassemblyList)
{
LogMessage(kMessageReassemblyDrop, *message, nullptr, kErrorNoFrameReceived);
LogMessage(kMessageReassemblyDrop, message, nullptr, kErrorNoFrameReceived);
if (message->GetType() == Message::kTypeIp6)
if (message.GetType() == Message::kTypeIp6)
{
mIpCounters.mRxFailure++;
}
}
mReassemblyList.DequeueAndFreeAll();
mReassemblyList.DequeueAndFree(message);
}
}
void MeshForwarder::HandleTimeTick(void)
@@ -1377,26 +1378,22 @@ void MeshForwarder::HandleTimeTick(void)
bool MeshForwarder::UpdateReassemblyList(void)
{
Message *next = nullptr;
for (Message *message = mReassemblyList.GetHead(); message; message = next)
for (Message &message : mReassemblyList)
{
next = message->GetNext();
if (message->GetTimeout() > 0)
if (message.GetTimeout() > 0)
{
message->DecrementTimeout();
message.DecrementTimeout();
}
else
{
LogMessage(kMessageReassemblyDrop, *message, nullptr, kErrorReassemblyTimeout);
LogMessage(kMessageReassemblyDrop, message, nullptr, kErrorReassemblyTimeout);
if (message->GetType() == Message::kTypeIp6)
if (message.GetType() == Message::kTypeIp6)
{
mIpCounters.mRxFailure++;
}
mReassemblyList.DequeueAndFree(*message);
mReassemblyList.DequeueAndFree(message);
}
}
+24 -36
View File
@@ -140,40 +140,37 @@ Error MeshForwarder::SendMessage(Message &aMessage)
void MeshForwarder::HandleResolved(const Ip6::Address &aEid, Error aError)
{
Message * cur, *next;
Ip6::Address ip6Dst;
bool enqueuedMessage = false;
for (cur = mResolvingQueue.GetHead(); cur; cur = next)
for (Message &message : mResolvingQueue)
{
next = cur->GetNext();
if (cur->GetType() != Message::kTypeIp6)
if (message.GetType() != Message::kTypeIp6)
{
continue;
}
IgnoreError(cur->Read(Ip6::Header::kDestinationFieldOffset, ip6Dst));
IgnoreError(message.Read(Ip6::Header::kDestinationFieldOffset, ip6Dst));
if (ip6Dst == aEid)
{
mResolvingQueue.Dequeue(*cur);
mResolvingQueue.Dequeue(message);
if (aError == kErrorNone)
{
#if OPENTHREAD_CONFIG_BACKBONE_ROUTER_ENABLE
// Pass back to IPv6 layer for DUA destination resolved by Backbone Query
if (ForwardDuaToBackboneLink(*cur, ip6Dst) != kErrorNone)
if (ForwardDuaToBackboneLink(message, ip6Dst) != kErrorNone)
#endif
{
mSendQueue.Enqueue(*cur);
mSendQueue.Enqueue(message);
enqueuedMessage = true;
}
}
else
{
LogMessage(kMessageDrop, *cur, nullptr, aError);
cur->Free();
LogMessage(kMessageDrop, message, nullptr, aError);
message.Free();
}
}
}
@@ -182,8 +179,6 @@ void MeshForwarder::HandleResolved(const Ip6::Address &aEid, Error aError)
{
mScheduleTransmissionTask.Post();
}
return;
}
#if OPENTHREAD_CONFIG_BACKBONE_ROUTER_ENABLE
@@ -280,30 +275,26 @@ exit:
void MeshForwarder::RemoveMessages(Child &aChild, Message::SubType aSubType)
{
Message *nextMessage;
for (Message *message = mSendQueue.GetHead(); message; message = nextMessage)
for (Message &message : mSendQueue)
{
nextMessage = message->GetNext();
if ((aSubType != Message::kSubTypeNone) && (aSubType != message->GetSubType()))
if ((aSubType != Message::kSubTypeNone) && (aSubType != message.GetSubType()))
{
continue;
}
if (mIndirectSender.RemoveMessageFromSleepyChild(*message, aChild) != kErrorNone)
if (mIndirectSender.RemoveMessageFromSleepyChild(message, aChild) != kErrorNone)
{
switch (message->GetType())
switch (message.GetType())
{
case Message::kTypeIp6:
{
Ip6::Header ip6header;
IgnoreError(message->Read(0, ip6header));
IgnoreError(message.Read(0, ip6header));
if (&aChild == static_cast<Child *>(Get<NeighborTable>().FindNeighbor(ip6header.GetDestination())))
{
message->ClearDirectTransmission();
message.ClearDirectTransmission();
}
break;
@@ -313,11 +304,11 @@ void MeshForwarder::RemoveMessages(Child &aChild, Message::SubType aSubType)
{
Lowpan::MeshHeader meshHeader;
IgnoreError(meshHeader.ParseFrom(*message));
IgnoreError(meshHeader.ParseFrom(message));
if (&aChild == static_cast<Child *>(Get<NeighborTable>().FindNeighbor(meshHeader.GetDestination())))
{
message->ClearDirectTransmission();
message.ClearDirectTransmission();
}
break;
@@ -328,41 +319,38 @@ void MeshForwarder::RemoveMessages(Child &aChild, Message::SubType aSubType)
}
}
RemoveMessageIfNoPendingTx(*message);
RemoveMessageIfNoPendingTx(message);
}
}
void MeshForwarder::RemoveDataResponseMessages(void)
{
Ip6::Header ip6Header;
Message * next;
for (Message *message = mSendQueue.GetHead(); message != nullptr; message = next)
for (Message &message : mSendQueue)
{
next = message->GetNext();
if (message->GetSubType() != Message::kSubTypeMleDataResponse)
if (message.GetSubType() != Message::kSubTypeMleDataResponse)
{
continue;
}
IgnoreError(message->Read(0, ip6Header));
IgnoreError(message.Read(0, ip6Header));
if (!(ip6Header.GetDestination().IsMulticast()))
{
for (Child &child : Get<ChildTable>().Iterate(Child::kInStateAnyExceptInvalid))
{
IgnoreError(mIndirectSender.RemoveMessageFromSleepyChild(*message, child));
IgnoreError(mIndirectSender.RemoveMessageFromSleepyChild(message, child));
}
}
if (mSendMessage == message)
if (mSendMessage == &message)
{
mSendMessage = nullptr;
}
LogMessage(kMessageDrop, *message, nullptr, kErrorNone);
mSendQueue.DequeueAndFree(*message);
LogMessage(kMessageDrop, message, nullptr, kErrorNone);
mSendQueue.DequeueAndFree(message);
}
}
+8 -15
View File
@@ -1982,15 +1982,12 @@ void Mle::HandleDelayedResponseTimer(void)
{
TimeMilli now = TimerMilli::GetNow();
TimeMilli nextSendTime = now.GetDistantFuture();
Message * nextMessage;
for (Message *message = mDelayedResponses.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mDelayedResponses)
{
DelayedResponseMetadata metadata;
nextMessage = message->GetNext();
metadata.ReadFrom(*message);
metadata.ReadFrom(message);
if (now < metadata.mSendTime)
{
@@ -2001,8 +1998,8 @@ void Mle::HandleDelayedResponseTimer(void)
}
else
{
mDelayedResponses.Dequeue(*message);
SendDelayedResponse(*message, metadata);
mDelayedResponses.Dequeue(message);
SendDelayedResponse(message, metadata);
}
}
@@ -2053,20 +2050,16 @@ void Mle::RemoveDelayedDataRequestMessage(const Ip6::Address &aDestination)
void Mle::RemoveDelayedMessage(Message::SubType aSubType, MessageType aMessageType, const Ip6::Address *aDestination)
{
Message *nextMessage;
for (Message *message = mDelayedResponses.GetHead(); message != nullptr; message = nextMessage)
for (Message &message : mDelayedResponses)
{
DelayedResponseMetadata metadata;
nextMessage = message->GetNext();
metadata.ReadFrom(message);
metadata.ReadFrom(*message);
if ((message->GetSubType() == aSubType) &&
if ((message.GetSubType() == aSubType) &&
((aDestination == nullptr) || (metadata.mDestination == *aDestination)))
{
mDelayedResponses.DequeueAndFree(*message);
mDelayedResponses.DequeueAndFree(message);
Log(kMessageRemoveDelayed, aMessageType, metadata.mDestination);
}
}
+2 -2
View File
@@ -3166,9 +3166,9 @@ Error MleRouter::SendChildUpdateRequest(Child &aChild)
{
uint16_t childIndex = Get<ChildTable>().GetChildIndex(aChild);
for (message = Get<MeshForwarder>().GetSendQueue().GetHead(); message; message = message->GetNext())
for (const Message &msg : Get<MeshForwarder>().GetSendQueue())
{
if (message->GetChildMask(childIndex) && message->GetSubType() == Message::kSubTypeMleChildUpdateRequest)
if (msg.GetChildMask(childIndex) && msg.GetSubType() == Message::kSubTypeMleChildUpdateRequest)
{
// No need to send the resync "Child Update Request" to the sleepy child
// if there is one already queued.
+70 -4
View File
@@ -46,9 +46,10 @@ static ot::MessagePool *sMessagePool;
// This function verifies the content of the message queue to match the passed in messages
void VerifyMessageQueueContent(ot::MessageQueue &aMessageQueue, int aExpectedLength, ...)
{
va_list args;
ot::Message *message;
ot::Message *msgArg;
const ot::MessageQueue &constQueue = aMessageQueue;
va_list args;
ot::Message * message;
ot::Message * msgArg;
va_start(args, aExpectedLength);
@@ -69,10 +70,34 @@ void VerifyMessageQueueContent(ot::MessageQueue &aMessageQueue, int aExpectedLen
aExpectedLength--;
}
VerifyOrQuit(aExpectedLength == 0, "less entries than expected");
VerifyOrQuit(aExpectedLength == 0, "fewer entries than expected");
}
va_end(args);
// Check range-based `for` loop iteration using non-const iterator
message = aMessageQueue.GetHead();
for (ot::Message &msg : aMessageQueue)
{
VerifyOrQuit(message == &msg, "`for` loop iteration does not match expected");
message = message->GetNext();
}
VerifyOrQuit(message == nullptr, "`for` loop iteration resulted in fewer entries than expected");
// Check range-base `for` iteration using const iterator
message = aMessageQueue.GetHead();
for (const ot::Message &constMsg : constQueue)
{
VerifyOrQuit(message == &constMsg, "`for` loop iteration does not match expected");
message = message->GetNext();
}
VerifyOrQuit(message == nullptr, "`for` loop iteration resulted in fewer entries than expected");
}
void TestMessageQueue(void)
@@ -174,6 +199,47 @@ void TestMessageQueue(void)
messageQueue.Dequeue(*messages[0]);
VerifyMessageQueueContent(messageQueue, 0);
// Range-based `for` and dequeue during iteration
for (uint16_t removeIndex = 0; removeIndex < 5; removeIndex++)
{
uint16_t index = 0;
messageQueue.Enqueue(*messages[0]);
messageQueue.Enqueue(*messages[1]);
messageQueue.Enqueue(*messages[2]);
messageQueue.Enqueue(*messages[3]);
messageQueue.Enqueue(*messages[4]);
VerifyMessageQueueContent(messageQueue, 5, messages[0], messages[1], messages[2], messages[3], messages[4]);
// While iterating over the queue remove the entry at `removeIndex`
for (ot::Message &message : messageQueue)
{
if (index == removeIndex)
{
messageQueue.Dequeue(message);
}
VerifyOrQuit(&message == messages[index++]);
}
index = 0;
// Iterate over the queue and remove all
for (ot::Message &message : messageQueue)
{
if (index == removeIndex)
{
index++;
}
VerifyOrQuit(&message == messages[index++]);
messageQueue.Dequeue(message);
}
VerifyMessageQueueContent(messageQueue, 0);
}
testFreeInstance(sInstance);
}
+80 -5
View File
@@ -42,11 +42,12 @@
// This function verifies the content of the priority queue to match the passed in messages
void VerifyPriorityQueueContent(ot::PriorityQueue &aPriorityQueue, int aExpectedLength, ...)
{
va_list args;
ot::Message *message;
ot::Message *msgArg;
int8_t curPriority = ot::Message::kNumPriorities;
uint16_t msgCount, bufCount;
const ot::PriorityQueue &constQueue = aPriorityQueue;
va_list args;
ot::Message * message;
ot::Message * msgArg;
int8_t curPriority = ot::Message::kNumPriorities;
uint16_t msgCount, bufCount;
// Check the `GetInfo`
aPriorityQueue.GetInfo(msgCount, bufCount);
@@ -106,6 +107,30 @@ void VerifyPriorityQueueContent(ot::PriorityQueue &aPriorityQueue, int aExpected
}
va_end(args);
// Check range-based `for` loop iteration using non-const iterator
message = aPriorityQueue.GetHead();
for (ot::Message &msg : aPriorityQueue)
{
VerifyOrQuit(message == &msg, "`for` loop iteration does not match expected");
message = message->GetNext();
}
VerifyOrQuit(message == nullptr, "`for` loop iteration resulted in fewer entries than expected");
// Check range-base `for` iteration using const iterator
message = aPriorityQueue.GetHead();
for (const ot::Message &constMsg : constQueue)
{
VerifyOrQuit(message == &constMsg, "`for` loop iteration does not match expected");
message = message->GetNext();
}
VerifyOrQuit(message == nullptr, "`for` loop iteration resulted in fewer entries than expected");
}
// This function verifies the content of the message queue to match the passed in messages
@@ -290,6 +315,56 @@ void TestPriorityQueue(void)
VerifyPriorityQueueContent(queue, 1, msgNor[0]);
VerifyMsgQueueContent(messageQueue, 1, msgNor[1]);
queue.Dequeue(*msgNor[0]);
VerifyPriorityQueueContent(queue, 0);
messageQueue.Dequeue(*msgNor[1]);
VerifyMsgQueueContent(messageQueue, 0);
for (ot::Message *message : msgNor)
{
SuccessOrQuit(message->SetPriority(ot::Message::kPriorityNormal));
}
// Range-based `for` and dequeue during iteration
for (uint16_t removeIndex = 0; removeIndex < 4; removeIndex++)
{
uint16_t index = 0;
queue.Enqueue(*msgNor[0]);
queue.Enqueue(*msgNor[1]);
queue.Enqueue(*msgNor[2]);
queue.Enqueue(*msgNor[3]);
VerifyPriorityQueueContent(queue, 4, msgNor[0], msgNor[1], msgNor[2], msgNor[3]);
// While iterating over the queue remove the entry at `removeIndex`
for (ot::Message &message : queue)
{
if (index == removeIndex)
{
queue.Dequeue(message);
}
VerifyOrQuit(&message == msgNor[index++]);
}
index = 0;
// Iterate over the queue and remove all
for (ot::Message &message : queue)
{
if (index == removeIndex)
{
index++;
}
VerifyOrQuit(&message == msgNor[index++]);
queue.Dequeue(message);
}
VerifyPriorityQueueContent(queue, 0);
}
testFreeInstance(instance);
}