From 237c91b9398bfe55426c45cbe5bcb5713c178ba3 Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Mon, 21 Mar 2022 09:28:12 -0700 Subject: [PATCH] [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. --- src/core/coap/coap.cpp | 72 ++++++-------- src/core/coap/coap_message.cpp | 10 ++ src/core/coap/coap_message.hpp | 43 +++++++- src/core/common/message.cpp | 77 ++++++++------- src/core/common/message.hpp | 130 +++++++++++++++++++++++-- src/core/net/dns_client.cpp | 20 ++-- src/core/net/ip6.cpp | 21 ++-- src/core/net/ip6_mpl.cpp | 22 ++--- src/core/net/sntp_client.cpp | 43 +++----- src/core/thread/indirect_sender.cpp | 30 +++--- src/core/thread/mesh_forwarder.cpp | 39 ++++---- src/core/thread/mesh_forwarder_ftd.cpp | 60 +++++------- src/core/thread/mle.cpp | 23 ++--- src/core/thread/mle_router.cpp | 4 +- tests/unit/test_message_queue.cpp | 74 +++++++++++++- tests/unit/test_priority_queue.cpp | 85 +++++++++++++++- 16 files changed, 503 insertions(+), 250 deletions(-) diff --git a/src/core/coap/coap.cpp b/src/core/coap/coap.cpp index 92375c729..0ffb75718 100644 --- a/src/core/coap/coap.cpp +++ b/src/core/coap/coap.cpp @@ -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; } diff --git a/src/core/coap/coap_message.cpp b/src/core/coap/coap_message.cpp index ae0c33f7e..5fba03f64 100644 --- a/src/core/coap/coap_message.cpp +++ b/src/core/coap/coap_message.cpp @@ -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; diff --git a/src/core/coap/coap_message.hpp b/src/core/coap/coap_message.hpp index cf56cb070..75f6cf492 100644 --- a/src/core/coap/coap_message.hpp +++ b/src/core/coap/coap_message.hpp @@ -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(ot::Message::ConstIterator::operator*()); } + const Message *operator->(void) + { + return static_cast(ot::Message::ConstIterator::operator->()); + } + }; + + class Iterator : public ot::Message::Iterator + { + public: + using ot::Message::Iterator::Iterator; + + Message &operator*(void) { return static_cast(ot::Message::Iterator::operator*()); } + Message *operator->(void) { return static_cast(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(ot::MessageQueue::GetHead()); } + Message *GetHead(void) { return static_cast(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(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(); } }; /** diff --git a/src/core/common/message.cpp b/src/core/common/message.cpp index 4c48204b5..409a0d33e 100644 --- a/src/core/common/message.cpp +++ b/src/core/common/message.cpp @@ -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(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(); } } diff --git a/src/core/common/message.hpp b/src/core/common/message.hpp index 3d7da03ed..37bcaaf37 100644 --- a/src/core/common/message.hpp +++ b/src/core/common/message.hpp @@ -42,10 +42,12 @@ #include #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 + { + friend class ItemPtrIterator; + + public: + ConstIterator(void) = default; + + explicit ConstIterator(const Message *aMessage) + : ItemPtrIterator(aMessage) + { + } + + private: + void Advance(void) { mItem = mItem->GetNext(); } + }; + + class Iterator : public ItemPtrIterator + { + friend class ItemPtrIterator; + + 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(mData); } - void SetTail(Message *aMessage) { mData = aMessage; } + Message * GetTail(void) { return static_cast(mData); } + const Message *GetTail(void) const { return static_cast(mData); } + void SetTail(Message *aMessage) { mData = aMessage; } }; /** * This class implements a priority queue. * */ -class PriorityQueue +class PriorityQueue : private Clearable { 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. }; diff --git a/src/core/net/dns_client.cpp b/src/core/net/dns_client.cpp index a26c0a34e..bb33d7072 100644 --- a/src/core/net/dns_client.cpp +++ b/src/core/net/dns_client.cpp @@ -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) diff --git a/src/core/net/ip6.cpp b/src/core/net/ip6.cpp index a13c99df7..205ac5705 100644 --- a/src/core/net/ip6.cpp +++ b/src/core/net/ip6.cpp @@ -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); } } } diff --git a/src/core/net/ip6_mpl.cpp b/src/core/net/ip6_mpl.cpp index a41932a92..3e0156ae6 100644 --- a/src/core/net/ip6_mpl.cpp +++ b/src/core/net/ip6_mpl.cpp @@ -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().EnqueueDatagram(*message); + metadata.RemoveFrom(message); + Get().EnqueueDatagram(message); } else { // Stop retransmitting if the number of timer expirations is already exceeded. - message->Free(); + message.Free(); } } } diff --git a/src/core/net/sntp_client.cpp b/src/core/net/sntp_client.cpp index 65641dae6..b2e82f200 100644 --- a/src/core/net/sntp_client.cpp +++ b/src/core/net/sntp_client.cpp @@ -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) diff --git a/src/core/thread/indirect_sender.cpp b/src/core/thread/indirect_sender.cpp index a6c88b90f..8fa0fd820 100644 --- a/src/core/thread/indirect_sender.cpp +++ b/src/core/thread/indirect_sender.cpp @@ -136,18 +136,13 @@ exit: void IndirectSender::ClearAllMessagesForSleepyChild(Child &aChild) { - Message *message; - Message *nextMessage; - VerifyOrExit(aChild.GetIndirectMessageCount() > 0); - for (message = Get().mSendQueue.GetHead(); message; message = nextMessage) + for (Message &message : Get().mSendQueue) { - nextMessage = message->GetNext(); + message.ClearChildMask(Get().GetChildIndex(aChild)); - message->ClearChildMask(Get().GetChildIndex(aChild)); - - Get().RemoveMessageIfNoPendingTx(*message); + Get().RemoveMessageIfNoPendingTx(message); } aChild.SetIndirectMessage(nullptr); @@ -186,12 +181,12 @@ void IndirectSender::HandleChildModeChange(Child &aChild, Mle::DeviceMode aOldMo { uint16_t childIndex = Get().GetChildIndex(aChild); - for (Message *message = Get().mSendQueue.GetHead(); message; message = message->GetNext()) + for (Message &message : Get().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().GetChildIndex(aChild); - for (message = Get().mSendQueue.GetHead(); message; message = message->GetNext()) + for (Message &message : Get().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) diff --git a/src/core/thread/mesh_forwarder.cpp b/src/core/thread/mesh_forwarder.cpp index cc13c3a3c..54939fc67 100644 --- a/src/core/thread/mesh_forwarder.cpp +++ b/src/core/thread/mesh_forwarder.cpp @@ -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); } } diff --git a/src/core/thread/mesh_forwarder_ftd.cpp b/src/core/thread/mesh_forwarder_ftd.cpp index 1e098dfe7..76775c511 100644 --- a/src/core/thread/mesh_forwarder_ftd.cpp +++ b/src/core/thread/mesh_forwarder_ftd.cpp @@ -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(Get().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(Get().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().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); } } diff --git a/src/core/thread/mle.cpp b/src/core/thread/mle.cpp index 6a9983d0d..23f6f4c62 100644 --- a/src/core/thread/mle.cpp +++ b/src/core/thread/mle.cpp @@ -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); } } diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index 86857581c..da679e106 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -3166,9 +3166,9 @@ Error MleRouter::SendChildUpdateRequest(Child &aChild) { uint16_t childIndex = Get().GetChildIndex(aChild); - for (message = Get().GetSendQueue().GetHead(); message; message = message->GetNext()) + for (const Message &msg : Get().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. diff --git a/tests/unit/test_message_queue.cpp b/tests/unit/test_message_queue.cpp index 46031d2ee..424702971 100644 --- a/tests/unit/test_message_queue.cpp +++ b/tests/unit/test_message_queue.cpp @@ -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); } diff --git a/tests/unit/test_priority_queue.cpp b/tests/unit/test_priority_queue.cpp index 30a3a77eb..5d7be54ef 100644 --- a/tests/unit/test_priority_queue.cpp +++ b/tests/unit/test_priority_queue.cpp @@ -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); }