From d95b44cda73ba28eadb9173e3932801639af7458 Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Fri, 31 May 2024 12:23:36 -0700 Subject: [PATCH] [mle] add helper methods to `RxMessage` (#10319) This commit adds helper methods to `Mle::RxMessage` for reading and processing specific TLVs: - `ReadModeTlv()`: Reads the Mode TLV as a bit mask and converts it to a `DeviceMode`. - `ReadVersionTlv()`: Reads the Version TLV and verifies that it is at least 1.1 (the minimum supported version). - `ReadAndMatchResponseTlvWith(const TxChallenge &)`: Reads the Response TLV and matches it against a given challenge. --- src/core/thread/mle.cpp | 48 ++++++++++++++++++++++++++++------ src/core/thread/mle.hpp | 3 +++ src/core/thread/mle_router.cpp | 28 ++++++-------------- 3 files changed, 51 insertions(+), 28 deletions(-) diff --git a/src/core/thread/mle.cpp b/src/core/thread/mle.cpp index fc1c73c07..9be918d8a 100644 --- a/src/core/thread/mle.cpp +++ b/src/core/thread/mle.cpp @@ -3124,7 +3124,6 @@ void Mle::HandleParentResponse(RxInfo &aRxInfo) { Error error = kErrorNone; int8_t rss = aRxInfo.mMessage.GetAverageRss(); - RxChallenge response; uint16_t version; uint16_t sourceAddress; LeaderData leaderData; @@ -3144,11 +3143,9 @@ void Mle::HandleParentResponse(RxInfo &aRxInfo) Log(kMessageReceive, kTypeParentResponse, aRxInfo.mMessageInfo.GetPeerAddr(), sourceAddress); - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, version)); - VerifyOrExit(version >= kThreadVersion1p1, error = kErrorParse); + SuccessOrExit(error = aRxInfo.mMessage.ReadVersionTlv(version)); - SuccessOrExit(error = aRxInfo.mMessage.ReadResponseTlv(response)); - VerifyOrExit(response == mParentRequestChallenge, error = kErrorParse); + SuccessOrExit(error = aRxInfo.mMessage.ReadAndMatchResponseTlvWith(mParentRequestChallenge)); aRxInfo.mMessageInfo.GetPeerAddr().GetIid().ConvertToExtAddress(extAddress); @@ -3532,7 +3529,7 @@ void Mle::HandleChildUpdateResponse(RxInfo &aRxInfo) { Error error = kErrorNone; uint8_t status; - uint8_t mode; + DeviceMode mode; RxChallenge response; uint32_t linkFrameCounter; uint32_t mleFrameCounter; @@ -3572,8 +3569,8 @@ void Mle::HandleChildUpdateResponse(RxInfo &aRxInfo) ExitNow(); } - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, mode)); - VerifyOrExit(DeviceMode(mode) == mDeviceMode, error = kErrorDrop); + SuccessOrExit(error = aRxInfo.mMessage.ReadModeTlv(mode)); + VerifyOrExit(mode == mDeviceMode, error = kErrorDrop); switch (mRole) { @@ -4951,6 +4948,29 @@ bool Mle::RxMessage::ContainsTlv(Tlv::Type aTlvType) const return Tlv::FindTlvValueOffset(*this, aTlvType, offset, length) == kErrorNone; } +Error Mle::RxMessage::ReadModeTlv(DeviceMode &aMode) const +{ + Error error; + uint8_t modeBitmask; + + SuccessOrExit(error = Tlv::Find(*this, modeBitmask)); + aMode.Set(modeBitmask); + +exit: + return error; +} + +Error Mle::RxMessage::ReadVersionTlv(uint16_t &aVersion) const +{ + Error error; + + SuccessOrExit(error = Tlv::Find(*this, aVersion)); + VerifyOrExit(aVersion >= kThreadVersion1p1, error = kErrorParse); + +exit: + return error; +} + Error Mle::RxMessage::ReadChallengeOrResponse(uint8_t aTlvType, RxChallenge &aRxChallenge) const { Error error; @@ -4974,6 +4994,18 @@ Error Mle::RxMessage::ReadResponseTlv(RxChallenge &aResponse) const return ReadChallengeOrResponse(Tlv::kResponse, aResponse); } +Error Mle::RxMessage::ReadAndMatchResponseTlvWith(const TxChallenge &aChallenge) const +{ + Error error; + RxChallenge response; + + SuccessOrExit(error = ReadResponseTlv(response)); + VerifyOrExit(response == aChallenge, error = kErrorSecurity); + +exit: + return error; +} + Error Mle::RxMessage::ReadFrameCounterTlvs(uint32_t &aLinkFrameCounter, uint32_t &aMleFrameCounter) const { Error error; diff --git a/src/core/thread/mle.hpp b/src/core/thread/mle.hpp index 2cffb15b3..aa52e8895 100644 --- a/src/core/thread/mle.hpp +++ b/src/core/thread/mle.hpp @@ -1061,8 +1061,11 @@ private: { public: bool ContainsTlv(Tlv::Type aTlvType) const; + Error ReadModeTlv(DeviceMode &aMode) const; + Error ReadVersionTlv(uint16_t &aVersion) const; Error ReadChallengeTlv(RxChallenge &aChallenge) const; Error ReadResponseTlv(RxChallenge &aResponse) const; + Error ReadAndMatchResponseTlvWith(const TxChallenge &aChallenge) const; Error ReadFrameCounterTlvs(uint32_t &aLinkFrameCounter, uint32_t &aMleFrameCounter) const; Error ReadTlvRequestTlv(TlvList &aTlvList) const; Error ReadLeaderDataTlv(LeaderData &aLeaderData) const; diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index 9a4c3dd03..363a40d96 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -686,8 +686,7 @@ void MleRouter::HandleLinkRequest(RxInfo &aRxInfo) SuccessOrExit(error = aRxInfo.mMessage.ReadChallengeTlv(challenge)); - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, version)); - VerifyOrExit(version >= kThreadVersion1p1, error = kErrorParse); + SuccessOrExit(error = aRxInfo.mMessage.ReadVersionTlv(version)); switch (aRxInfo.mMessage.ReadLeaderDataTlv(leaderData)) { @@ -926,8 +925,7 @@ Error MleRouter::HandleLinkAccept(RxInfo &aRxInfo, bool aRequest) RemoveNeighbor(*aRxInfo.mNeighbor); } - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, version)); - VerifyOrExit(version >= kThreadVersion1p1, error = kErrorParse); + SuccessOrExit(error = aRxInfo.mMessage.ReadVersionTlv(version)); SuccessOrExit(error = aRxInfo.mMessage.ReadFrameCounterTlvs(linkFrameCounter, mleFrameCounter)); @@ -1372,7 +1370,6 @@ void MleRouter::HandleParentRequest(RxInfo &aRxInfo) uint8_t scanMask; RxChallenge challenge; Child *child; - uint8_t modeBitmask; DeviceMode mode; Log(kMessageReceive, kTypeParentRequest, aRxInfo.mMessageInfo.GetPeerAddr()); @@ -1402,8 +1399,7 @@ void MleRouter::HandleParentRequest(RxInfo &aRxInfo) aRxInfo.mMessageInfo.GetPeerAddr().GetIid().ConvertToExtAddress(extAddr); - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, version)); - VerifyOrExit(version >= kThreadVersion1p1, error = kErrorParse); + SuccessOrExit(error = aRxInfo.mMessage.ReadVersionTlv(version)); SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, scanMask)); @@ -1437,9 +1433,8 @@ void MleRouter::HandleParentRequest(RxInfo &aRxInfo) #if OPENTHREAD_CONFIG_TIME_SYNC_ENABLE child->SetTimeSyncEnabled(Tlv::Find(aRxInfo.mMessage, nullptr, 0) == kErrorNone); #endif - if (Tlv::Find(aRxInfo.mMessage, modeBitmask) == kErrorNone) + if (aRxInfo.mMessage.ReadModeTlv(mode) == kErrorNone) { - mode.Set(modeBitmask); child->SetDeviceMode(mode); child->SetVersion(version); } @@ -1972,10 +1967,8 @@ void MleRouter::HandleChildIdRequest(RxInfo &aRxInfo) Error error = kErrorNone; Mac::ExtAddress extAddr; uint16_t version; - RxChallenge response; uint32_t linkFrameCounter; uint32_t mleFrameCounter; - uint8_t modeBitmask; DeviceMode mode; uint32_t timeout; TlvList tlvList; @@ -1995,11 +1988,9 @@ void MleRouter::HandleChildIdRequest(RxInfo &aRxInfo) child = mChildTable.FindChild(extAddr, Child::kInStateAnyExceptInvalid); VerifyOrExit(child != nullptr, error = kErrorAlready); - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, version)); - VerifyOrExit(version >= kThreadVersion1p1, error = kErrorParse); + SuccessOrExit(error = aRxInfo.mMessage.ReadVersionTlv(version)); - SuccessOrExit(error = aRxInfo.mMessage.ReadResponseTlv(response)); - VerifyOrExit(response == child->GetChallenge(), error = kErrorSecurity); + SuccessOrExit(error = aRxInfo.mMessage.ReadAndMatchResponseTlvWith(child->GetChallenge())); Get().RemoveMessages(*child, Message::kSubTypeMleGeneral); Get().RemoveMessages(*child, Message::kSubTypeMleChildIdRequest); @@ -2008,8 +1999,7 @@ void MleRouter::HandleChildIdRequest(RxInfo &aRxInfo) SuccessOrExit(error = aRxInfo.mMessage.ReadFrameCounterTlvs(linkFrameCounter, mleFrameCounter)); - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, modeBitmask)); - mode.Set(modeBitmask); + SuccessOrExit(error = aRxInfo.mMessage.ReadModeTlv(mode)); SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, timeout)); @@ -2140,7 +2130,6 @@ void MleRouter::HandleChildUpdateRequest(RxInfo &aRxInfo) { Error error = kErrorNone; Mac::ExtAddress extAddr; - uint8_t modeBitmask; DeviceMode mode; RxChallenge challenge; LeaderData leaderData; @@ -2154,8 +2143,7 @@ void MleRouter::HandleChildUpdateRequest(RxInfo &aRxInfo) Log(kMessageReceive, kTypeChildUpdateRequestOfChild, aRxInfo.mMessageInfo.GetPeerAddr()); - SuccessOrExit(error = Tlv::Find(aRxInfo.mMessage, modeBitmask)); - mode.Set(modeBitmask); + SuccessOrExit(error = aRxInfo.mMessage.ReadModeTlv(mode)); switch (aRxInfo.mMessage.ReadChallengeTlv(challenge)) {