From cf5637f3cabf70b39d153cc4ba61be6eb1b6461b Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Tue, 16 Jul 2024 09:11:10 -0700 Subject: [PATCH] [mesh-forwarder] `RemoveMessagesForChild()` to use a predicate fn (#10514) This commit updates `RemoveMessagesForChild()` to use a predicate function to determine which messages should be removed. This replaces the `Message::SubType` filtering and allows messages matching multiple sub-types to be removed together. --- src/core/thread/mesh_forwarder.hpp | 26 +++++++++++++++----- src/core/thread/mesh_forwarder_ftd.cpp | 34 +++++++++----------------- src/core/thread/mle_router.cpp | 31 +++++++++++++++++++---- src/core/thread/mle_router.hpp | 3 +++ 4 files changed, 60 insertions(+), 34 deletions(-) diff --git a/src/core/thread/mesh_forwarder.hpp b/src/core/thread/mesh_forwarder.hpp index 99b622fbd..2790be9ba 100644 --- a/src/core/thread/mesh_forwarder.hpp +++ b/src/core/thread/mesh_forwarder.hpp @@ -241,16 +241,30 @@ public: void SetRxOnWhenIdle(bool aRxOnWhenIdle); #if OPENTHREAD_FTD + /** - * Frees any messages queued for an existing child. + * Represents a predicate function for checking if a given `Message` meets specific criteria. * - * @param[in] aChild A reference to the child. - * @param[in] aSubType The message sub-type to remove. - * Use Message::kSubTypeNone remove all messages for @p aChild. + * @param[in] aMessage The message to evaluate. + * + * @return TRUE If the @p aMessage satisfies the predicate condition. + * @return FALSE If the @p aMessage does not satisfy the predicate condition. * */ - void RemoveMessages(Child &aChild, Message::SubType aSubType); -#endif + typedef bool (&MessageChecker)(const Message &aMessage); + + /** + * Removes and frees messages queued for a child, based on a given predicate. + * + * The `aChild` can be either sleepy or non-sleepy. + * + * @param[in] aChild The child whose messages are to be evaluated. + * @param[in] aMessageChecker The predicate function to filter messages. + * + */ + void RemoveMessagesForChild(Child &aChild, MessageChecker aMessageChecker); + +#endif // OPENTHREAD_FTD /** * Frees unicast/multicast MLE Data Responses from Send Message Queue if any. diff --git a/src/core/thread/mesh_forwarder_ftd.cpp b/src/core/thread/mesh_forwarder_ftd.cpp index 603c7ed14..904286439 100644 --- a/src/core/thread/mesh_forwarder_ftd.cpp +++ b/src/core/thread/mesh_forwarder_ftd.cpp @@ -266,49 +266,37 @@ exit: return error; } -void MeshForwarder::RemoveMessages(Child &aChild, Message::SubType aSubType) +void MeshForwarder::RemoveMessagesForChild(Child &aChild, MessageChecker &aMessageChecker) { for (Message &message : mSendQueue) { - if ((aSubType != Message::kSubTypeNone) && (aSubType != message.GetSubType())) + if (!aMessageChecker(message)) { continue; } if (mIndirectSender.RemoveMessageFromSleepyChild(message, aChild) != kErrorNone) { - switch (message.GetType()) - { - case Message::kTypeIp6: + const Neighbor *neighbor = nullptr; + + if (message.GetType() == Message::kTypeIp6) { Ip6::Header ip6header; IgnoreError(message.Read(0, ip6header)); - - if (&aChild == Get().FindNeighbor(ip6header.GetDestination())) - { - message.ClearDirectTransmission(); - } - - break; + neighbor = Get().FindNeighbor(ip6header.GetDestination()); } - - case Message::kType6lowpan: + else if (message.GetType() == Message::kType6lowpan) { Lowpan::MeshHeader meshHeader; IgnoreError(meshHeader.ParseFrom(message)); - - if (&aChild == Get().FindNeighbor(meshHeader.GetDestination())) - { - message.ClearDirectTransmission(); - } - - break; + neighbor = Get().FindNeighbor(meshHeader.GetDestination()); } - default: - break; + if (&aChild == neighbor) + { + message.ClearDirectTransmission(); } } diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index a776e533e..d4d1878d5 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -1960,6 +1960,30 @@ exit: } #endif // OPENTHREAD_CONFIG_TMF_PROXY_DUA_ENABLE +bool MleRouter::IsMessageMleSubType(const Message &aMessage) +{ + bool isMle = false; + + switch (aMessage.GetSubType()) + { + case Message::kSubTypeMleGeneral: + case Message::kSubTypeMleChildIdRequest: + case Message::kSubTypeMleChildUpdateRequest: + case Message::kSubTypeMleDataResponse: + isMle = true; + break; + default: + break; + } + + return isMle; +} + +bool MleRouter::IsMessageChildUpdateRequest(const Message &aMessage) +{ + return aMessage.GetSubType() == Message::kSubTypeMleChildUpdateRequest; +} + void MleRouter::HandleChildIdRequest(RxInfo &aRxInfo) { Error error = kErrorNone; @@ -1990,10 +2014,7 @@ void MleRouter::HandleChildIdRequest(RxInfo &aRxInfo) SuccessOrExit(error = aRxInfo.mMessage.ReadAndMatchResponseTlvWith(child->GetChallenge())); - Get().RemoveMessages(*child, Message::kSubTypeMleGeneral); - Get().RemoveMessages(*child, Message::kSubTypeMleChildIdRequest); - Get().RemoveMessages(*child, Message::kSubTypeMleChildUpdateRequest); - Get().RemoveMessages(*child, Message::kSubTypeMleDataResponse); + Get().RemoveMessagesForChild(*child, IsMessageMleSubType); SuccessOrExit(error = aRxInfo.mMessage.ReadFrameCounterTlvs(linkFrameCounter, mleFrameCounter)); @@ -2883,7 +2904,7 @@ Error MleRouter::SendChildUpdateRequest(Child &aChild) // Remove queued outdated "Child Update Request" when // there is newer Network Data is to send. - Get().RemoveMessages(aChild, Message::kSubTypeMleChildUpdateRequest); + Get().RemoveMessagesForChild(aChild, IsMessageChildUpdateRequest); break; } } diff --git a/src/core/thread/mle_router.hpp b/src/core/thread/mle_router.hpp index ac39177cf..8a43b1941 100644 --- a/src/core/thread/mle_router.hpp +++ b/src/core/thread/mle_router.hpp @@ -608,6 +608,9 @@ private: void HandleNetworkDataUpdateRouter(void); void HandleDiscoveryRequest(RxInfo &aRxInfo); + static bool IsMessageMleSubType(const Message &aMessage); + static bool IsMessageChildUpdateRequest(const Message &aMessage); + Error ProcessRouteTlv(const RouteTlv &aRouteTlv, RxInfo &aRxInfo); Error ReadAndProcessRouteTlvOnFed(RxInfo &aRxInfo, uint8_t aParentId);