diff --git a/include/openthread/instance.h b/include/openthread/instance.h index 62b578dac..efa9bab18 100644 --- a/include/openthread/instance.h +++ b/include/openthread/instance.h @@ -52,7 +52,7 @@ extern "C" { * * @note This number versions both OpenThread platform and user APIs. */ -#define OPENTHREAD_API_VERSION (484) +#define OPENTHREAD_API_VERSION (485) /** * @addtogroup api-instance diff --git a/include/openthread/mdns.h b/include/openthread/mdns.h index c77b4c71d..3193aa24c 100644 --- a/include/openthread/mdns.h +++ b/include/openthread/mdns.h @@ -550,6 +550,40 @@ typedef otPlatDnssdAddressAndTtl otMdnsAddressAndTtl; */ typedef otPlatDnssdAddressResult otMdnsAddressResult; +/** + * Represents a record query result. + */ +typedef struct otMdnsRecordResult +{ + const char *mFirstLabel; ///< The first label of the name to be queried. + const char *mNextLabels; ///< The rest of the name labels. Does not include domain name. Can be NULL. + uint16_t mRecordType; ///< The record type. + const uint8_t *mRecordData; ///< The record data bytes. + uint16_t mRecordDataLength; ///< Number of bytes in record data. + uint32_t mTtl; ///< TTL in seconds. Zero TTL indicates removal the data. + uint32_t mInfraIfIndex; ///< The infrastructure network interface index. +} otMdnsRecordResult; + +/** + * Represents the callback function used to report a record querier result. + * + * @param[in] aInstance The OpenThread instance. + * @param[in] aResult The record querier result. + */ +typedef void (*otMdnsRecordCallback)(otInstance *aInstance, const otMdnsRecordResult *aResult); + +/** + * Represents a record querier. + */ +typedef struct otMdnsRecordQuerier +{ + const char *mFirstLabel; ///< The first label of the name to be queried. MUST NOT be NULL. + const char *mNextLabels; ///< The rest of name labels, excluding domain name. Can be NULL. + uint16_t mRecordType; ///< The record type to query. + uint32_t mInfraIfIndex; ///< The infrastructure network interface index. + otMdnsRecordCallback mCallback; ///< The callback to report result. +} otMdnsRecordQuerier; + /** * Starts a service browser. * @@ -734,6 +768,57 @@ otError otMdnsStartIp4AddressResolver(otInstance *aInstance, const otMdnsAddress */ otError otMdnsStopIp4AddressResolver(otInstance *aInstance, const otMdnsAddressResolver *aResolver); +/** + * Starts a record querier. + * + * Initiates a continuous query for a given `mRecordType` as specified in @p aQuerier. The queried name is specified + * by the combination of `mFirstLabel` and `mNextLabels` (optional rest of the labels) in @p aQuerier. The + * `mFirstLabel` MUST be non-NULL but `mNextLabels` can be `NULL` if there are no other labels. The `mNextLabels` + * MUST NOT include the domain name. The reason for a separate first label is to allow it to include a dot `.` + * character (as allowed for service instance labels). + * + * Discovered results are reported through the `mCallback` function in @p aQuerier, providing the raw record + * data bytes. A removed record data is indicated with a TTL value of zero. The callback may be invoked immediately + * with cached information (if available) and potentially before this function returns. When cached results are used, + * the reported TTL value will reflect the original TTL from the last received response. + * + * Multiple querier instances can be started for the same name, provided they use different callback functions. + * + * The record querier MUST not be used for record types PTR, SRV, TXT, A, and AAAA. Otherwise, `OT_ERROR_INVALID_ARGS` + * will be returned. For these, browsers/resolvers can be used. This design is intentional to enable the implementation + * of an "opportunistic cache mechanism", where, depending on currently active service browsers/resolvers, the mDNS + * implementation will also monitor and cache related records (e.g., when a service is resolved, the address records + * associated with its host name are cached even if there is no active address resolver for this hostname). + * + * The @p aQuerier and all its contained information (strings) are only valid during this call. The platform MUST save + * a copy of the information if it wants to retain the information after returning from this function. + * + * @param[in] aInstance The OpenThread instance. + * @param[in] aQuerier The record querier to be started. + * + * @retval OT_ERROR_NONE Record @p aQuerier started successfully. + * @retval OT_ERROR_INVALID_STATE mDNS module is not enabled. + * @retval OT_ERROR_ALREADY An identical querier (same name, record type, and callback) is already active. + * @retval OT_ERROR_INVALID_ARGS The `mRecordType` in @p aQuerier is invalid. MUST use browser/resolvers. + */ +otError otMdnsStartRecordQuerier(otInstance *aInstance, const otMdnsRecordQuerier *aQuerier); + +/** + * Stops a record querier. + * + * No action is performed if no matching querier with the same name and callback is currently active. + * + * The @p aQuerier and all its contained information (strings) are only valid during this call. The platform MUST save + * a copy of the information if it wants to retain the information after returning from this function. + * + * @param[in] aInstance The OpenThread instance. + * @param[in] aQuerier The record querier to be stopped. + * + * @retval OT_ERROR_NONE Querier stopped successfully. + * @retval OT_ERROR_INVALID_STATE mDNS module is not enabled. + */ +otError otMdnsStopRecordQuerier(otInstance *aInstance, const otMdnsRecordQuerier *aQuerier); + /** * Represents additional information about a browser/resolver and its cached results. */ @@ -862,6 +947,30 @@ otError otMdnsGetNextIp4AddressResolver(otInstance *aInstance, otMdnsAddressResolver *aResolver, otMdnsCacheInfo *aInfo); +/** + * Iterates over record querier entries. + * + * Requires `OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE`. + * + * On success, @p aQuerier is populated with information about the next querier . The `mCallback` field is always + * set to `NULL` as there may be multiple active querier with different callbacks. Other pointers within the + * `otMdnsRecordQuerier` structure remain valid until the next call to any OpenThread stack's public or platform + * API/callback. + * + * @param[in] aInstance The OpenThread instance. + * @param[in] aIterator Pointer to the iterator. + * @param[out] aQuerier Pointer to an `otMdnsRecordQuerier` to return the information about the next one. + * @param[out] aInfo Pointer to an `otMdnsCacheInfo` to return additional information. + * + * @retval OT_ERROR_NONE @p aQuerier, @p aInfo, & @p aIterator are updated successfully. + * @retval OT_ERROR_NOT_FOUND Reached the end of the list. + * @retval OT_ERROR_INVALID_ARG @p aIterator is not valid. + */ +otError otMdnsGetNextRecordQuerier(otInstance *aInstance, + otMdnsIterator *aIterator, + otMdnsRecordQuerier *aQuerier, + otMdnsCacheInfo *aInfo); + /** * @} */ diff --git a/src/cli/cli_mdns.cpp b/src/cli/cli_mdns.cpp index 8a4e08ee4..e04adc4b4 100644 --- a/src/cli/cli_mdns.cpp +++ b/src/cli/cli_mdns.cpp @@ -879,6 +879,64 @@ void Mdns::HandleIp4AddressResult(otInstance *aInstance, const otMdnsAddressResu Interpreter::GetInterpreter().mMdns.HandleAddressResult(*aResult, kIp4Address); } +template <> otError Mdns::Process(Arg aArgs[]) +{ + // mdns recordquerier start|stop [] + + otError error; + otMdnsRecordQuerier querier; + bool isStart; + + ClearAllBytes(querier); + + SuccessOrExit(error = ParseStartOrStop(aArgs[0], isStart)); + + SuccessOrExit(error = aArgs[1].ParseAsUint16(querier.mRecordType)); + + VerifyOrExit(!aArgs[2].IsEmpty(), error = OT_ERROR_INVALID_ARGS); + querier.mFirstLabel = aArgs[2].GetCString(); + + if (!aArgs[3].IsEmpty()) + { + querier.mNextLabels = aArgs[3].GetCString(); + VerifyOrExit(aArgs[4].IsEmpty(), error = OT_ERROR_INVALID_ARGS); + } + + querier.mInfraIfIndex = mInfraIfIndex; + querier.mCallback = HandleRecordResult; + + if (isStart) + { + error = otMdnsStartRecordQuerier(GetInstancePtr(), &querier); + } + else + { + error = otMdnsStopRecordQuerier(GetInstancePtr(), &querier); + } + +exit: + return error; +} + +void Mdns::HandleRecordResult(otInstance *aInstance, const otMdnsRecordResult *aResult) +{ + OT_UNUSED_VARIABLE(aInstance); + + Interpreter::GetInterpreter().mMdns.HandleRecordResult(*aResult); +} + +void Mdns::HandleRecordResult(const otMdnsRecordResult &aResult) +{ + OutputLine("mDNS result for record %u and name %s %s", aResult.mRecordType, aResult.mFirstLabel, + aResult.mNextLabels == nullptr ? "" : aResult.mNextLabels); + + OutputFormat(kIndentSize, "data: "); + OutputBytesLine(aResult.mRecordData, aResult.mRecordDataLength); + + OutputLine(kIndentSize, "ttl: %lu", ToUlong(aResult.mTtl)); + OutputLine(kIndentSize, "if-index: %lu", ToUlong(aResult.mInfraIfIndex)); +} + #if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE template <> otError Mdns::Process(Arg aArgs[]) @@ -1083,6 +1141,46 @@ exit: return error; } +template <> otError Mdns::Process(Arg aArgs[]) +{ + // mdns recordqueriers + + otError error; + otMdnsIterator *iterator = nullptr; + otMdnsCacheInfo info; + otMdnsRecordQuerier querier; + + VerifyOrExit(aArgs[0].IsEmpty(), error = OT_ERROR_INVALID_ARGS); + + iterator = otMdnsAllocateIterator(GetInstancePtr()); + VerifyOrExit(iterator != nullptr, error = OT_ERROR_NO_BUFS); + + while (true) + { + error = otMdnsGetNextRecordQuerier(GetInstancePtr(), iterator, &querier, &info); + + if (error == OT_ERROR_NOT_FOUND) + { + error = OT_ERROR_NONE; + ExitNow(); + } + + SuccessOrExit(error); + + OutputLine("Record querier for type %u and name %s %s", querier.mRecordType, querier.mFirstLabel, + querier.mNextLabels == nullptr ? "" : querier.mNextLabels); + OutputCacheInfo(info); + } + +exit: + if (iterator != nullptr) + { + otMdnsFreeIterator(GetInstancePtr(), iterator); + } + + return error; +} + #endif // OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE otError Mdns::Process(Arg aArgs[]) @@ -1110,6 +1208,10 @@ otError Mdns::Process(Arg aArgs[]) #if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE CmdEntry("ip6resolvers"), CmdEntry("keys"), +#endif + CmdEntry("recordquerier"), +#if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE + CmdEntry("recordqueriers"), #endif CmdEntry("register"), #if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE diff --git a/src/cli/cli_mdns.hpp b/src/cli/cli_mdns.hpp index a2fb03e46..1b8555cab 100644 --- a/src/cli/cli_mdns.hpp +++ b/src/cli/cli_mdns.hpp @@ -117,6 +117,7 @@ private: void HandleSrvResult(const otMdnsSrvResult &aResult); void HandleTxtResult(const otMdnsTxtResult &aResult); void HandleAddressResult(const otMdnsAddressResult &aResult, IpAddressType aType); + void HandleRecordResult(const otMdnsRecordResult &aResult); static otError ParseStartOrStop(const Arg &aArg, bool &aIsStart); static void HandleRegisterationDone(otInstance *aInstance, otMdnsRequestId aRequestId, otError aError); @@ -125,6 +126,7 @@ private: static void HandleTxtResult(otInstance *aInstance, const otMdnsTxtResult *aResult); static void HandleIp6AddressResult(otInstance *aInstance, const otMdnsAddressResult *aResult); static void HandleIp4AddressResult(otInstance *aInstance, const otMdnsAddressResult *aResult); + static void HandleRecordResult(otInstance *aInstance, const otMdnsRecordResult *aResult); static otError ParseServiceArgs(Arg aArgs[], otMdnsService &aService, Buffers &aBuffers); diff --git a/src/core/api/mdns_api.cpp b/src/core/api/mdns_api.cpp index 7d737211f..0085fc750 100644 --- a/src/core/api/mdns_api.cpp +++ b/src/core/api/mdns_api.cpp @@ -228,6 +228,20 @@ otError otMdnsStopIp4AddressResolver(otInstance *aInstance, const otMdnsAddressR return AsCoreType(aInstance).Get().StopIp4AddressResolver(*aResolver); } +otError otMdnsStartRecordQuerier(otInstance *aInstance, const otMdnsRecordQuerier *aQuerier) +{ + AssertPointerIsNotNull(aQuerier); + + return AsCoreType(aInstance).Get().StartRecordQuerier(*aQuerier); +} + +otError otMdnsStopRecordQuerier(otInstance *aInstance, const otMdnsRecordQuerier *aQuerier) +{ + AssertPointerIsNotNull(aQuerier); + + return AsCoreType(aInstance).Get().StopRecordQuerier(*aQuerier); +} + #if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE otError otMdnsGetNextBrowser(otInstance *aInstance, @@ -290,6 +304,18 @@ otError otMdnsGetNextIp4AddressResolver(otInstance *aInstance, return AsCoreType(aInstance).Get().GetNextIp4AddressResolver(*aIterator, *aResolver, *aInfo); } +otError otMdnsGetNextRecordQuerier(otInstance *aInstance, + otMdnsIterator *aIterator, + otMdnsRecordQuerier *aQuerier, + otMdnsCacheInfo *aInfo) +{ + AssertPointerIsNotNull(aIterator); + AssertPointerIsNotNull(aQuerier); + AssertPointerIsNotNull(aInfo); + + return AsCoreType(aInstance).Get().GetNextRecordQuerier(*aIterator, *aQuerier, *aInfo); +} + #endif // OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE #endif // OPENTHREAD_CONFIG_MULTICAST_DNS_ENABLE && OPENTHREAD_CONFIG_MULTICAST_DNS_PUBLIC_API_ENABLE diff --git a/src/core/net/dns_types.cpp b/src/core/net/dns_types.cpp index 41ed989a5..6c037e27b 100644 --- a/src/core/net/dns_types.cpp +++ b/src/core/net/dns_types.cpp @@ -111,8 +111,11 @@ bool Name::Matches(const char *aFirstLabel, const char *aLabels, const char *aDo VerifyOrExit(matches); } - matches = CompareAndSkipLabels(namePtr, aLabels, kLabelSeparatorChar); - VerifyOrExit(matches); + if (aLabels != nullptr) + { + matches = CompareAndSkipLabels(namePtr, aLabels, kLabelSeparatorChar); + VerifyOrExit(matches); + } matches = CompareAndSkipLabels(namePtr, aDomain, kNullChar); } @@ -125,7 +128,11 @@ bool Name::Matches(const char *aFirstLabel, const char *aLabels, const char *aDo SuccessOrExit(CompareLabel(*mMessage, offset, aFirstLabel)); } - SuccessOrExit(CompareMultipleLabels(*mMessage, offset, aLabels)); + if (aLabels != nullptr) + { + SuccessOrExit(CompareMultipleLabels(*mMessage, offset, aLabels)); + } + SuccessOrExit(CompareName(*mMessage, offset, aDomain)); matches = true; } diff --git a/src/core/net/dns_types.hpp b/src/core/net/dns_types.hpp index 23f16bc76..156677bb0 100644 --- a/src/core/net/dns_types.hpp +++ b/src/core/net/dns_types.hpp @@ -606,15 +606,15 @@ public: * @p aFirstLabel can be `nullptr` if not needed. But if non-null, it is treated as a single label and can itself * include dot `.` character. * - * The @p aLabels MUST NOT be `nullptr` and MUST follow "..", i.e., a sequence of one or - * more labels separated by dot '.' char, and it MUST NOT end with dot `.`. + * The @p aLabels can be `nullptr`. If it is provided it MUST follow "..", i.e., a + * sequence of one or more labels separated by dot '.' char, and it MUST NOT end with dot `.`. * * @p aDomain MUST NOT be `nullptr` and MUST have at least one label and MUST always end with a dot `.` character. * * If the above conditions are not satisfied, the behavior of this method is undefined. * * @param[in] aFirstLabel A first label to check. Can be `nullptr`. - * @param[in] aLabels A string of dot separated labels, MUST NOT end with dot. + * @param[in] aLabels A string of dot separated labels, MUST NOT end with dot. Can be `nullptr` * @param[in] aDomain Domain name. MUST end with dot. * * @retval TRUE The name matches the given components. diff --git a/src/core/net/mdns.cpp b/src/core/net/mdns.cpp index 977513567..bd7774bf4 100644 --- a/src/core/net/mdns.cpp +++ b/src/core/net/mdns.cpp @@ -115,6 +115,7 @@ Error Core::SetEnabled(bool aEnable, uint32_t aInfraIfIndex) mTxtCacheList.Clear(); mIp6AddrCacheList.Clear(); mIp4AddrCacheList.Clear(); + mRecordCacheList.Clear(); mCacheTimer.Stop(); } @@ -243,6 +244,11 @@ Error Core::GetNextIp4AddressResolver(Iterator &aIterator, AddressResolver &aRes return static_cast(aIterator).GetNextIp4AddressResolver(aResolver, aInfo); } +Error Core::GetNextRecordQuerier(Iterator &aIterator, RecordQuerier &aQuerier, CacheInfo &aInfo) const +{ + return static_cast(aIterator).GetNextRecordQuerier(aQuerier, aInfo); +} + #endif // OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE void Core::InvokeConflictCallback(const char *aName, const char *aServiceType) @@ -4092,6 +4098,16 @@ void Core::RxMessage::ProcessResponse(void) addrCache.CommitNewResponseEntries(); } } + + if (!Get().mRecordCacheList.IsEmpty()) + { + IterateOnAllRecordsInResponse(&RxMessage::ProcessOtherRecord); + + for (RecordCache &recordCache : Get().mRecordCacheList) + { + recordCache.CommitNewResponseEntries(); + } + } } void Core::RxMessage::IterateOnAllRecordsInResponse(RecordProcessor aRecordProcessor) @@ -4225,6 +4241,19 @@ exit: return; } +void Core::RxMessage::ProcessOtherRecord(const Name &aName, const ResourceRecord &aRecord, uint16_t aRecordOffset) +{ + RecordCache *recordCache; + + recordCache = Get().mRecordCacheList.FindMatching(aName, aRecord.GetType()); + VerifyOrExit(recordCache != nullptr); + + recordCache->ProcessResponseRecord(*mMessagePtr, aRecord, aRecordOffset); + +exit: + return; +} + //--------------------------------------------------------------------------------------------------------------------- // Core::RxMessage::Question @@ -4429,6 +4458,9 @@ void Core::TxMessageHistory::HandleTimer(void) mTimer.FireAtIfEarlier(nextTime); } +//--------------------------------------------------------------------------------------------------------------------- +// Core + template Error Core::Start(const BrowserResolverType &aBrowserOrResolver) { @@ -4489,6 +4521,28 @@ Error Core::StartIp6AddressResolver(const AddressResolver &aResolver) return Start(aResolver); } +Error Core::StartRecordQuerier(const RecordQuerier &aQuerier) +{ + Error error; + + switch (aQuerier.mRecordType) + { + case ResourceRecord::kTypePtr: + case ResourceRecord::kTypeSrv: + case ResourceRecord::kTypeTxt: + case ResourceRecord::kTypeAaaa: + case ResourceRecord::kTypeA: + error = kErrorInvalidArgs; + break; + + default: + error = Start(aQuerier); + break; + } + + return error; +} + Error Core::StopIp6AddressResolver(const AddressResolver &aResolver) { return Stop(aResolver); @@ -4504,6 +4558,8 @@ Error Core::StopIp4AddressResolver(const AddressResolver &aResolver) return Stop(aResolver); } +Error Core::StopRecordQuerier(const RecordQuerier &aQuerier) { return Stop(aQuerier); } + void Core::AddPassiveSrvTxtCache(const char *aServiceInstance, const char *aServiceType) { ServiceName serviceName(aServiceInstance, aServiceType); @@ -4548,6 +4604,7 @@ void Core::HandleCacheTimer(void) mTxtCacheList.RemoveAndFreeAllMatching(expireChecker); mIp6AddrCacheList.RemoveAndFreeAllMatching(expireChecker); mIp4AddrCacheList.RemoveAndFreeAllMatching(expireChecker); + mRecordCacheList.RemoveAndFreeAllMatching(expireChecker); // Process cache types in a specific order to optimize name // compression when constructing query messages. @@ -4577,6 +4634,11 @@ void Core::HandleCacheTimer(void) addrCache.HandleTimer(context); } + for (RecordCache &recordCache : mRecordCacheList) + { + recordCache.HandleTimer(context); + } + context.mQueryMessage.Send(); mCacheTimer.FireAtIfEarlier(context.mNextFireTime); @@ -4612,6 +4674,11 @@ void Core::HandleCacheTask(void) { addrCache.ClearEmptyCallbacks(); } + + for (RecordCache &recordCache : mRecordCacheList) + { + recordCache.ClearEmptyCallbacks(); + } } TimeMilli Core::RandomizeFirstProbeTxTime(void) @@ -4681,6 +4748,14 @@ void Core::ResultCallback::Invoke(Instance &aInstance, const AddressResult &aRes } } +void Core::ResultCallback::Invoke(Instance &aInstance, const RecordResult &aResult) const +{ + if (mSharedCallback.mRecord != nullptr) + { + mSharedCallback.mRecord(&aInstance, &aResult); + } +} + //--------------------------------------------------------------------------------------------------------------------- // Core::CacheContext @@ -4915,6 +4990,7 @@ Error Core::CacheEntry::Add(const ResultCallback &aCallback) switch (mType) { case kBrowseCache: + case kRecordCache: shouldStart = true; break; case kSrvCache: @@ -4954,6 +5030,9 @@ Error Core::CacheEntry::Add(const ResultCallback &aCallback) case kIp4AddrCache: As().ReportResultsTo(*callback); break; + case kRecordCache: + As().ReportResultsTo(*callback); + break; } exit: @@ -5010,6 +5089,9 @@ void Core::CacheEntry::HandleTimer(CacheContext &aContext) // compress offset since the host name would not be used // in any other query question. break; + + case kRecordCache: + break; } VerifyOrExit(HasFireTime()); @@ -5042,6 +5124,9 @@ void Core::CacheEntry::HandleTimer(CacheContext &aContext) case kIp4AddrCache: As().ProcessExpiredRecords(aContext.GetNow()); break; + case kRecordCache: + As().ProcessExpiredRecords(aContext.GetNow()); + break; } DetermineNextFireTime(); @@ -5069,6 +5154,9 @@ Core::ResultCallback *Core::CacheEntry::FindCallbackMatching(const ResultCallbac case kIp4AddrCache: callback = mCallbacks.FindMatching(aCallback.mSharedCallback.mAddress); break; + case kRecordCache: + callback = mCallbacks.FindMatching(aCallback.mSharedCallback.mRecord); + break; } return callback; @@ -5107,6 +5195,9 @@ void Core::CacheEntry::DetermineNextFireTime(void) case kIp4AddrCache: As().DetermineRecordFireTime(); break; + case kRecordCache: + As().DetermineRecordFireTime(); + break; } } @@ -5139,6 +5230,9 @@ void Core::CacheEntry::PrepareQuery(CacheContext &aContext) case kIp4AddrCache: As().PrepareAQuestion(query); break; + case kRecordCache: + As().PrepareQueryQuestion(query); + break; } query.CheckSizeLimitToPrepareAgain(prepareAgain); @@ -5169,6 +5263,8 @@ void Core::CacheEntry::PrepareQuery(CacheContext &aContext) case kIp4AddrCache: As().UpdateRecordStateAfterQuery(aContext.GetNow()); break; + case kRecordCache: + As().UpdateRecordStateAfterQuery(aContext.GetNow()); } } @@ -6404,6 +6500,311 @@ exit: void Core::Ip4AddrCache::PrepareAQuestion(TxMessage &aQuery) { PrepareQueryQuestion(aQuery, ResourceRecord::kTypeA); } +//--------------------------------------------------------------------------------------------------------------------- +// Core::RecordCache + +Error Core::RecordCache::Init(Instance &aInstance, const RecordQuerier &aQuerier) +{ + Error error; + + CacheEntry::Init(aInstance, kRecordCache); + + mNext = nullptr; + mShouldFlush = false; + SuccessOrExit(error = mFirstLabel.Set(aQuerier.mFirstLabel)); + SuccessOrExit(error = mNextLabels.Set(aQuerier.mNextLabels)); + mRecordType = aQuerier.mRecordType; + +exit: + return error; +} + +bool Core::RecordCache::Matches(const Name &aFullName, uint16_t aRecordType) const +{ + return (mRecordType == aRecordType) && + aFullName.Matches(mFirstLabel.AsCString(), mNextLabels.AsCString(), kLocalDomain); +} + +bool Core::RecordCache::Matches(const RecordQuerier &aQuerier) const +{ + bool matches = false; + + VerifyOrExit(mRecordType == aQuerier.mRecordType); + + VerifyOrExit(NameMatch(mFirstLabel, aQuerier.mFirstLabel)); + + if (mNextLabels.IsNull()) + { + VerifyOrExit(aQuerier.mNextLabels == nullptr); + } + else + { + VerifyOrExit(NameMatch(mNextLabels, aQuerier.mNextLabels)); + } + + matches = true; + +exit: + return matches; +} + +bool Core::RecordCache::Matches(const ExpireChecker &aExpireChecker) const { return ShouldDelete(aExpireChecker.mNow); } + +Error Core::RecordCache::Add(const RecordQuerier &aQuerier) +{ + return CacheEntry::Add(ResultCallback(aQuerier.mCallback)); +} + +void Core::RecordCache::Remove(const RecordQuerier &aQuerier) +{ + return CacheEntry::Remove(ResultCallback(aQuerier.mCallback)); +} + +void Core::RecordCache::PrepareQueryQuestion(TxMessage &aQuery) +{ + Question question; + + question.SetType(mRecordType); + question.SetClass(ResourceRecord::kClassInternet); + + AppendNameTo(aQuery, kQuestionSection); + SuccessOrAssert(aQuery.SelectMessageFor(kQuestionSection).Append(question)); + + aQuery.IncrementRecordCount(kQuestionSection); +} + +void Core::RecordCache::AppendNameTo(TxMessage &aTxMessage, Section aSection) +{ + uint16_t compressOffset = kUnspecifiedOffset; + AppendOutcome outcome; + + outcome = aTxMessage.AppendLabel(aSection, mFirstLabel.AsCString(), compressOffset); + VerifyOrExit(outcome != kAppendedFullNameAsCompressed); + + if (!mNextLabels.IsNull()) + { + compressOffset = kUnspecifiedOffset; + outcome = aTxMessage.AppendMultipleLabels(aSection, mNextLabels.AsCString(), compressOffset); + VerifyOrExit(outcome != kAppendedFullNameAsCompressed); + } + + aTxMessage.AppendDomainName(aSection); + +exit: + return; +} + +void Core::RecordCache::UpdateRecordStateAfterQuery(TimeMilli aNow) +{ + for (RecordDataEntry &entry : mCommittedEntries) + { + entry.mRecord.UpdateStateAfterQuery(aNow); + } +} + +void Core::RecordCache::ProcessResponseRecord(const Message &aMessage, + const ResourceRecord &aRecord, + uint16_t aRecordOffset) +{ + // Name and record type in `aMessage` are already matched. + + // Adds a new record data to `mNewEntries` list. This called as + // the records in a received response are processed one by one. + // Once all records are processed `CommitNewResponseEntries()` is + // called to update the list. + + Heap::Data data; + RecordDataEntry *entry; + + SuccessOrExit(data.SetFrom(aMessage, aRecordOffset + sizeof(ResourceRecord), aRecord.GetLength())); + + if (aRecord.GetClass() & kClassCacheFlushFlag) + { + mShouldFlush = true; + } + + // Check for duplicates in the same response. + + entry = mNewEntries.FindMatching(data); + + if (entry == nullptr) + { + entry = RecordDataEntry::Allocate(data); + OT_ASSERT(entry != nullptr); + mNewEntries.Push(*entry); + } + + entry->mRecord.RefreshTtl(aRecord.GetTtl()); + +exit: + return; +} + +void Core::RecordCache::CommitNewResponseEntries(void) +{ + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + // Invoke callbacks if there is any change. + + // If we need to flush, check any entry in the previously + // `mCommittedEntries` that does not appear in the new list + // and signal their removal. + + if (mShouldFlush) + { + for (RecordDataEntry &exitingEntry : mCommittedEntries) + { + if (!mNewEntries.ContainsMatching(exitingEntry.mData)) + { + exitingEntry.mRecord.RefreshTtl(0); + PrepareResultAndInvokeCallbacks(exitingEntry); + } + } + } + + // Signal addition of any new entries or if there is any + // change to an existing entry (TTL value changed). + + for (const RecordDataEntry &newEntry : mNewEntries) + { + RecordDataEntry *exitingEntry = mCommittedEntries.FindMatching(newEntry.mData); + bool shouldSignal = false; + + if (exitingEntry == nullptr) + { + shouldSignal = (newEntry.GetTtl() > 0); + } + else + { + shouldSignal = (exitingEntry->GetTtl() != newEntry.GetTtl()); + } + + if (shouldSignal) + { + PrepareResultAndInvokeCallbacks(newEntry); + } + } + + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + // Now merge the new entries into the `mCommittedEntries` list. + + if (mShouldFlush) + { + mCommittedEntries.Clear(); + StopInitialQueries(); + mShouldFlush = false; + } + + while (!mNewEntries.IsEmpty()) + { + OwnedPtr newEntry = mNewEntries.Pop(); + RecordDataEntry *entry; + + entry = mCommittedEntries.FindMatching(newEntry->mData); + + if (entry != nullptr) + { + entry->mRecord.RefreshTtl(newEntry->GetTtl()); + } + else + { + mCommittedEntries.Push(*newEntry.Release()); + } + } + + mCommittedEntries.RemoveAndFreeAllMatching(EmptyChecker()); + + DetermineNextFireTime(); + ScheduleTimer(); +} + +void Core::RecordCache::DetermineRecordFireTime(void) +{ + for (RecordDataEntry &entry : mCommittedEntries) + { + entry.mRecord.UpdateQueryAndFireTimeOn(*this); + } +} + +void Core::RecordCache::ProcessExpiredRecords(TimeMilli aNow) +{ + OwningList expiredEntries; + + mCommittedEntries.RemoveAllMatching(expiredEntries, ExpireChecker(aNow)); + + for (RecordDataEntry &entry : expiredEntries) + { + entry.mRecord.RefreshTtl(0); + PrepareResultAndInvokeCallbacks(entry); + } +} + +void Core::RecordCache::ReportResultsTo(ResultCallback &aCallback) const +{ + for (const RecordDataEntry &entry : mCommittedEntries) + { + RecordResult result; + + PreareResultFor(entry, result); + aCallback.Invoke(GetInstance(), result); + } +} + +void Core::RecordCache::PreareResultFor(const RecordDataEntry &aEntry, RecordResult &aResult) const +{ + ClearAllBytes(aResult); + aResult.mFirstLabel = mFirstLabel.AsCString(); + aResult.mNextLabels = mNextLabels.AsCString(); + aResult.mRecordType = mRecordType; + aResult.mRecordData = aEntry.mData.GetBytes(); + aResult.mRecordDataLength = aEntry.mData.GetLength(); + aResult.mTtl = aEntry.mRecord.GetTtl(); + aResult.mInfraIfIndex = Get().mInfraIfIndex; +} + +void Core::RecordCache::PrepareResultAndInvokeCallbacks(const RecordDataEntry &aEntry) +{ + RecordResult result; + + PreareResultFor(aEntry, result); + InvokeCallbacks(result); +} + +#if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE + +void Core::RecordCache::CopyInfoTo(RecordQuerier &aQuerier, CacheInfo &aInfo) const +{ + aQuerier.mFirstLabel = mFirstLabel.AsCString(); + aQuerier.mNextLabels = mNextLabels.AsCString(); + aQuerier.mRecordType = mRecordType; + aQuerier.mInfraIfIndex = Get().mInfraIfIndex; + aQuerier.mCallback = nullptr; + aInfo.mIsActive = IsActive(); + aInfo.mHasCachedResults = !mCommittedEntries.IsEmpty(); +} + +#endif + +//--------------------------------------------------------------------------------------------------------------------- +// Core::RecordCache::RecordDataEntry + +Core::RecordCache::RecordDataEntry::RecordDataEntry(Heap::Data &aData) + : mNext(nullptr) + , mData(static_cast(aData)) +{ +} + +bool Core::RecordCache::RecordDataEntry::Matches(const ExpireChecker &aExpireChecker) const +{ + return mRecord.ShouldExpire(aExpireChecker.mNow); +} + +bool Core::RecordCache::RecordDataEntry::Matches(EmptyChecker aChecker) const +{ + OT_UNUSED_VARIABLE(aChecker); + + return !mRecord.IsPresent(); +} + //--------------------------------------------------------------------------------------------------------------------- // Core::Iterator @@ -6618,6 +7019,29 @@ exit: return error; } +Error Core::EntryIterator::GetNextRecordQuerier(RecordQuerier &aQuerier, CacheInfo &aInfo) +{ + Error error = kErrorNone; + + if (mType == kUnspecified) + { + mRecordCache = Get().mRecordCacheList.GetHead(); + mType = kRecordQuerier; + } + else + { + VerifyOrExit(mType == kRecordQuerier, error = kErrorInvalidArgs); + } + + VerifyOrExit(mRecordCache != nullptr, error = kErrorNotFound); + + mRecordCache->CopyInfoTo(aQuerier, aInfo); + mRecordCache = mRecordCache->GetNext(); + +exit: + return error; +} + #endif // OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE } // namespace Multicast diff --git a/src/core/net/mdns.hpp b/src/core/net/mdns.hpp index 549fc0a2a..cfd96e635 100644 --- a/src/core/net/mdns.hpp +++ b/src/core/net/mdns.hpp @@ -115,6 +115,9 @@ public: typedef otMdnsAddressCallback AddressCallback; ///< Address callback typedef otMdnsAddressResult AddressResult; ///< Address result. typedef otMdnsAddressAndTtl AddressAndTtl; ///< Address and TTL. + typedef otMdnsRecordResult RecordResult; ///< Record query result + typedef otMdnsRecordCallback RecordCallback; ///< Record query callback. + typedef otMdnsRecordQuerier RecordQuerier; ///< Record querier. typedef otMdnsIterator Iterator; ///< An entry iterator. typedef otMdnsCacheInfo CacheInfo; ///< Cache information. @@ -537,6 +540,50 @@ public: */ Error StopIp4AddressResolver(const AddressResolver &aResolver); + /** + * Starts a record querier. + * + * Initiates a continuous query for a given `mRecordType` as specified in @p aQuerier. The queried name is specified + * by the combination of `mFirstLabel` and `mNextLabels` (optional rest of the labels) in @p aQuerier. The + * `mFirstLabel` MUST be non-NULL but `mNextLabels` can be `NULL` if there are no other labels. The `mNextLabels` + * MUST NOT include the domain name. The reason for a separate first label is to allow it to include a dot `.` + * character (as allowed for service instance labels). + * + * Discovered results are reported through the `mCallback` function in @p aQuerier, providing the raw record + * data bytes. A removed record data is indicated with a TTL value of zero. The callback may be invoked immediately + * with cached information (if available) and potentially before this function returns. When cached results are + * used, the reported TTL value will reflect the original TTL from the last received response. + * + * Multiple querier instances can be started for the same name, provided they use different callback functions. + * + * The record querier MUST not be used for record types PTR, SRV, TXT, A, and AAAA. Otherwise, `kErrorInvalidArgs` + * will be returned. For these, browsers/resolvers can be used. This design is intentional to enable the + * implementation of the "opportunistic cache mechanism", where, depending on currently active service + * browsers/resolvers, the mDNS implementation will also monitor and cache related records (e.g., when a + * service is resolved, the address records associated with its host name are cached even if there is no active + * address resolver for this hostname). + * + * @param[in] aQuerier The record querier to be started. + * + * @retval kErrorNone Record @p aQuerier started successfully. + * @retval kErrorInvalidState mDNS module is not enabled. + * @retval kErrorAlready An identical querier (same name, record type, and callback) is already active. + * @retval kErrorInvalidArg The `mRecordType` in @p aQuerier is invalid. MUST use browser/resolvers. + */ + Error StartRecordQuerier(const RecordQuerier &aQuerier); + + /** + * Stops a record querier. + * + * No action is performed if no matching querier with the same name and callback is currently active. + * + * @param[in] aQuerier The record querier to be stopped. + * + * @retval kErrorNone Querier stopped successfully. + * @retval kErrorInvalidStat mDNS module is not enabled. + */ + Error StopRecordQuerier(const RecordQuerier &aQuerier); + /** * Sets the max size threshold for mDNS messages. * @@ -691,6 +738,24 @@ public: */ Error GetNextIp4AddressResolver(Iterator &aIterator, AddressResolver &aResolver, CacheInfo &aInfo) const; + /** + * Iterates over record querier entries. + * + * On success, @p aQuerier is populated with information about the next querier . The `mCallback` field is always + * set to `nullptr` as there may be multiple active querier with different callbacks. Other pointers within the + * `RecordQuerier` structure remain valid until the next call to any OpenThread stack's public or platform + * API/callback. + * + * @param[in] aIterator The iterator to use + * @param[out] aQuerier A `RecordQuerier` to return the information about the next querier. + * @param[out] aInfo A `CacheInfo` to return additional information. + * + * @retval kErrorNone @p aQuerier, @p aInfo, & @p aIterator are updated successfully. + * @retval kErrorNotFound Reached the end of the list. + * @retval kErrorInvalidArg @p aIterator is not valid. + */ + Error GetNextRecordQuerier(Iterator &aIterator, RecordQuerier &aQuerier, CacheInfo &aInfo) const; + #endif // OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE private: @@ -1354,6 +1419,7 @@ private: void ProcessTxtRecord(const Name &aName, const ResourceRecord &aRecord, uint16_t aRecordOffset); void ProcessAaaaRecord(const Name &aName, const ResourceRecord &aRecord, uint16_t aRecordOffset); void ProcessARecord(const Name &aName, const ResourceRecord &aRecord, uint16_t aRecordOffset); + void ProcessOtherRecord(const Name &aName, const ResourceRecord &aRecord, uint16_t aRecordOffset); RxMessage *mNext; TimeMilli mRxTime; @@ -1480,12 +1546,14 @@ private: bool Matches(SrvCallback aCallback) const { return mSharedCallback.mSrv == aCallback; } bool Matches(TxtCallback aCallback) const { return mSharedCallback.mTxt == aCallback; } bool Matches(AddressCallback aCallback) const { return mSharedCallback.mAddress == aCallback; } + bool Matches(RecordCallback aCallback) const { return mSharedCallback.mRecord == aCallback; } bool Matches(EmptyChecker) const { return (mSharedCallback.mSrv == nullptr); } void Invoke(Instance &aInstance, const BrowseResult &aResult) const; void Invoke(Instance &aInstance, const SrvResult &aResult) const; void Invoke(Instance &aInstance, const TxtResult &aResult) const; void Invoke(Instance &aInstance, const AddressResult &aResult) const; + void Invoke(Instance &aInstance, const RecordResult &aResult) const; void ClearCallback(void) { mSharedCallback.Clear(); } @@ -1496,6 +1564,7 @@ private: explicit SharedCallback(SrvCallback aCallback) { mSrv = aCallback; } explicit SharedCallback(TxtCallback aCallback) { mTxt = aCallback; } explicit SharedCallback(AddressCallback aCallback) { mAddress = aCallback; } + explicit SharedCallback(RecordCallback aCallback) { mRecord = aCallback; } void Clear(void) { mBrowse = nullptr; } @@ -1503,6 +1572,7 @@ private: SrvCallback mSrv; TxtCallback mTxt; AddressCallback mAddress; + RecordCallback mRecord; }; ResultCallback *mNext; @@ -1574,6 +1644,7 @@ private: kTxtCache, kIp6AddrCache, kIp4AddrCache, + kRecordCache, }; void Init(Instance &aInstance, Type aType); @@ -1609,7 +1680,7 @@ private: uint8_t mInitalQueries; // Number initial queries sent already. 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 for this entry. + bool mIsActive : 1; // Whether there is any active resolver/browser/querier for this entry. 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`. @@ -1884,6 +1955,61 @@ private: // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + class RecordCache : public CacheEntry, public LinkedListEntry, public Heap::Allocatable + { + friend class CacheEntry; + friend class LinkedListEntry; + friend class Heap::Allocatable; + + public: + bool Matches(const Name &aFullName, uint16_t aRecordType) const; + bool Matches(const RecordQuerier &aQuerier) const; + bool Matches(const ExpireChecker &aExpireChecker) const; + Error Add(const RecordQuerier &aQuerier); + void Remove(const RecordQuerier &aQuerier); + void ProcessResponseRecord(const Message &aMessage, const ResourceRecord &aRecord, uint16_t aRecordOffset); + void CommitNewResponseEntries(void); +#if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE + void CopyInfoTo(RecordQuerier &aQuerier, CacheInfo &aInfo) const; +#endif + + private: + struct RecordDataEntry : public LinkedListEntry, public Heap::Allocatable + { + explicit RecordDataEntry(Heap::Data &aData); + bool Matches(const Heap::Data &aData) const { return (mData == aData); } + bool Matches(const ExpireChecker &aExpireChecker) const; + bool Matches(EmptyChecker aChecker) const; + uint32_t GetTtl(void) const { return mRecord.GetTtl(); } + + RecordDataEntry *mNext; + Heap::Data mData; + CacheRecordInfo mRecord; + }; + + // Called by base class `CacheEntry` + void PrepareQueryQuestion(TxMessage &aQuery); + void UpdateRecordStateAfterQuery(TimeMilli aNow); + void DetermineRecordFireTime(void); + void ProcessExpiredRecords(TimeMilli aNow); + void ReportResultsTo(ResultCallback &aCallback) const; + + Error Init(Instance &aInstance, const RecordQuerier &aQuerier); + void AppendNameTo(TxMessage &aTxMessage, Section aSection); + void PreareResultFor(const RecordDataEntry &aEntry, RecordResult &aResult) const; + void PrepareResultAndInvokeCallbacks(const RecordDataEntry &aEntry); + + RecordCache *mNext; + Heap::String mFirstLabel; + Heap::String mNextLabels; + uint16_t mRecordType; + OwningList mCommittedEntries; + OwningList mNewEntries; + bool mShouldFlush; + }; + + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + #if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE class EntryIterator : public Iterator, public InstanceLocator, public Heap::Allocatable @@ -1900,6 +2026,7 @@ private: Error GetNextTxtResolver(TxtResolver &aResolver, CacheInfo &aInfo); Error GetNextIp6AddressResolver(AddressResolver &aResolver, CacheInfo &aInfo); Error GetNextIp4AddressResolver(AddressResolver &aResolver, CacheInfo &aInfo); + Error GetNextRecordQuerier(RecordQuerier &aQuerier, CacheInfo &aInfo); private: static constexpr uint16_t kArrayCapacityIncrement = 32; @@ -1916,6 +2043,7 @@ private: kTxtResolver, kIp6AddrResolver, kIp4AddrResolver, + kRecordQuerier, }; explicit EntryIterator(Instance &aInstance); @@ -1931,6 +2059,7 @@ private: const TxtCache *mTxtCache; const Ip6AddrCache *mIp6AddrCache; const Ip4AddrCache *mIp4AddrCache; + const RecordCache *mRecordCache; }; Heap::Array mSubTypeArray; @@ -2005,6 +2134,7 @@ private: OwningList mTxtCacheList; OwningList mIp6AddrCacheList; OwningList mIp4AddrCacheList; + OwningList mRecordCacheList; TimeMilli mNextQueryTxTime; CacheTimer mCacheTimer; CacheTask mCacheTask; @@ -2040,6 +2170,11 @@ template <> inline OwningList &Core::GetCacheList inline OwningList &Core::GetCacheList(void) +{ + return mRecordCacheList; +} + } // namespace Multicast } // namespace Dns diff --git a/tests/unit/test_mdns.cpp b/tests/unit/test_mdns.cpp index eef4bc39c..bb3edf809 100644 --- a/tests/unit/test_mdns.cpp +++ b/tests/unit/test_mdns.cpp @@ -937,6 +937,25 @@ struct DnsMessage : public Allocatable, public LinkedListEntryGet().Allocate(Message::kTypeOther); + VerifyOrQuit(message != nullptr); + + header.Clear(); + header.SetType(Header::kTypeResponse); + header.SetAnswerCount(aNumRecords); + + SuccessOrQuit(message->Append(header)); + + for (uint16_t index = 0; index < aNumRecords; index++) + { + SuccessOrQuit(Name::AppendName(aName, *message)); + + rr.Init(aRecordType); + + if (aCacheFlush) + { + rr.SetClass(rr.GetClass() | kClassCacheFlushFlag); + } + + rr.SetTtl(aRecords[index].mTtl); + rr.SetLength(aRecords[index].mLength); + SuccessOrQuit(message->Append(rr)); + SuccessOrQuit(message->AppendBytes(aRecords[index].mData, aRecords[index].mLength)); + } + + SuccessOrQuit(AsCoreType(&senderAddrInfo.mAddress).FromString(kDeviceIp6Address)); + senderAddrInfo.mPort = kMdnsPort; + senderAddrInfo.mInfraIfIndex = 0; + + Log("Sending record %u response for %s, num-records %u", aRecordType, aName, aNumRecords); + + otPlatMdnsHandleReceive(sInstance, message, /* aIsUnicast */ false, &senderAddrInfo); +} + static void SendResponseWithEmptyKey(const char *aName, Section aSection) { Message *message; @@ -1674,6 +1746,7 @@ Core *InitTest(void) static const uint8_t kKey1[] = {0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77}; static const uint8_t kKey2[] = {0x12, 0x34, 0x56}; +static const uint8_t kKey3[] = {0xaa, 0xbb, 0xcc, 0xdd}; static const uint8_t kTxtData1[] = {3, 'a', '=', '1', 0}; static const uint8_t kTxtData2[] = {1, 'b', 0}; static const uint8_t kEmptyTxtData[] = {0}; @@ -4848,10 +4921,29 @@ struct AddrCallback : public Allocatable, public LinkedListEntry, public LinkedListEntry +{ + static constexpr uint16_t kMaxRecordDataLength = 256; + + template bool MatchesData(const uint8_t (&aData)[kSize]) const + { + return (mRecordDataLength == kSize) && (memcmp(mRecordData, aData, kSize) == 0); + } + + RecordCallback *mNext; + DnsName mFirstLabel; + DnsName mNextLabels; + uint16_t mRecordType; + uint8_t mRecordData[kMaxRecordDataLength]; + uint16_t mRecordDataLength; + uint32_t mTtl; +}; + OwningList sBrowseCallbacks; OwningList sSrvCallbacks; OwningList sTxtCallbacks; OwningList sAddrCallbacks; +OwningList sRecordCallbacks; void HandleBrowseResult(otInstance *aInstance, const otMdnsBrowseResult *aResult) { @@ -5007,6 +5099,41 @@ void HandleAddrResultAlternate(otInstance *aInstance, const otMdnsAddressResult HandleAddrResult(aInstance, aResult); } +void HandleRecordResult(otInstance *aInstance, const otMdnsRecordResult *aResult) +{ + RecordCallback *entry; + + VerifyOrQuit(aInstance == sInstance); + VerifyOrQuit(aResult != nullptr); + VerifyOrQuit(aResult->mFirstLabel != nullptr); + VerifyOrQuit(aResult->mRecordData != nullptr); + VerifyOrQuit(aResult->mInfraIfIndex == kInfraIfIndex); + + VerifyOrQuit(aResult->mRecordDataLength <= RecordCallback::kMaxRecordDataLength); + + Log("Record callback: %s %s type:%u -> rlen:%u ttl:%lu", aResult->mFirstLabel, + (aResult->mNextLabels != nullptr) ? aResult->mNextLabels : "(null)", aResult->mRecordType, + aResult->mRecordDataLength, ToUlong(aResult->mTtl)); + + entry = RecordCallback::Allocate(); + VerifyOrQuit(entry != nullptr); + + entry->mFirstLabel.CopyFrom(aResult->mFirstLabel); + entry->mNextLabels.CopyFrom(aResult->mNextLabels); + entry->mRecordType = aResult->mRecordType; + entry->mRecordDataLength = aResult->mRecordDataLength; + memcpy(entry->mRecordData, aResult->mRecordData, aResult->mRecordDataLength); + entry->mTtl = aResult->mTtl; + + sRecordCallbacks.PushAfterTail(*entry); +} + +void HandleRecordResultAlternate(otInstance *aInstance, const otMdnsRecordResult *aResult) +{ + Log("Alternate record callback is called"); + HandleRecordResult(aInstance, aResult); +} + //--------------------------------------------------------------------------------------------------------------------- void TestBrowser(void) @@ -6745,6 +6872,462 @@ void TestIp6AddrResolver(void) testFreeInstance(sInstance); } +void TestRecordQuerier(void) +{ + static constexpr uint8_t kMaxResponseRecords = 4; + + Core *mdns = InitTest(); + Core::RecordQuerier querier; + Core::RecordQuerier querier2; + Core::Iterator *iterator; + Core::CacheInfo cacheInfo; + const DnsMessage *dnsMsg; + const RecordCallback *recordCallback; + uint16_t heapAllocations; + RecordData records[kMaxResponseRecords]; + + Log("-------------------------------------------------------------------------------------------"); + Log("TestRecordQuerier"); + + AdvanceTime(1); + + heapAllocations = sHeapAllocatedPtrs.GetLength(); + SuccessOrQuit(mdns->SetEnabled(true, kInfraIfIndex)); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Start a record querier. Validate initial queries."); + + ClearAllBytes(querier); + + querier.mFirstLabel = "mysrv"; + querier.mNextLabels = "_srv._udp"; + querier.mRecordType = ResourceRecord::kTypeKey; + querier.mInfraIfIndex = kInfraIfIndex; + querier.mCallback = HandleRecordResult; + + sDnsMessages.Clear(); + SuccessOrQuit(mdns->StartRecordQuerier(querier)); + + for (uint8_t queryCount = 0; queryCount < kNumInitalQueries; queryCount++) + { + sDnsMessages.Clear(); + + AdvanceTime((queryCount == 0) ? 125 : (1U << (queryCount - 1)) * 1000); + + VerifyOrQuit(!sDnsMessages.IsEmpty()); + dnsMsg = sDnsMessages.GetHead(); + dnsMsg->ValidateHeader(kMulticastQuery, /* Q */ 1, /* Ans */ 0, /* Auth */ 0, /* Addnl */ 0); + dnsMsg->ValidateAsQueryFor(querier); + VerifyOrQuit(dnsMsg->GetNext() == nullptr); + } + + sDnsMessages.Clear(); + + AdvanceTime(20 * 1000); + VerifyOrQuit(sDnsMessages.IsEmpty()); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Send a response. Validate callback result."); + + records[0].mData = kKey1; + records[0].mLength = sizeof(kKey1); + records[0].mTtl = 120; + + sRecordCallbacks.Clear(); + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ false, 1, records); + + AdvanceTime(1); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey1)); + VerifyOrQuit(recordCallback->mTtl == 120); + VerifyOrQuit(recordCallback->GetNext() == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Send a second response (without cache-flush). Validate callback result."); + + records[0].mData = kKey2; + records[0].mLength = sizeof(kKey2); + records[0].mTtl = 120; + + sRecordCallbacks.Clear(); + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ false, 1, records); + + AdvanceTime(1); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey2)); + VerifyOrQuit(recordCallback->mTtl == 120); + VerifyOrQuit(recordCallback->GetNext() == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Start another record querier for the same name and record type with different callback."); + + ClearAllBytes(querier2); + + querier2.mFirstLabel = "mysrv"; + querier2.mNextLabels = "_srv._udp"; + querier2.mRecordType = ResourceRecord::kTypeKey; + querier2.mInfraIfIndex = kInfraIfIndex; + querier2.mCallback = HandleRecordResultAlternate; + + sRecordCallbacks.Clear(); + + SuccessOrQuit(mdns->StartRecordQuerier(querier2)); + + AdvanceTime(1); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Validate callback result from cache for the new querier"); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + + for (uint8_t num = 2; num > 0; num--) + { + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey2) || recordCallback->MatchesData(kKey1)); + VerifyOrQuit(recordCallback->mTtl == 120); + recordCallback = recordCallback->GetNext(); + } + + VerifyOrQuit(recordCallback == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Stop the second querier."); + + SuccessOrQuit(mdns->StopRecordQuerier(querier2)); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Send a response (without cache-flush) with one previous record and a new record."); + + records[0].mData = kKey1; + records[0].mLength = sizeof(kKey1); + records[0].mTtl = 120; + + records[1].mData = kKey3; + records[1].mLength = sizeof(kKey3); + records[1].mTtl = 120; + + sRecordCallbacks.Clear(); + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ false, 2, records); + + AdvanceTime(1); + + // Only key3 (which is new) should be reported. + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey3)); + VerifyOrQuit(recordCallback->mTtl == 120); + VerifyOrQuit(recordCallback->GetNext() == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + AdvanceTime(5000); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Send a response (with cache-flush) with only one record, `key3`."); + + records[0].mData = kKey3; + records[0].mLength = sizeof(kKey3); + records[0].mTtl = 120; + + sRecordCallbacks.Clear(); + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ true, 1, records); + + AdvanceTime(1); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Validate callback result indicating the two other two keys are removed."); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + + for (uint8_t num = 2; num > 0; num--) + { + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey1) || recordCallback->MatchesData(kKey2)); + VerifyOrQuit(recordCallback->mTtl == 0); + recordCallback = recordCallback->GetNext(); + } + + VerifyOrQuit(recordCallback == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + AdvanceTime(500); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Send a response removing key3 and other keys."); + + records[0].mData = kKey1; + records[0].mLength = sizeof(kKey1); + records[0].mTtl = 0; + + records[1].mData = kKey2; + records[1].mLength = sizeof(kKey2); + records[1].mTtl = 0; + + records[2].mData = kKey3; + records[2].mLength = sizeof(kKey3); + records[2].mTtl = 0; + + sRecordCallbacks.Clear(); + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ false, 3, records); + + AdvanceTime(1); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Validate callback result indicating key3 is now removed."); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey3)); + VerifyOrQuit(recordCallback->mTtl == 0); + VerifyOrQuit(recordCallback->GetNext() == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Send a response adding two keys"); + + records[0].mData = kKey1; + records[0].mLength = sizeof(kKey1); + records[0].mTtl = 500; + + records[1].mData = kKey2; + records[1].mLength = sizeof(kKey2); + records[1].mTtl = 500; + + sRecordCallbacks.Clear(); + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ true, 2, records); + + AdvanceTime(1); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Validate callback results"); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + + for (uint8_t num = 2; num > 0; num--) + { + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey1) || recordCallback->MatchesData(kKey2)); + VerifyOrQuit(recordCallback->mTtl == 500); + recordCallback = recordCallback->GetNext(); + } + + VerifyOrQuit(recordCallback == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + AdvanceTime(5000); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Send a response changing the TTL for key1"); + + records[0].mData = kKey1; + records[0].mLength = sizeof(kKey1); + records[0].mTtl = 120; + + sRecordCallbacks.Clear(); + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ false, 1, records); + + AdvanceTime(1); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Validate callback results indicating key1 TTL change"); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey1)); + VerifyOrQuit(recordCallback->mTtl == 120); + VerifyOrQuit(recordCallback->GetNext() == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + sRecordCallbacks.Clear(); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Check query is sent at 80 percentage of TTL and then respond to it."); + + // First query should be sent at 80-82% of TTL of 120 second (96.0-98.4 sec). + // We wait for 100 second. Note that 5 seconds already passed in the + // previous step. + + AdvanceTime(96 * 1000 - 1); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + AdvanceTime(4 * 1000 + 1); + + VerifyOrQuit(!sDnsMessages.IsEmpty()); + dnsMsg = sDnsMessages.GetHead(); + dnsMsg->ValidateHeader(kMulticastQuery, /* Q */ 1, /* Ans */ 0, /* Auth */ 0, /* Addnl */ 0); + dnsMsg->ValidateAsQueryFor(querier); + VerifyOrQuit(dnsMsg->GetNext() == nullptr); + + sDnsMessages.Clear(); + VerifyOrQuit(sRecordCallbacks.IsEmpty()); + + AdvanceTime(10); + + SendRecordResponse("mysrv._srv._udp.local.", ResourceRecord::kTypeKey, /* aCacheFlush */ false, 1, records); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Check queries are sent at 80, 85, 90, 95 percentages of TTL."); + + for (uint8_t queryCount = 0; queryCount < kNumRefreshQueries; queryCount++) + { + if (queryCount == 0) + { + // First query is expected in 80-82% of TTL, so + // 80% of 120 = 96.0, 82% of 120 = 98.4 + + AdvanceTime(96 * 1000 - 1); + } + else + { + // Next query should happen within 3%-5% of TTL + // from previous query. We wait 3% of TTL here. + AdvanceTime(3600 - 1); + } + + VerifyOrQuit(sDnsMessages.IsEmpty()); + + // Wait for 2% of TTL of 120 which is 2.4 sec. + + AdvanceTime(2400 + 1); + + VerifyOrQuit(!sDnsMessages.IsEmpty()); + dnsMsg = sDnsMessages.GetHead(); + dnsMsg->ValidateHeader(kMulticastQuery, /* Q */ 1, /* Ans */ 0, /* Auth */ 0, /* Addnl */ 0); + dnsMsg->ValidateAsQueryFor(querier); + VerifyOrQuit(dnsMsg->GetNext() == nullptr); + + sDnsMessages.Clear(); + VerifyOrQuit(sRecordCallbacks.IsEmpty()); + } + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Check TTL timeout and callback result."); + + AdvanceTime(6 * 1000); + + VerifyOrQuit(!sRecordCallbacks.IsEmpty()); + recordCallback = sRecordCallbacks.GetHead(); + VerifyOrQuit(recordCallback->mFirstLabel.Matches("mysrv")); + VerifyOrQuit(recordCallback->mNextLabels.Matches("_srv._udp")); + VerifyOrQuit(recordCallback->mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(recordCallback->MatchesData(kKey1)); + VerifyOrQuit(recordCallback->mTtl == 0); + VerifyOrQuit(recordCallback->GetNext() == nullptr); + + VerifyOrQuit(sDnsMessages.IsEmpty()); + +#if OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Check the list of `RecordQueier` entries and the cache-info"); + + iterator = mdns->AllocateIterator(); + VerifyOrQuit(iterator != nullptr); + + SuccessOrQuit(mdns->GetNextRecordQuerier(*iterator, querier2, cacheInfo)); + VerifyOrQuit(querier2.mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(StringMatch(querier2.mFirstLabel, "mysrv", kStringCaseInsensitiveMatch)); + VerifyOrQuit(StringMatch(querier2.mNextLabels, "_srv._udp", kStringCaseInsensitiveMatch)); + + VerifyOrQuit(cacheInfo.mIsActive); + VerifyOrQuit(cacheInfo.mHasCachedResults); + + VerifyOrQuit(mdns->GetNextRecordQuerier(*iterator, querier2, cacheInfo) == kErrorNotFound); + + mdns->FreeIterator(*iterator); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Stop the record querier"); + + SuccessOrQuit(mdns->StopRecordQuerier(querier)); + + sDnsMessages.Clear(); + + AdvanceTime(10); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Check the list of `RecordQueier` entries and cache-info after stop (no longer active)"); + + iterator = mdns->AllocateIterator(); + VerifyOrQuit(iterator != nullptr); + + SuccessOrQuit(mdns->GetNextRecordQuerier(*iterator, querier2, cacheInfo)); + VerifyOrQuit(querier2.mRecordType == ResourceRecord::kTypeKey); + VerifyOrQuit(StringMatch(querier2.mFirstLabel, "mysrv", kStringCaseInsensitiveMatch)); + VerifyOrQuit(StringMatch(querier2.mNextLabels, "_srv._udp", kStringCaseInsensitiveMatch)); + + VerifyOrQuit(!cacheInfo.mIsActive); + VerifyOrQuit(cacheInfo.mHasCachedResults); + + VerifyOrQuit(mdns->GetNextRecordQuerier(*iterator, querier2, cacheInfo) == kErrorNotFound); + + mdns->FreeIterator(*iterator); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + Log("Check the `RecordQuerier` is correctly removed after 'remove timeout' of 7 minutes"); + + AdvanceTime(7 * 60 * 1000); + VerifyOrQuit(sDnsMessages.IsEmpty()); + + iterator = mdns->AllocateIterator(); + VerifyOrQuit(iterator != nullptr); + + VerifyOrQuit(mdns->GetNextRecordQuerier(*iterator, querier2, cacheInfo) == kErrorNotFound); + + mdns->FreeIterator(*iterator); + +#endif // OPENTHREAD_CONFIG_MULTICAST_DNS_ENTRY_ITERATION_API_ENABLE + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -"); + + SuccessOrQuit(mdns->SetEnabled(false, kInfraIfIndex)); + VerifyOrQuit(sHeapAllocatedPtrs.GetLength() <= heapAllocations); + + Log("End of test"); + + testFreeInstance(sInstance); +} + void TestPassiveCache(void) { static const char *const kSubTypes[] = {"_sub1", "_xyzw"}; @@ -7286,6 +7869,7 @@ int main(void) ot::Dns::Multicast::TestSrvResolver(); ot::Dns::Multicast::TestTxtResolver(); ot::Dns::Multicast::TestIp6AddrResolver(); + ot::Dns::Multicast::TestRecordQuerier(); ot::Dns::Multicast::TestPassiveCache(); ot::Dns::Multicast::TestLegacyUnicastResponse();