From cba1bebf1f1ef3531bfe5397c75c3c6e045218b6 Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Wed, 24 Aug 2022 23:35:21 -0700 Subject: [PATCH] [child-table] check neighbor is in `ChildTable` before cast as `Child` (#8071) This commit adds a check `Get().Contains(Neighbor &)` to ensure that the neighbor is from the child table before casting the neighbor entry to `Child`. This adds safety and protection against potential corner-case where we have a neighbor entry which is `Parent` or `ParentCandidate` that is a REED. --- src/core/thread/mesh_forwarder.cpp | 4 ++-- src/core/thread/mesh_forwarder_ftd.cpp | 28 ++++++++++++++------------ src/core/thread/mle_router.cpp | 5 +++-- src/core/thread/neighbor_table.cpp | 1 + 4 files changed, 21 insertions(+), 17 deletions(-) diff --git a/src/core/thread/mesh_forwarder.cpp b/src/core/thread/mesh_forwarder.cpp index 8cdfa0165..07e7dcfa1 100644 --- a/src/core/thread/mesh_forwarder.cpp +++ b/src/core/thread/mesh_forwarder.cpp @@ -1834,8 +1834,8 @@ uint16_t MeshForwarder::CalcFrameVersion(const Neighbor *aNeighbor, bool aIePres version = Mac::Frame::kFcfFrameVersion2015; } #if OPENTHREAD_FTD && OPENTHREAD_CONFIG_MAC_CSL_TRANSMITTER_ENABLE - else if (aNeighbor != nullptr && !Mle::IsActiveRouter(aNeighbor->GetRloc16()) && - Get().IsRouterOrLeader() && static_cast(aNeighbor)->IsCslSynchronized()) + else if ((aNeighbor != nullptr) && Get().Contains(*aNeighbor) && + static_cast(aNeighbor)->IsCslSynchronized()) { version = Mac::Frame::kFcfFrameVersion2015; } diff --git a/src/core/thread/mesh_forwarder_ftd.cpp b/src/core/thread/mesh_forwarder_ftd.cpp index 3132e199c..14e1a8d6b 100644 --- a/src/core/thread/mesh_forwarder_ftd.cpp +++ b/src/core/thread/mesh_forwarder_ftd.cpp @@ -50,7 +50,6 @@ Error MeshForwarder::SendMessage(Message &aMessage) { Mle::MleRouter &mle = Get(); Error error = kErrorNone; - Neighbor * neighbor; aMessage.SetOffset(0); aMessage.SetDatagramTag(0); @@ -104,17 +103,20 @@ Error MeshForwarder::SendMessage(Message &aMessage) } } } - else if ((neighbor = Get().FindNeighbor(ip6Header.GetDestination())) != nullptr && - !neighbor->IsRxOnWhenIdle() && !aMessage.IsDirectTransmission()) + else // Destination is unicast { - // destined for a sleepy child - Child &child = *static_cast(neighbor); - mIndirectSender.AddMessageForSleepyChild(aMessage, child); - } - else - { - // schedule direct transmission - aMessage.SetDirectTransmission(); + Neighbor *neighbor = Get().FindNeighbor(ip6Header.GetDestination()); + + if ((neighbor != nullptr) && !neighbor->IsRxOnWhenIdle() && !aMessage.IsDirectTransmission() && + Get().Contains(*neighbor)) + { + // Destined for a sleepy child + mIndirectSender.AddMessageForSleepyChild(aMessage, *static_cast(neighbor)); + } + else + { + aMessage.SetDirectTransmission(); + } } break; @@ -283,7 +285,7 @@ void MeshForwarder::RemoveMessages(Child &aChild, Message::SubType aSubType) IgnoreError(message.Read(0, ip6header)); - if (&aChild == static_cast(Get().FindNeighbor(ip6header.GetDestination()))) + if (&aChild == Get().FindNeighbor(ip6header.GetDestination())) { message.ClearDirectTransmission(); } @@ -297,7 +299,7 @@ void MeshForwarder::RemoveMessages(Child &aChild, Message::SubType aSubType) IgnoreError(meshHeader.ParseFrom(message)); - if (&aChild == static_cast(Get().FindNeighbor(meshHeader.GetDestination()))) + if (&aChild == Get().FindNeighbor(meshHeader.GetDestination())) { message.ClearDirectTransmission(); } diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index a91ae96ae..bb9137594 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -2724,7 +2724,8 @@ void MleRouter::HandleChildUpdateResponse(RxInfo &aRxInfo) Child * child; uint16_t addressRegistrationOffset = 0; - if ((aRxInfo.mNeighbor == nullptr) || IsActiveRouter(aRxInfo.mNeighbor->GetRloc16())) + if ((aRxInfo.mNeighbor == nullptr) || IsActiveRouter(aRxInfo.mNeighbor->GetRloc16()) || + !Get().Contains(*aRxInfo.mNeighbor)) { Log(kMessageReceive, kTypeChildUpdateResponseOfUnknownChild, aRxInfo.mMessageInfo.GetPeerAddr()); ExitNow(error = kErrorNotFound); @@ -3529,7 +3530,7 @@ void MleRouter::RemoveNeighbor(Neighbor &aNeighbor) } else if (!IsActiveRouter(aNeighbor.GetRloc16())) { - OT_ASSERT(mChildTable.GetChildIndex(static_cast(aNeighbor)) < kMaxChildren); + OT_ASSERT(mChildTable.Contains(aNeighbor)); if (aNeighbor.IsStateValidOrRestoring()) { diff --git a/src/core/thread/neighbor_table.cpp b/src/core/thread/neighbor_table.cpp index f7856a11d..e399b3dea 100644 --- a/src/core/thread/neighbor_table.cpp +++ b/src/core/thread/neighbor_table.cpp @@ -271,6 +271,7 @@ void NeighborTable::Signal(Event aEvent, const Neighbor &aNeighbor) case kChildRemoved: case kChildModeChanged: #if OPENTHREAD_FTD + OT_ASSERT(Get().Contains(aNeighbor)); static_cast(info.mInfo.mChild).SetFrom(static_cast(aNeighbor)); #endif break;