diff --git a/src/core/thread/mle.cpp b/src/core/thread/mle.cpp index 1df7d228e..50fdb0259 100644 --- a/src/core/thread/mle.cpp +++ b/src/core/thread/mle.cpp @@ -2750,18 +2750,6 @@ void Mle::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageIn neighbor->SetMleFrameCounter(frameCounter + 1); } - else - { - if (!(command == Header::kCommandLinkRequest || command == Header::kCommandLinkAccept || - command == Header::kCommandLinkAcceptAndRequest || command == Header::kCommandAdvertisement || - command == Header::kCommandParentRequest || command == Header::kCommandParentResponse || - command == Header::kCommandChildIdRequest || command == Header::kCommandChildUpdateRequest || - command == Header::kCommandChildUpdateResponse || command == Header::kCommandAnnounce)) - { - otLogDebgMle("mle sequence unknown! %d", command); - ExitNow(); - } - } switch (command) { @@ -2782,11 +2770,11 @@ void Mle::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageIn break; case Header::kCommandDataRequest: - Get().HandleDataRequest(aMessage, aMessageInfo); + Get().HandleDataRequest(aMessage, aMessageInfo, neighbor); break; case Header::kCommandDataResponse: - HandleDataResponse(aMessage, aMessageInfo); + HandleDataResponse(aMessage, aMessageInfo, neighbor); break; case Header::kCommandParentRequest: @@ -2802,7 +2790,7 @@ void Mle::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageIn break; case Header::kCommandChildIdResponse: - HandleChildIdResponse(aMessage, aMessageInfo); + HandleChildIdResponse(aMessage, aMessageInfo, neighbor); break; case Header::kCommandChildUpdateRequest: @@ -2835,7 +2823,7 @@ void Mle::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageIn #if OPENTHREAD_CONFIG_TIME_SYNC_ENABLE case Header::kCommandTimeSync: - Get().HandleTimeSync(aMessage, aMessageInfo); + Get().HandleTimeSync(aMessage, aMessageInfo, neighbor); break; #endif } @@ -2943,12 +2931,16 @@ exit: return error; } -otError Mle::HandleDataResponse(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo) +otError Mle::HandleDataResponse(const Message & aMessage, + const Ip6::MessageInfo &aMessageInfo, + const Neighbor * aNeighbor) { otError error; LogMleMessage("Receive Data Response", aMessageInfo.GetPeerAddr()); + VerifyOrExit(aNeighbor && aNeighbor->IsStateValid(), error = OT_ERROR_SECURITY); + error = HandleLeaderData(aMessage, aMessageInfo); if (error != OT_ERROR_NONE) @@ -2966,6 +2958,7 @@ otError Mle::HandleDataResponse(const Message &aMessage, const Ip6::MessageInfo IgnoreReturnValue(Get().StopFastPolls()); } +exit: return error; } @@ -3383,7 +3376,9 @@ exit: return error; } -otError Mle::HandleChildIdResponse(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo) +otError Mle::HandleChildIdResponse(const Message & aMessage, + const Ip6::MessageInfo &aMessageInfo, + const Neighbor * aNeighbor) { OT_UNUSED_VARIABLE(aMessageInfo); @@ -3404,6 +3399,8 @@ otError Mle::HandleChildIdResponse(const Message &aMessage, const Ip6::MessageIn LogMleMessage("Receive Child ID Response", aMessageInfo.GetPeerAddr(), sourceAddress.GetRloc16()); + VerifyOrExit(aNeighbor && aNeighbor->IsStateValid(), error = OT_ERROR_SECURITY); + VerifyOrExit(mAttachState == kAttachStateChildIdRequest); // Leader Data diff --git a/src/core/thread/mle.hpp b/src/core/thread/mle.hpp index b20e0a7d6..7ab63742a 100644 --- a/src/core/thread/mle.hpp +++ b/src/core/thread/mle.hpp @@ -1707,14 +1707,18 @@ private: void ScheduleMessageTransmissionTimer(void); otError HandleAdvertisement(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); - otError HandleChildIdResponse(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + otError HandleChildIdResponse(const Message & aMessage, + const Ip6::MessageInfo &aMessageInfo, + const Neighbor * aNeighbor); otError HandleChildUpdateRequest(const Message & aMessage, const Ip6::MessageInfo &aMessageInfo, Neighbor * aNeighbor); otError HandleChildUpdateResponse(const Message & aMessage, const Ip6::MessageInfo &aMessageInfo, const Neighbor * aNeighbor); - otError HandleDataResponse(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + otError HandleDataResponse(const Message & aMessage, + const Ip6::MessageInfo &aMessageInfo, + const Neighbor * aNeighbor); otError HandleParentResponse(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence); otError HandleAnnounce(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); otError HandleDiscoveryResponse(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index e22d6da4d..973202833 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -2483,7 +2483,9 @@ exit: return error; } -otError MleRouter::HandleDataRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo) +otError MleRouter::HandleDataRequest(const Message & aMessage, + const Ip6::MessageInfo &aMessageInfo, + const Neighbor * aNeighbor) { otError error = OT_ERROR_NONE; TlvRequestTlv tlvRequest; @@ -2494,6 +2496,8 @@ otError MleRouter::HandleDataRequest(const Message &aMessage, const Ip6::Message LogMleMessage("Receive Data Request", aMessageInfo.GetPeerAddr()); + VerifyOrExit(aNeighbor && aNeighbor->IsStateValid(), error = OT_ERROR_SECURITY); + // TLV Request SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kTlvRequest, sizeof(tlvRequest), tlvRequest)); VerifyOrExit(tlvRequest.IsValid() && tlvRequest.GetLength() <= sizeof(tlvs), error = OT_ERROR_PARSE); @@ -4657,11 +4661,16 @@ bool MleRouter::IsSleepyChildSubscribed(const Ip6::Address &aAddress, Child &aCh } #if OPENTHREAD_CONFIG_TIME_SYNC_ENABLE -void MleRouter::HandleTimeSync(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo) +void MleRouter::HandleTimeSync(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, const Neighbor *aNeighbor) { LogMleMessage("Receive Time Sync", aMessageInfo.GetPeerAddr()); + VerifyOrExit(aNeighbor && aNeighbor->IsStateValid()); + Get().HandleTimeSyncMessage(aMessage); + +exit: + return; } otError MleRouter::SendTimeSync(void) diff --git a/src/core/thread/mle_router_ftd.hpp b/src/core/thread/mle_router_ftd.hpp index 92794c13b..bb88a82ba 100644 --- a/src/core/thread/mle_router_ftd.hpp +++ b/src/core/thread/mle_router_ftd.hpp @@ -698,11 +698,11 @@ private: const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence, Neighbor * aNeighbor); - otError HandleDataRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + otError HandleDataRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, const Neighbor *aNeighbor); void HandleNetworkDataUpdateRouter(void); otError HandleDiscoveryRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); #if OPENTHREAD_CONFIG_TIME_SYNC_ENABLE - void HandleTimeSync(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + void HandleTimeSync(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, const Neighbor *aNeighbor); #endif otError ProcessRouteTlv(const RouteTlv &aRoute); diff --git a/src/core/thread/mle_router_mtd.hpp b/src/core/thread/mle_router_mtd.hpp index aa4a3e067..9ca7bf68a 100644 --- a/src/core/thread/mle_router_mtd.hpp +++ b/src/core/thread/mle_router_mtd.hpp @@ -147,14 +147,14 @@ private: { return OT_ERROR_DROP; } - otError HandleDataRequest(const Message &, const Ip6::MessageInfo &) { return OT_ERROR_DROP; } + otError HandleDataRequest(const Message &, const Ip6::MessageInfo &, const Neighbor *) { return OT_ERROR_DROP; } void HandleNetworkDataUpdateRouter(void) {} otError HandleDiscoveryRequest(const Message &, const Ip6::MessageInfo &) { return OT_ERROR_DROP; } void HandlePartitionChange(void) {} void StopAdvertiseTimer(void) {} otError ProcessRouteTlv(const RouteTlv &) { return OT_ERROR_NONE; } #if OPENTHREAD_CONFIG_TIME_SYNC_ENABLE - otError HandleTimeSync(const Message &, const Ip6::MessageInfo &) { return OT_ERROR_DROP; } + otError HandleTimeSync(const Message &, const Ip6::MessageInfo &, const Neighbor *) { return OT_ERROR_DROP; } #endif ChildTable mChildTable;