From 4aab013e063e78353a41a75fdc300780498edea5 Mon Sep 17 00:00:00 2001 From: Jonathan Hui Date: Fri, 25 Oct 2019 10:27:32 -0700 Subject: [PATCH] [mle] fix MleRouter::HandleLinkAccept message validation (#4297) This commit moves the message security check to the top of MleRouter::HandleLinkAccept. --- src/core/thread/mle.cpp | 4 +- src/core/thread/mle_router.cpp | 83 +++++++++++++++--------------- src/core/thread/mle_router_ftd.hpp | 9 +++- src/core/thread/mle_router_mtd.hpp | 12 +++-- 4 files changed, 60 insertions(+), 48 deletions(-) diff --git a/src/core/thread/mle.cpp b/src/core/thread/mle.cpp index d2536f5e3..1df7d228e 100644 --- a/src/core/thread/mle.cpp +++ b/src/core/thread/mle.cpp @@ -2770,11 +2770,11 @@ void Mle::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageIn break; case Header::kCommandLinkAccept: - Get().HandleLinkAccept(aMessage, aMessageInfo, keySequence); + Get().HandleLinkAccept(aMessage, aMessageInfo, keySequence, neighbor); break; case Header::kCommandLinkAcceptAndRequest: - Get().HandleLinkAcceptAndRequest(aMessage, aMessageInfo, keySequence); + Get().HandleLinkAcceptAndRequest(aMessage, aMessageInfo, keySequence, neighbor); break; case Header::kCommandAdvertisement: diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index 6975d1a74..e22d6da4d 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -768,21 +768,24 @@ exit: otError MleRouter::HandleLinkAccept(const Message & aMessage, const Ip6::MessageInfo &aMessageInfo, - uint32_t aKeySequence) + uint32_t aKeySequence, + Neighbor * aNeighbor) { - return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, false); + return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, aNeighbor, false); } otError MleRouter::HandleLinkAcceptAndRequest(const Message & aMessage, const Ip6::MessageInfo &aMessageInfo, - uint32_t aKeySequence) + uint32_t aKeySequence, + Neighbor * aNeighbor) { - return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, true); + return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, aNeighbor, true); } otError MleRouter::HandleLinkAccept(const Message & aMessage, const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence, + Neighbor * aNeighbor, bool aRequest) { static const uint8_t dataRequestTlvs[] = {Tlv::kNetworkData}; @@ -790,7 +793,6 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage, otError error = OT_ERROR_NONE; const otThreadLinkInfo *linkInfo = static_cast(aMessageInfo.GetLinkInfo()); Router * router; - Neighbor * neighbor; Neighbor::State neighborState; Mac::ExtAddress macAddr; VersionTlv version; @@ -803,10 +805,6 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage, RouteTlv route; LeaderDataTlv leaderData; LinkMarginTlv linkMargin; - ChallengeTlv challenge; - TlvRequestTlv tlvRequest; - - aMessageInfo.GetPeerAddr().ToExtAddress(macAddr); // Source Address SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kSourceAddress, sizeof(sourceAddress), sourceAddress)); @@ -821,20 +819,45 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage, LogMleMessage("Receive Link Accept", aMessageInfo.GetPeerAddr(), sourceAddress.GetRloc16()); } - // Version - SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kVersion, sizeof(version), version)); - VerifyOrExit(version.IsValid(), error = OT_ERROR_PARSE); + VerifyOrExit(IsActiveRouter(sourceAddress.GetRloc16()), error = OT_ERROR_PARSE); + + routerId = GetRouterId(sourceAddress.GetRloc16()); + router = mRouterTable.GetRouter(routerId); + neighborState = (router != NULL) ? router->GetState() : Neighbor::kStateInvalid; // Response SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kResponse, sizeof(response), response)); VerifyOrExit(response.IsValid(), error = OT_ERROR_PARSE); - // Remove stale neighbors - if ((neighbor = GetNeighbor(macAddr)) != NULL && neighbor->GetRloc16() != sourceAddress.GetRloc16()) + // verify response + switch (neighborState) { - RemoveNeighbor(*neighbor); + case Neighbor::kStateLinkRequest: + VerifyOrExit(memcmp(router->GetChallenge(), response.GetResponse(), router->GetChallengeSize()) == 0, + error = OT_ERROR_SECURITY); + break; + + case Neighbor::kStateInvalid: + VerifyOrExit((mChallengeTimeout > 0) && (memcmp(mChallenge, response.GetResponse(), sizeof(mChallenge)) == 0), + error = OT_ERROR_SECURITY); + + case Neighbor::kStateValid: + break; + + default: + ExitNow(error = OT_ERROR_SECURITY); } + // Remove stale neighbors + if (aNeighbor && aNeighbor->GetRloc16() != sourceAddress.GetRloc16()) + { + RemoveNeighbor(*aNeighbor); + } + + // Version + SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kVersion, sizeof(version), version)); + VerifyOrExit(version.IsValid(), error = OT_ERROR_PARSE); + // Link-Layer Frame Counter SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kLinkFrameCounter, sizeof(linkFrameCounter), linkFrameCounter)); VerifyOrExit(linkFrameCounter.IsValid(), error = OT_ERROR_PARSE); @@ -863,32 +886,6 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage, linkMargin.SetLinkMargin(0); } - VerifyOrExit(IsActiveRouter(sourceAddress.GetRloc16()), error = OT_ERROR_PARSE); - - routerId = GetRouterId(sourceAddress.GetRloc16()); - router = mRouterTable.GetRouter(routerId); - neighborState = (router != NULL) ? router->GetState() : Neighbor::kStateInvalid; - - // verify response - switch (neighborState) - { - case Neighbor::kStateLinkRequest: - VerifyOrExit(memcmp(router->GetChallenge(), response.GetResponse(), router->GetChallengeSize()) == 0, - error = OT_ERROR_SECURITY); - break; - - case Neighbor::kStateInvalid: - VerifyOrExit((mChallengeTimeout > 0) && (memcmp(mChallenge, response.GetResponse(), sizeof(mChallenge)) == 0), - error = OT_ERROR_SECURITY); - break; - - case Neighbor::kStateValid: - break; - - default: - ExitNow(error = OT_ERROR_INVALID_STATE); - } - switch (mRole) { case OT_DEVICE_ROLE_DISABLED: @@ -968,6 +965,7 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage, } // finish link synchronization + aMessageInfo.GetPeerAddr().ToExtAddress(macAddr); router->SetExtAddress(macAddr); router->SetRloc16(sourceAddress.GetRloc16()); router->SetLinkFrameCounter(linkFrameCounter.GetFrameCounter()); @@ -986,6 +984,9 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage, if (aRequest) { + ChallengeTlv challenge; + TlvRequestTlv tlvRequest; + // Challenge SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kChallenge, sizeof(challenge), challenge)); VerifyOrExit(challenge.IsValid(), error = OT_ERROR_PARSE); diff --git a/src/core/thread/mle_router_ftd.hpp b/src/core/thread/mle_router_ftd.hpp index dd36e85dd..92794c13b 100644 --- a/src/core/thread/mle_router_ftd.hpp +++ b/src/core/thread/mle_router_ftd.hpp @@ -675,14 +675,19 @@ private: void HandleDetachStart(void); otError HandleChildStart(AttachMode aMode); otError HandleLinkRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, Neighbor *aNeighbor); - otError HandleLinkAccept(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence); otError HandleLinkAccept(const Message & aMessage, const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence, + Neighbor * aNeighbor); + otError HandleLinkAccept(const Message & aMessage, + const Ip6::MessageInfo &aMessageInfo, + uint32_t aKeySequence, + Neighbor * aNeighbor, bool aRequest); otError HandleLinkAcceptAndRequest(const Message & aMessage, const Ip6::MessageInfo &aMessageInfo, - uint32_t aKeySequence); + uint32_t aKeySequence, + Neighbor * aNeighbor); otError HandleAdvertisement(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); otError HandleParentRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); otError HandleChildIdRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence); diff --git a/src/core/thread/mle_router_mtd.hpp b/src/core/thread/mle_router_mtd.hpp index 65492677c..aa4a3e067 100644 --- a/src/core/thread/mle_router_mtd.hpp +++ b/src/core/thread/mle_router_mtd.hpp @@ -130,9 +130,15 @@ private: void HandleDetachStart(void) {} otError HandleChildStart(AttachMode) { return OT_ERROR_NONE; } otError HandleLinkRequest(const Message &, const Ip6::MessageInfo &, Neighbor *) { return OT_ERROR_DROP; } - otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t) { return OT_ERROR_DROP; } - otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t, bool) { return OT_ERROR_DROP; } - otError HandleLinkAcceptAndRequest(const Message &, const Ip6::MessageInfo &, uint32_t) { return OT_ERROR_DROP; } + otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t, Neighbor *) { return OT_ERROR_DROP; } + otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t, Neighbor *, bool) + { + return OT_ERROR_DROP; + } + otError HandleLinkAcceptAndRequest(const Message &, const Ip6::MessageInfo &, uint32_t, Neighbor *) + { + return OT_ERROR_DROP; + } otError HandleAdvertisement(const Message &, const Ip6::MessageInfo &) { return OT_ERROR_DROP; } otError HandleParentRequest(const Message &, const Ip6::MessageInfo &) { return OT_ERROR_DROP; } otError HandleChildIdRequest(const Message &, const Ip6::MessageInfo &, uint32_t) { return OT_ERROR_DROP; }