diff --git a/src/core/thread/mesh_forwarder.cpp b/src/core/thread/mesh_forwarder.cpp index 0c42a3a9b..6648a760d 100644 --- a/src/core/thread/mesh_forwarder.cpp +++ b/src/core/thread/mesh_forwarder.cpp @@ -237,6 +237,79 @@ void MeshForwarder::UpdateIndirectMessages(void) } } +void MeshForwarder::RemoveMessages(Child &aChild, uint8_t aSubType) +{ + ThreadNetif &netif = GetNetif(); + Message *nextMessage; + + for (Message *message = mSendQueue.GetHead(); message; message = nextMessage) + { + uint8_t childIndex = netif.GetMle().GetChildIndex(aChild); + + nextMessage = message->GetNext(); + + if ((aSubType != Message::kSubTypeNone) && (aSubType != message->GetSubType())) + { + continue; + } + + if (message->GetChildMask(childIndex)) + { + message->ClearChildMask(childIndex); + mSourceMatchController.DecrementMessageCount(aChild); + } + else + { + switch (message->GetType()) + { + case Message::kTypeIp6: + { + Ip6::Header ip6header; + + IgnoreReturnValue(message->Read(0, sizeof(ip6header), &ip6header)); + + if (&aChild == static_cast(netif.GetMle().GetNeighbor(ip6header.GetDestination()))) + { + message->ClearDirectTransmission(); + } + + break; + } + + case Message::kType6lowpan: + { + Lowpan::MeshHeader meshHeader; + + IgnoreReturnValue(meshHeader.Init(*message)); + + if (&aChild == static_cast(netif.GetMle().GetNeighbor(meshHeader.GetDestination()))) + { + message->ClearDirectTransmission(); + } + + break; + } + + default: + { + break; + } + } + } + + if (!message->IsChildPending() && !message->GetDirectTransmission()) + { + if (mSendMessage == message) + { + mSendMessage = NULL; + } + + mSendQueue.Dequeue(*message); + message->Free(); + } + } +} + void MeshForwarder::ScheduleTransmissionTask(Tasklet &aTasklet) { GetOwner(aTasklet).ScheduleTransmissionTask(); diff --git a/src/core/thread/mesh_forwarder.hpp b/src/core/thread/mesh_forwarder.hpp index 0f2fc83c9..c21b1f1da 100644 --- a/src/core/thread/mesh_forwarder.hpp +++ b/src/core/thread/mesh_forwarder.hpp @@ -164,6 +164,16 @@ public: */ void UpdateIndirectMessages(void); + /** + * This method frees any messages queued for an existing child. + * + * @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. + * + */ + void RemoveMessages(Child &aChild, uint8_t aSubType); + /** * This method returns a reference to the send queue. * diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index 685ad791c..3d01972a0 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -1706,13 +1706,6 @@ otError MleRouter::HandleParentRequest(const Message &aMessage, const Ip6::Messa child = FindChild(macAddr); - if (child != NULL && !child->IsFullThreadDevice()) - { - // Parent Request from a MTD child means that the child had detached. It can be safely removed. - RemoveNeighbor(*child); - child = NULL; - } - if (child == NULL) { VerifyOrExit((child = NewChild()) != NULL); @@ -1919,6 +1912,8 @@ otError MleRouter::SendParentResponse(Child *aChild, const ChallengeTlv &challen uint16_t delay; VerifyOrExit((message = NewMleMessage()) != NULL, error = OT_ERROR_NO_BUFS); + message->SetDirectTransmission(); + SuccessOrExit(error = AppendHeader(*message, Header::kCommandParentResponse)); SuccessOrExit(error = AppendSourceAddress(*message)); SuccessOrExit(error = AppendLeaderData(*message)); @@ -2084,6 +2079,10 @@ otError MleRouter::HandleChildIdRequest(const Message &aMessage, const Ip6::Mess memcmp(response.GetResponse(), child->GetChallenge(), child->GetChallengeSize()) == 0, error = OT_ERROR_SECURITY); + // Remove existing MLE messages + netif.GetMeshForwarder().RemoveMessages(*child, Message::kSubTypeMleGeneral); + netif.GetMeshForwarder().RemoveMessages(*child, Message::kSubTypeMleChildUpdateRequest); + // Link-Layer Frame Counter SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kLinkFrameCounter, sizeof(linkFrameCounter), linkFrameCounter)); @@ -2163,7 +2162,6 @@ otError MleRouter::HandleChildIdRequest(const Message &aMessage, const Ip6::Mess } } - child->SetLastHeard(TimerMilli::GetNow()); child->SetLinkFrameCounter(linkFrameCounter.GetFrameCounter()); child->SetMleFrameCounter(mleFrameCounter.GetFrameCounter());