diff --git a/src/core/net/mdns.cpp b/src/core/net/mdns.cpp index 0f34ea69b..53da710d9 100644 --- a/src/core/net/mdns.cpp +++ b/src/core/net/mdns.cpp @@ -5569,12 +5569,14 @@ void Core::CacheEntry::Init(Instance &aInstance, Type aType) { InstanceLocatorInit::Init(aInstance); - mType = aType; - mInitalQueries = 0; - mQueryPending = false; - mLastQueryTimeValid = false; - mIsActive = false; - mDeleteTime = TimerMilli::GetNow() + kNonActiveDeleteTimeout; + mType = aType; + mContinuousRetry = false; + mQueryPending = false; + mLastQueryTimeValid = false; + mIsActive = false; + mDeleteTime = TimerMilli::GetNow() + kNonActiveDeleteTimeout; + mRetryInterval = 0; + mJitteredRetryInterval = 0; } void Core::CacheEntry::SetIsActive(bool aIsActive) @@ -5612,9 +5614,11 @@ bool Core::CacheEntry::ShouldDelete(TimeMilli aNow) const { return !mIsActive && void Core::CacheEntry::StartInitialQueries(void) { - mInitalQueries = 0; - mLastQueryTimeValid = false; - mLastQueryTime = Get().RandomizeInitialQueryTxTime(); + mContinuousRetry = true; + mRetryInterval = 0; + mJitteredRetryInterval = 0; + mLastQueryTimeValid = false; + mLastQueryTime = Get().RandomizeInitialQueryTxTime(); ScheduleQuery(mLastQueryTime); } @@ -5850,11 +5854,9 @@ void Core::CacheEntry::DetermineNextFireTime(void) { mQueryPending = false; - if (mInitalQueries < kNumberOfInitalQueries) + if (mContinuousRetry) { - uint32_t interval = (mInitalQueries == 0) ? 0 : (1U << (mInitalQueries - 1)) * kInitialQueryInterval; - - ScheduleQuery(mLastQueryTime + interval); + ScheduleQuery(mLastQueryTime + mJitteredRetryInterval); } if (!mIsActive) @@ -5885,6 +5887,26 @@ void Core::CacheEntry::DetermineNextFireTime(void) } } +void Core::CacheEntry::UpdateQueryRetryInterval(void) +{ + uint16_t maxJitter; + + VerifyOrExit(mContinuousRetry); + + mRetryInterval *= kQueryRetryGrowthFactor; + mRetryInterval = Clamp(mRetryInterval, kMinQueryRetryInterval, kMaxQueryRetryInterval); + + // We pre-calculate the jittered retry interval to ensure + // `DetermineNextFireTime()` uses a consistent value. + + maxJitter = ClampToUint16(mRetryInterval / kQueryRetryJitterDivisor); + + mJitteredRetryInterval = Random::NonCrypto::AddJitter(mRetryInterval, maxJitter); + +exit: + return; +} + void Core::CacheEntry::ScheduleTimer(void) { ScheduleFireTimeOn(Get().mCacheTimer); } void Core::CacheEntry::PrepareQuery(CacheContext &aContext) @@ -5926,10 +5948,7 @@ void Core::CacheEntry::PrepareQuery(CacheContext &aContext) mLastQueryTimeValid = true; mLastQueryTime = aContext.GetNow(); - if (mInitalQueries < kNumberOfInitalQueries) - { - mInitalQueries++; - } + UpdateQueryRetryInterval(); // Let the cache entry super-classes update their state // after query was sent. @@ -6477,7 +6496,7 @@ void Core::SrvCache::ProcessResponseRecord(const Message &aMessage, uint16_t aRe if (mRecord.IsPresent()) { - StopInitialQueries(); + StopQueryRetries(); // If not present already, we add a passive `TxtCache` for the // same service name, and an `Ip6AddrCache` for the host name. @@ -6653,7 +6672,7 @@ void Core::TxtCache::ProcessResponseRecord(const Message &aMessage, uint16_t aRe if (mRecord.IsPresent()) { - StopInitialQueries(); + StopQueryRetries(); } ConvertTo(result); @@ -7066,7 +7085,7 @@ void Core::AddrCache::CommitNewResponseEntries(void) } } - StopInitialQueries(); + StopQueryRetries(); // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - // Invoke callbacks if there is any change. @@ -7438,7 +7457,7 @@ void Core::RecordCache::CommitNewEntriesForType(uint16_t aRecordType) if (mRecordType != ResourceRecord::kTypeAny) { - StopInitialQueries(); + StopQueryRetries(); } } diff --git a/src/core/net/mdns.hpp b/src/core/net/mdns.hpp index cacebb908..94bb982e9 100644 --- a/src/core/net/mdns.hpp +++ b/src/core/net/mdns.hpp @@ -896,8 +896,10 @@ private: static constexpr uint8_t kNumberOfAnnounces = 3; static constexpr uint32_t kAnnounceInterval = 1000; // In msec - time between first two announces - static constexpr uint8_t kNumberOfInitalQueries = 3; - static constexpr uint32_t kInitialQueryInterval = 1000; // In msec - time between first two queries + static constexpr uint32_t kMinQueryRetryInterval = Time::kOneSecondInMsec; // In msec + static constexpr uint32_t kMaxQueryRetryInterval = Time::kOneHourInMsec; // In msec + static constexpr uint32_t kQueryRetryGrowthFactor = 2; + static constexpr uint32_t kQueryRetryJitterDivisor = 32; static constexpr uint32_t kMinInitialQueryDelay = 20; // msec static constexpr uint32_t kMaxInitialQueryDelay = 120; // msec @@ -1843,7 +1845,7 @@ private: bool IsActive(void) const { return mIsActive; } bool ShouldDelete(TimeMilli aNow) const; void StartInitialQueries(void); - void StopInitialQueries(void) { mInitalQueries = kNumberOfInitalQueries; } + void StopQueryRetries(void) { mContinuousRetry = false; } Error Add(const ResultCallback &aCallback); void Remove(const ResultCallback &aCallback); void DetermineNextFireTime(void); @@ -1863,7 +1865,7 @@ private: bool ShouldQuery(TimeMilli aNow); void PrepareQuery(CacheContext &aContext); void ProcessExpiredRecords(TimeMilli aNow); - void DetermineNextInitialQueryTime(void); + void UpdateQueryRetryInterval(void); ResultCallback *FindCallbackMatching(const ResultCallback &aCallback); @@ -1871,10 +1873,12 @@ private: template const CacheType &As(void) const { return *static_cast(this); } Type mType; // Cache entry type. - uint8_t mInitalQueries; // Number initial queries sent already. + bool mContinuousRetry : 1; // Whether to continue sending queries. bool mQueryPending : 1; // Whether a query tx request is pending. bool mLastQueryTimeValid : 1; // Whether `mLastQueryTime` is valid. bool mIsActive : 1; // Whether there is any active resolver/browser/querier for this entry. + uint32_t mRetryInterval; // The current query retry interval (in msec). + uint32_t mJitteredRetryInterval; // The current query retry interval with added random jitter (in msec). TimeMilli mNextQueryTime; // The next query tx time when `mQueryPending`. TimeMilli mLastQueryTime; // The last query tx time or the upcoming tx time of first initial query. TimeMilli mDeleteTime; // The time to delete the entry when not `mIsActive`. diff --git a/tests/unit/test_mdns.cpp b/tests/unit/test_mdns.cpp index f5ce29f22..64f86183a 100644 --- a/tests/unit/test_mdns.cpp +++ b/tests/unit/test_mdns.cpp @@ -68,7 +68,7 @@ static constexpr uint16_t kClassMask = 0x7fff; static constexpr uint16_t kStringSize = 300; static constexpr uint16_t kMaxDataSize = 400; static constexpr uint16_t kNumAnnounces = 3; -static constexpr uint16_t kNumInitalQueries = 3; +static constexpr uint16_t kNumInitialQueries = 15; static constexpr uint16_t kNumRefreshQueries = 4; static constexpr bool kCacheFlush = true; static constexpr uint16_t kMdnsPort = 5353; @@ -5825,6 +5825,28 @@ void HandleRecordResultAlternate(otInstance *aInstance, const otMdnsRecordResult HandleRecordResult(aInstance, aResult); } +uint32_t DetermineQueryWaitTime(uint8_t aQueryCount) +{ + uint32_t interval = 125; + + if (aQueryCount == 0) + { + interval = 125; + } + else if (aQueryCount <= 12) + { + interval = (1U << (aQueryCount - 1)) * 1000; + interval += (interval / 32) + 1; + } + else + { + interval = 3600 * 1000; + interval += (interval / 32) + 1; + } + + return interval; +} + //--------------------------------------------------------------------------------------------------------------------- void TestBrowser(void) @@ -5857,11 +5879,11 @@ void TestBrowser(void) sDnsMessages.Clear(); SuccessOrQuit(mdns->StartBrowser(browser)); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount < kNumInitialQueries; queryCount++) { sDnsMessages.Clear(); - AdvanceTime((queryCount == 0) ? 125 : (1U << (queryCount - 1)) * 1000); + AdvanceTime(DetermineQueryWaitTime(queryCount)); VerifyOrQuit(!sDnsMessages.IsEmpty()); dnsMsg = sDnsMessages.GetHead(); @@ -6157,7 +6179,7 @@ void TestBrowser(void) sDnsMessages.Clear(); - SendPtrResponse("_srv._udp.local.", "mysrv._srv._udp.local.", 120, kInAnswerSection); + SendPtrResponse("_srv._udp.local.", "mysrv._srv._udp.local.", 20 * 3600, kInAnswerSection); AdvanceTime(1); @@ -6166,19 +6188,19 @@ void TestBrowser(void) VerifyOrQuit(browseCallback->mServiceType.Matches("_srv._udp")); VerifyOrQuit(!browseCallback->mIsSubType); VerifyOrQuit(browseCallback->mServiceInstance.Matches("mysrv")); - VerifyOrQuit(browseCallback->mTtl == 120); + VerifyOrQuit(browseCallback->mTtl == 20 * 3600); VerifyOrQuit(browseCallback->GetNext() == nullptr); sBrowseCallbacks.Clear(); Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); - Log("Validate initial esquires are still sent and include known-answer"); + Log("Validate initial queries are still sent and include known-answer"); - for (uint8_t queryCount = 1; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 1; queryCount < kNumInitialQueries; queryCount++) { sDnsMessages.Clear(); - AdvanceTime((1U << (queryCount - 1)) * 1000); + AdvanceTime(DetermineQueryWaitTime(queryCount)); VerifyOrQuit(!sDnsMessages.IsEmpty()); dnsMsg = sDnsMessages.GetHead(); @@ -6231,11 +6253,11 @@ void TestSrvResolver(void) sDnsMessages.Clear(); SuccessOrQuit(mdns->StartSrvResolver(resolver)); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount < kNumInitialQueries; queryCount++) { sDnsMessages.Clear(); - AdvanceTime((queryCount == 0) ? 125 : (1U << (queryCount - 1)) * 1000); + AdvanceTime(DetermineQueryWaitTime(queryCount)); VerifyOrQuit(!sDnsMessages.IsEmpty()); dnsMsg = sDnsMessages.GetHead(); @@ -6610,11 +6632,16 @@ void TestSrvResolver(void) sSrvCallbacks.Clear(); - AdvanceTime(15 * 1000); + AdvanceTime(20 * 1000); + + // Initial query intervals use exponential backoff: 1, 2, 4, + // 8, ... seconds. Cumulative send times would be 0, 1, 3, 7, + // 15, 31, ... So within 20 seconds, we should see a total of 5 + // queries. dnsMsg = sDnsMessages.GetHead(); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount <= 4; queryCount++) { VerifyOrQuit(dnsMsg != nullptr); dnsMsg->ValidateHeader(kMulticastQuery, /* Q */ 1, /* Ans */ 0, /* Auth */ 0, /* Addnl */ 0); @@ -6664,11 +6691,11 @@ void TestTxtResolver(void) sDnsMessages.Clear(); SuccessOrQuit(mdns->StartTxtResolver(resolver)); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount < kNumInitialQueries; queryCount++) { sDnsMessages.Clear(); - AdvanceTime((queryCount == 0) ? 125 : (1U << (queryCount - 1)) * 1000); + AdvanceTime(DetermineQueryWaitTime(queryCount)); VerifyOrQuit(!sDnsMessages.IsEmpty()); dnsMsg = sDnsMessages.GetHead(); @@ -7031,11 +7058,16 @@ void TestTxtResolver(void) sTxtCallbacks.Clear(); - AdvanceTime(15 * 1000); + AdvanceTime(20 * 1000); + + // Initial query intervals use exponential backoff: 1, 2, 4, + // 8, ... seconds. Cumulative send times would be 0, 1, 3, 7, + // 15, 31, ... So within 20 seconds, we should see a total of 5 + // queries. dnsMsg = sDnsMessages.GetHead(); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount <= 4; queryCount++) { VerifyOrQuit(dnsMsg != nullptr); dnsMsg->ValidateHeader(kMulticastQuery, /* Q */ 1, /* Ans */ 0, /* Auth */ 0, /* Addnl */ 0); @@ -7085,11 +7117,11 @@ void TestIp6AddrResolver(void) sDnsMessages.Clear(); SuccessOrQuit(mdns->StartIp6AddressResolver(resolver)); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount < kNumInitialQueries; queryCount++) { sDnsMessages.Clear(); - AdvanceTime((queryCount == 0) ? 125 : (1U << (queryCount - 1)) * 1000); + AdvanceTime(DetermineQueryWaitTime(queryCount)); VerifyOrQuit(!sDnsMessages.IsEmpty()); dnsMsg = sDnsMessages.GetHead(); @@ -7539,11 +7571,16 @@ void TestIp6AddrResolver(void) sAddrCallbacks.Clear(); - AdvanceTime(15 * 1000); + AdvanceTime(20 * 1000); + + // Initial query intervals use exponential backoff: 1, 2, 4, + // 8, ... seconds. Cumulative send times would be 0, 1, 3, 7, + // 15, 31, ... So within 20 seconds, we should see a total of 5 + // queries. dnsMsg = sDnsMessages.GetHead(); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount <= 4; queryCount++) { VerifyOrQuit(dnsMsg != nullptr); dnsMsg->ValidateHeader(kMulticastQuery, /* Q */ 1, /* Ans */ 0, /* Auth */ 0, /* Addnl */ 0); @@ -7599,11 +7636,11 @@ void TestRecordQuerier(void) sDnsMessages.Clear(); SuccessOrQuit(mdns->StartRecordQuerier(querier)); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount < kNumInitialQueries; queryCount++) { sDnsMessages.Clear(); - AdvanceTime((queryCount == 0) ? 125 : (1U << (queryCount - 1)) * 1000); + AdvanceTime(DetermineQueryWaitTime(queryCount)); VerifyOrQuit(!sDnsMessages.IsEmpty()); dnsMsg = sDnsMessages.GetHead(); @@ -8088,11 +8125,11 @@ void TestRecordQuerierForAny(void) sDnsMessages.Clear(); SuccessOrQuit(mdns->StartRecordQuerier(querier)); - for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + for (uint8_t queryCount = 0; queryCount < kNumInitialQueries; queryCount++) { sDnsMessages.Clear(); - AdvanceTime((queryCount == 0) ? 125 : (1U << (queryCount - 1)) * 1000); + AdvanceTime(DetermineQueryWaitTime(queryCount)); VerifyOrQuit(!sDnsMessages.IsEmpty()); dnsMsg = sDnsMessages.GetHead();