diff --git a/src/core/api/thread_api.cpp b/src/core/api/thread_api.cpp index 97e5508f0..b5661df64 100644 --- a/src/core/api/thread_api.cpp +++ b/src/core/api/thread_api.cpp @@ -275,7 +275,8 @@ uint32_t otThreadGetKeySequenceCounter(otInstance *aInstance) void otThreadSetKeySequenceCounter(otInstance *aInstance, uint32_t aKeySequenceCounter) { - AsCoreType(aInstance).Get().SetCurrentKeySequence(aKeySequenceCounter, KeyManager::kForceUpdate); + AsCoreType(aInstance).Get().SetCurrentKeySequence( + aKeySequenceCounter, KeyManager::kForceUpdate | KeyManager::kGuardTimerUnchanged); } uint16_t otThreadGetKeySwitchGuardTime(otInstance *aInstance) diff --git a/src/core/mac/mac.cpp b/src/core/mac/mac.cpp index 9ca72fdba..d1fbd0a8e 100644 --- a/src/core/mac/mac.cpp +++ b/src/core/mac/mac.cpp @@ -1641,7 +1641,7 @@ Error Mac::ProcessReceiveSecurity(RxFrame &aFrame, const Address &aSrcAddr, Neig if (keySequence > keyManager.GetCurrentKeySequence()) { - keyManager.SetCurrentKeySequence(keySequence, KeyManager::kApplyKeySwitchGuard); + keyManager.SetCurrentKeySequence(keySequence, KeyManager::kApplySwitchGuard | KeyManager::kResetGuardTimer); } } diff --git a/src/core/thread/key_manager.cpp b/src/core/thread/key_manager.cpp index d38fdb8b5..9f40c9ec2 100644 --- a/src/core/thread/key_manager.cpp +++ b/src/core/thread/key_manager.cpp @@ -368,11 +368,11 @@ void KeyManager::UpdateKeyMaterial(void) #endif } -void KeyManager::SetCurrentKeySequence(uint32_t aKeySequence, KeySequenceUpdateMode aUpdateMode) +void KeyManager::SetCurrentKeySequence(uint32_t aKeySequence, KeySeqUpdateFlags aFlags) { VerifyOrExit(aKeySequence != mKeySequence, Get().SignalIfFirst(kEventThreadKeySeqCounterChanged)); - if (aUpdateMode == kApplyKeySwitchGuard) + if (aFlags & kApplySwitchGuard) { VerifyOrExit(mKeySwitchGuardTimer == 0); } @@ -384,7 +384,11 @@ void KeyManager::SetCurrentKeySequence(uint32_t aKeySequence, KeySequenceUpdateM mMleFrameCounter = 0; ResetKeyRotationTimer(); - mKeySwitchGuardTimer = mKeySwitchGuardTime; + + if (aFlags & kResetGuardTimer) + { + mKeySwitchGuardTimer = mKeySwitchGuardTime; + } Get().Signal(kEventThreadKeySeqCounterChanged); @@ -528,7 +532,7 @@ void KeyManager::CheckForKeyRotation(void) { if (mHoursSinceKeyRotation >= mSecurityPolicy.mRotationTime) { - SetCurrentKeySequence(mKeySequence + 1, kForceUpdate); + SetCurrentKeySequence(mKeySequence + 1, kForceUpdate | kResetGuardTimer); } } diff --git a/src/core/thread/key_manager.hpp b/src/core/thread/key_manager.hpp index 099854c45..56a64acb3 100644 --- a/src/core/thread/key_manager.hpp +++ b/src/core/thread/key_manager.hpp @@ -221,16 +221,24 @@ class KeyManager : public InstanceLocator, private NonCopyable { public: /** - * Determines whether to apply or ignore key switch guard when updating the key sequence. + * Defines bit-flag constants specifying how to handle key sequence update used in `KeySeqUpdateFlags`. + * + */ + enum KeySeqUpdateFlag : uint8_t + { + kApplySwitchGuard = (1 << 0), ///< Apply key switch guard check. + kForceUpdate = (0 << 0), ///< Ignore key switch guard check and forcibly update. + kResetGuardTimer = (1 << 1), ///< On key seq change, reset the guard timer. + kGuardTimerUnchanged = (0 << 1), ///< On key seq change, leave guard timer unchanged. + }; + + /** + * Represents a combination of `KeySeqUpdateFlag` bits. * * Used as input by `SetCurrentKeySequence()`. * */ - enum KeySequenceUpdateMode : uint8_t - { - kApplyKeySwitchGuard, ///< Apply key switch guard check before setting the new key sequence. - kForceUpdate, ///< Ignore key switch guard check and forcibly update the key sequence to new value. - }; + typedef uint8_t KeySeqUpdateFlags; /** * Initializes the object. @@ -342,14 +350,12 @@ public: /** * Sets the current key sequence value. * - * If @p aMode is `kApplyKeySwitchGuard`, the current key switch guard timer is checked and only if it is zero, key - * sequence will be updated. - * * @param[in] aKeySequence The key sequence value. - * @param[in] aUpdateMode Whether or not to apply the key switch guard. + * @param[in] aFlags Specify behavior when updating the key sequence, i.e., whether or not to apply the + * key switch guard or reset guard timer upon change. * */ - void SetCurrentKeySequence(uint32_t aKeySequence, KeySequenceUpdateMode aUpdateMode); + void SetCurrentKeySequence(uint32_t aKeySequence, KeySeqUpdateFlags aFlags); #if OPENTHREAD_CONFIG_RADIO_LINK_TREL_ENABLE /** diff --git a/src/core/thread/mle.cpp b/src/core/thread/mle.cpp index b45d07529..76a397fdd 100644 --- a/src/core/thread/mle.cpp +++ b/src/core/thread/mle.cpp @@ -371,7 +371,8 @@ void Mle::Restore(void) SuccessOrExit(Get().Read(networkInfo)); - Get().SetCurrentKeySequence(networkInfo.GetKeySequence(), KeyManager::kForceUpdate); + Get().SetCurrentKeySequence(networkInfo.GetKeySequence(), + KeyManager::kForceUpdate | KeyManager::kGuardTimerUnchanged); Get().SetMleFrameCounter(networkInfo.GetMleFrameCounter()); Get().SetAllMacFrameCounters(networkInfo.GetMacFrameCounter(), /* aSetIfLarger */ false); @@ -2721,33 +2722,43 @@ void Mle::ProcessKeySequence(RxInfo &aRxInfo) // neighbor. // Otherwise larger key seq MUST NOT be adopted. + bool isNextKeySeq; + KeyManager::KeySeqUpdateFlags flags = 0; + VerifyOrExit(aRxInfo.mKeySequence > Get().GetCurrentKeySequence()); + isNextKeySeq = (aRxInfo.mKeySequence - Get().GetCurrentKeySequence() == 1); + switch (aRxInfo.mClass) { case RxInfo::kAuthoritativeMessage: - Get().SetCurrentKeySequence(aRxInfo.mKeySequence, KeyManager::kForceUpdate); + flags = KeyManager::kForceUpdate; break; case RxInfo::kPeerMessage: - if ((aRxInfo.mNeighbor != nullptr) && aRxInfo.mNeighbor->IsStateValid()) + VerifyOrExit(aRxInfo.IsNeighborStateValid()); + + if (!isNextKeySeq) { - if (aRxInfo.mKeySequence - Get().GetCurrentKeySequence() == 1) - { - Get().SetCurrentKeySequence(aRxInfo.mKeySequence, KeyManager::kApplyKeySwitchGuard); - } - else - { - LogInfo("Large key seq jump in peer class msg from 0x%04x ", aRxInfo.mNeighbor->GetRloc16()); - ReestablishLinkWithNeighbor(*aRxInfo.mNeighbor); - } + LogInfo("Large key seq jump in peer class msg from 0x%04x ", aRxInfo.mNeighbor->GetRloc16()); + ReestablishLinkWithNeighbor(*aRxInfo.mNeighbor); + ExitNow(); } + + flags = KeyManager::kApplySwitchGuard; break; case RxInfo::kUnknown: - break; + ExitNow(); } + if (isNextKeySeq) + { + flags |= KeyManager::kResetGuardTimer; + } + + Get().SetCurrentKeySequence(aRxInfo.mKeySequence, flags); + exit: return; }