diff --git a/src/core/net/dnssd_server.cpp b/src/core/net/dnssd_server.cpp index 5e81902cc..63fced4d7 100644 --- a/src/core/net/dnssd_server.cpp +++ b/src/core/net/dnssd_server.cpp @@ -99,13 +99,9 @@ exit: void Server::Stop(void) { - // Abort all query transactions - for (QueryTransaction &query : mQueryTransactions) + for (ProxyQuery &query : mProxyQueries) { - if (query.IsValid()) - { - query.Finalize(kErrorFailed); - } + Finalize(query, Header::kResponseServerFailure); } #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE @@ -161,25 +157,21 @@ exit: void Server::ProcessQuery(Request &aRequest) { - Error error = kErrorNone; - Response response; - bool shouldSendResponse = true; - ResponseCode rcode = Header::kResponseSuccess; + ResponseCode rcode = Header::kResponseSuccess; + Response response(GetInstance()); #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE if (mEnableUpstreamQuery && ShouldForwardToUpstream(aRequest)) { - error = ResolveByUpstream(aRequest); + Error error = ResolveByUpstream(aRequest); if (error == kErrorNone) { - shouldSendResponse = false; ExitNow(); } LogWarn("Error forwarding to upstream: %s", ErrorToString(error)); - error = kErrorNone; rcode = Header::kResponseServerFailure; // Continue to allocate and prepare the response message @@ -187,22 +179,7 @@ void Server::ProcessQuery(Request &aRequest) } #endif - response.mMessage = Get().mSocket.NewMessage(); - VerifyOrExit(response.mMessage != nullptr, error = kErrorNoBufs); - - // Prepare DNS response header - response.mHeader.SetType(Header::kTypeResponse); - response.mHeader.SetMessageId(aRequest.mHeader.GetMessageId()); - response.mHeader.SetQueryType(aRequest.mHeader.GetQueryType()); - - if (aRequest.mHeader.IsRecursionDesiredFlagSet()) - { - response.mHeader.SetRecursionDesiredFlag(); - } - - // Append the empty header to reserve room for it in the message. - // Header will be updated in the message before sending it. - SuccessOrExit(error = response.mMessage->Append(response.mHeader)); + SuccessOrExit(response.AllocateAndInitFrom(aRequest)); #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE // Forwarding the query to the upstream may have already set the @@ -234,34 +211,60 @@ void Server::ProcessQuery(Request &aRequest) } #endif - if (ResolveByQueryCallbacks(response, *aRequest.mMessageInfo) == kErrorNone) - { - // `ResolveByQueryCallbacks()` will take ownership of the - // allocated `response.mMessage` on success. Therefore, - // there is no need to free it at `exit`. - - shouldSendResponse = false; - } + ResolveByProxy(response, *aRequest.mMessageInfo); exit: - if ((error == kErrorNone) && shouldSendResponse) + if (rcode != Header::kResponseSuccess) { - if (rcode != Header::kResponseSuccess) - { - response.SetResponseCode(rcode); - } - - response.Send(*aRequest.mMessageInfo); + response.SetResponseCode(rcode); } - FreeMessageOnError(response.mMessage, error); + response.Send(*aRequest.mMessageInfo); +} + +Server::Response::Response(Instance &aInstance) + : InstanceLocator(aInstance) +{ + // `mHeader` constructors already clears it + + mOffsets.Clear(); +} + +Error Server::Response::AllocateAndInitFrom(const Request &aRequest) +{ + Error error = kErrorNone; + + mMessage.Reset(Get().mSocket.NewMessage()); + VerifyOrExit(!mMessage.IsNull(), error = kErrorNoBufs); + + mHeader.SetType(Header::kTypeResponse); + mHeader.SetMessageId(aRequest.mHeader.GetMessageId()); + mHeader.SetQueryType(aRequest.mHeader.GetQueryType()); + + if (aRequest.mHeader.IsRecursionDesiredFlagSet()) + { + mHeader.SetRecursionDesiredFlag(); + } + + // Append the empty header to reserve room for it in the message. + // Header will be updated in the message before sending it. + error = mMessage->Append(mHeader); + +exit: + if (error != kErrorNone) + { + mMessage.Free(); + } + + return error; } void Server::Response::Send(const Ip6::MessageInfo &aMessageInfo) { - Error error; ResponseCode rcode = mHeader.GetResponseCode(); + VerifyOrExit(!mMessage.IsNull()); + if (rcode == Header::kResponseServerFailure) { mHeader.SetQuestionCount(0); @@ -272,19 +275,19 @@ void Server::Response::Send(const Ip6::MessageInfo &aMessageInfo) mMessage->Write(0, mHeader); - error = Get().mSocket.SendTo(*mMessage, aMessageInfo); + SuccessOrExit(Get().mSocket.SendTo(*mMessage, aMessageInfo)); - if (error != kErrorNone) - { - mMessage->Free(); - LogWarn("Failed to send reply: %s", ErrorToString(error)); - } - else - { - LogInfo("Send response, rcode:%u", rcode); - } + // When `SendTo()` returns success it takes over ownership of + // the given message, so we release ownership of `mMessage`. + + mMessage.Release(); + + LogInfo("Send response, rcode:%u", rcode); Get().UpdateResponseCounters(rcode); + +exit: + return; } Server::ResponseCode Server::Request::ParseQuestions(uint8_t aTestMode) @@ -419,19 +422,19 @@ Error Server::Response::ParseQueryName(void) switch (mType) { case kPtrQuery: - // `mServiceOffset` may be updated as we read labels and if we + // `mOffsets.mServiceName` may be updated as we read labels and if we // determine that the query name is a sub-type service. - mServiceOffset = sizeof(Header); + mOffsets.mServiceName = sizeof(Header); break; case kSrvQuery: case kTxtQuery: case kSrvTxtQuery: - mInstanceOffset = sizeof(Header); + mOffsets.mInstanceName = sizeof(Header); break; case kAaaaQuery: - mHostOffset = sizeof(Header); + mOffsets.mHostName = sizeof(Header); break; } @@ -451,14 +454,14 @@ Error Server::Response::ParseQueryName(void) if ((mType == kPtrQuery) && StringMatch(label, kSubLabel, kStringCaseInsensitiveMatch)) { - mServiceOffset = offset; + mOffsets.mServiceName = offset; } comapreOffset = offset; if (Name::CompareName(*mMessage, comapreOffset, kDefaultDomainName) == kErrorNone) { - mDomainOffset = offset; + mOffsets.mDomainName = offset; ExitNow(); } } @@ -469,24 +472,11 @@ exit: return error; } -void Server::Response::ReadQueryName(DnsName &aName) const -{ - // Query name is always present immediately after `Header` in the - // question section +void Server::Response::ReadQueryName(DnsName &aName) const { Server::ReadQueryName(*mMessage, aName); } - uint16_t offset = sizeof(Header); +bool Server::Response::QueryNameMatches(const char *aName) const { return Server::QueryNameMatches(*mMessage, aName); } - IgnoreError(Name::ReadName(*mMessage, offset, aName, sizeof(aName))); -} - -bool Server::Response::QueryNameMatches(const char *aName) const -{ - uint16_t offset = sizeof(Header); - - return (Name::CompareName(*mMessage, offset, aName) == kErrorNone); -} - -Error Server::Response::AppendQueryName(void) const { return Name::AppendPointerLabel(sizeof(Header), *mMessage); } +Error Server::Response::AppendQueryName(void) { return Name::AppendPointerLabel(sizeof(Header), *mMessage); } Error Server::Response::AppendPtrRecord(const char *aInstanceLabel, uint32_t aTtl) { @@ -502,9 +492,9 @@ Error Server::Response::AppendPtrRecord(const char *aInstanceLabel, uint32_t aTt recordOffset = mMessage->GetLength(); SuccessOrExit(error = mMessage->Append(ptrRecord)); - mInstanceOffset = mMessage->GetLength(); + mOffsets.mInstanceName = mMessage->GetLength(); SuccessOrExit(error = Name::AppendLabel(aInstanceLabel, *mMessage)); - SuccessOrExit(error = Name::AppendPointerLabel(mServiceOffset, *mMessage)); + SuccessOrExit(error = Name::AppendPointerLabel(mOffsets.mServiceName, *mMessage)); UpdateRecordLength(ptrRecord, recordOffset); @@ -549,14 +539,14 @@ Error Server::Response::AppendSrvRecord(const char *aHostName, srvRecord.SetWeight(aWeight); srvRecord.SetPort(aPort); - SuccessOrExit(error = Name::AppendPointerLabel(mInstanceOffset, *mMessage)); + SuccessOrExit(error = Name::AppendPointerLabel(mOffsets.mInstanceName, *mMessage)); recordOffset = mMessage->GetLength(); SuccessOrExit(error = mMessage->Append(srvRecord)); - mHostOffset = mMessage->GetLength(); + mOffsets.mHostName = mMessage->GetLength(); SuccessOrExit(error = Name::AppendMultipleLabels(hostLabels, *mMessage)); - SuccessOrExit(error = Name::AppendPointerLabel(mDomainOffset, *mMessage)); + SuccessOrExit(error = Name::AppendPointerLabel(mOffsets.mDomainName, *mMessage)); UpdateRecordLength(srvRecord, recordOffset); @@ -602,7 +592,7 @@ Error Server::Response::AppendHostAddresses(const Ip6::Address *aAddrs, uint16_t aaaaRecord.SetTtl(aTtl); aaaaRecord.SetAddress(aAddrs[index]); - SuccessOrExit(error = Name::AppendPointerLabel(mHostOffset, *mMessage)); + SuccessOrExit(error = Name::AppendPointerLabel(mOffsets.mHostName, *mMessage)); SuccessOrExit(error = mMessage->Append(aaaaRecord)); IncResourceRecordCount(); @@ -641,7 +631,7 @@ Error Server::Response::AppendTxtRecord(const void *aTxtData, uint16_t aTxtLengt txtRecord.SetTtl(aTtl); txtRecord.SetLength(aTxtLength); - SuccessOrExit(error = Name::AppendPointerLabel(mInstanceOffset, *mMessage)); + SuccessOrExit(error = Name::AppendPointerLabel(mOffsets.mInstanceName, *mMessage)); SuccessOrExit(error = mMessage->Append(txtRecord)); SuccessOrExit(error = mMessage->AppendBytes(aTxtData, aTxtLength)); @@ -651,7 +641,7 @@ exit: return error; } -void Server::Response::UpdateRecordLength(ResourceRecord &aRecord, uint16_t aOffset) const +void Server::Response::UpdateRecordLength(ResourceRecord &aRecord, uint16_t aOffset) { // Calculates RR DATA length and updates and re-writes it in the // response message. This should be called immediately @@ -837,24 +827,6 @@ exit: #endif // OPENTHREAD_CONFIG_SRP_SERVER_ENABLE -Error Server::ResolveByQueryCallbacks(Response &aResponse, const Ip6::MessageInfo &aMessageInfo) -{ - Error error = kErrorNone; - QueryTransaction *query = nullptr; - DnsName name; - - VerifyOrExit(mQuerySubscribe.IsSet(), error = kErrorFailed); - - query = NewQuery(aResponse, aMessageInfo); - VerifyOrExit(query != nullptr, error = kErrorNoBufs); - - query->ReadQueryName(name); - mQuerySubscribe.Invoke(name); - -exit: - return error; -} - #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE bool Server::ShouldForwardToUpstream(const Request &aRequest) { @@ -939,71 +911,107 @@ exit: } #endif // OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE -Server::QueryTransaction *Server::NewQuery(Response &aResponse, const Ip6::MessageInfo &aMessageInfo) +void Server::ResolveByProxy(Response &aResponse, const Ip6::MessageInfo &aMessageInfo) { - QueryTransaction *newQuery = nullptr; + ProxyQuery *query; + ProxyQueryInfo info; + DnsName name; - for (QueryTransaction &query : mQueryTransactions) + VerifyOrExit(mQuerySubscribe.IsSet()); + + // We try to convert `aResponse.mMessage` to a `ProxyQuery` by + // appending `ProxyQueryInfo` to it. + + info.mType = aResponse.mType; + info.mMessageInfo = aMessageInfo; + info.mExpireTime = TimerMilli::GetNow() + kQueryTimeout; + info.mOffsets = aResponse.mOffsets; + + if (aResponse.mMessage->Append(info) != kErrorNone) { - if (!query.IsValid()) - { - newQuery = &query; - break; - } + aResponse.SetResponseCode(Header::kResponseServerFailure); + ExitNow(); } - VerifyOrExit(newQuery != nullptr); + // Take over the ownership of `aResponse.mMessage` and add it as a + // `ProxyQuery` in `mProxyQueries` list. - *static_cast(newQuery) = aResponse; - newQuery->mMessageInfo = aMessageInfo; - newQuery->mExpireTime = TimerMilli::GetNow() + kQueryTimeout; + query = aResponse.mMessage.Release(); - mTimer.FireAtIfEarlier(newQuery->mExpireTime); + query->Write(0, aResponse.mHeader); + mProxyQueries.Enqueue(*query); + + mTimer.FireAtIfEarlier(info.mExpireTime); + + ReadQueryName(*query, name); + mQuerySubscribe.Invoke(name); exit: - return newQuery; + return; } -bool Server::QueryTransaction::CanAnswer(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo) const +void Server::ReadQueryName(const Message &aQuery, DnsName &aName) { - bool canAnswer = false; + uint16_t offset = sizeof(Header); - switch (mType) - { - case kPtrQuery: - canAnswer = QueryNameMatches(aServiceFullName); - break; - - case kSrvQuery: - case kTxtQuery: - case kSrvTxtQuery: - canAnswer = QueryNameMatches(aInstanceInfo.mFullName); - break; - - case kAaaaQuery: - break; - } - - return canAnswer; + IgnoreError(Name::ReadName(aQuery, offset, aName, sizeof(aName))); } -bool Server::QueryTransaction::CanAnswer(const char *aHostFullName) const +bool Server::QueryNameMatches(const Message &aQuery, const char *aName) { - return (mType == kAaaaQuery) && QueryNameMatches(aHostFullName); + uint16_t offset = sizeof(Header); + + return (Name::CompareName(aQuery, offset, aName) == kErrorNone); } -Error Server::QueryTransaction::ExtractServiceInstanceLabel(const char *aInstanceName, DnsLabel &aLabel) +void Server::ProxyQueryInfo::ReadFrom(const ProxyQuery &aQuery) +{ + SuccessOrAssert(aQuery.Read(aQuery.GetLength() - sizeof(ProxyQueryInfo), *this)); +} + +void Server::ProxyQueryInfo::RemoveFrom(ProxyQuery &aQuery) const +{ + SuccessOrAssert(aQuery.SetLength(aQuery.GetLength() - sizeof(ProxyQueryInfo))); +} + +void Server::ProxyQueryInfo::UpdateIn(ProxyQuery &aQuery) const +{ + aQuery.Write(aQuery.GetLength() - sizeof(ProxyQueryInfo), *this); +} + +Error Server::Response::ExtractServiceInstanceLabel(const char *aInstanceName, DnsLabel &aLabel) { uint16_t offset; DnsName serviceName; - offset = mServiceOffset; + offset = mOffsets.mServiceName; IgnoreError(Name::ReadName(*mMessage, offset, serviceName, sizeof(serviceName))); return Name::ExtractLabels(aInstanceName, serviceName, aLabel, sizeof(aLabel)); } -void Server::QueryTransaction::Answer(const ServiceInstanceInfo &aInstanceInfo) +void Server::RemoveQueryAndPrepareResponse(ProxyQuery &aQuery, const ProxyQueryInfo &aInfo, Response &aResponse) +{ + DnsName name; + + mProxyQueries.Dequeue(aQuery); + aInfo.RemoveFrom(aQuery); + + ReadQueryName(aQuery, name); + mQueryUnsubscribe.InvokeIfSet(name); + + aResponse.InitFrom(aQuery, aInfo); +} + +void Server::Response::InitFrom(ProxyQuery &aQuery, const ProxyQueryInfo &aInfo) +{ + mMessage.Reset(&aQuery); + IgnoreError(mMessage->Read(0, mHeader)); + mType = aInfo.mType; + mOffsets = aInfo.mOffsets; +} + +void Server::Response::Answer(const ServiceInstanceInfo &aInstanceInfo, const Ip6::MessageInfo &aMessageInfo) { static const Section kSections[] = {kAnswerSection, kAdditionalDataSection}; @@ -1043,17 +1051,24 @@ void Server::QueryTransaction::Answer(const ServiceInstanceInfo &aInstanceInfo) error = AppendHostAddresses(aInstanceInfo); exit: - Finalize(error); + if (error != kErrorNone) + { + SetResponseCode(Header::kResponseServerFailure); + } + + Send(aMessageInfo); } -void Server::QueryTransaction::Answer(const HostInfo &aHostInfo) +void Server::Response::Answer(const HostInfo &aHostInfo, const Ip6::MessageInfo &aMessageInfo) { - Error error; - mSection = kAnswerSection; - error = AppendHostAddresses(aHostInfo); - Finalize(error); + if (AppendHostAddresses(aHostInfo) != kErrorNone) + { + SetResponseCode(Header::kResponseServerFailure); + } + + Send(aMessageInfo); } void Server::SetQueryCallbacks(SubscribeCallback aSubscribe, UnsubscribeCallback aUnsubscribe, void *aContext) @@ -1070,11 +1085,38 @@ void Server::HandleDiscoveredServiceInstance(const char *aServiceFullName, const OT_ASSERT(StringEndsWith(aInstanceInfo.mFullName, Name::kLabelSeparatorChar)); OT_ASSERT(StringEndsWith(aInstanceInfo.mHostName, Name::kLabelSeparatorChar)); - for (QueryTransaction &query : mQueryTransactions) + // It is safe to remove entries from `mProxyQueries` as we iterate + // over it since it is a `MessageQueue`. + + for (ProxyQuery &query : mProxyQueries) { - if (query.IsValid() && query.CanAnswer(aServiceFullName, aInstanceInfo)) + bool canAnswer = false; + ProxyQueryInfo info; + + info.ReadFrom(query); + + switch (info.mType) { - query.Answer(aInstanceInfo); + case kPtrQuery: + canAnswer = QueryNameMatches(query, aServiceFullName); + break; + + case kSrvQuery: + case kTxtQuery: + case kSrvTxtQuery: + canAnswer = QueryNameMatches(query, aInstanceInfo.mFullName); + break; + + case kAaaaQuery: + break; + } + + if (canAnswer) + { + Response response(GetInstance()); + + RemoveQueryAndPrepareResponse(query, info, response); + response.Answer(aInstanceInfo, info.mMessageInfo); } } } @@ -1083,57 +1125,41 @@ void Server::HandleDiscoveredHost(const char *aHostFullName, const HostInfo &aHo { OT_ASSERT(StringEndsWith(aHostFullName, Name::kLabelSeparatorChar)); - for (QueryTransaction &query : mQueryTransactions) + for (ProxyQuery &query : mProxyQueries) { - if (query.IsValid() && query.CanAnswer(aHostFullName)) + ProxyQueryInfo info; + + info.ReadFrom(query); + + if ((info.mType == kAaaaQuery) && QueryNameMatches(query, aHostFullName)) { - query.Answer(aHostInfo); + Response response(GetInstance()); + + RemoveQueryAndPrepareResponse(query, info, response); + response.Answer(aHostInfo, info.mMessageInfo); } } } const otDnssdQuery *Server::GetNextQuery(const otDnssdQuery *aQuery) const { - const QueryTransaction *cur = &mQueryTransactions[0]; - const QueryTransaction *found = nullptr; - const QueryTransaction *query = static_cast(aQuery); + const ProxyQuery *query = static_cast(aQuery); - if (aQuery != nullptr) - { - cur = query + 1; - } - - for (; cur < GetArrayEnd(mQueryTransactions); cur++) - { - if (cur->IsValid()) - { - found = cur; - break; - } - } - - return static_cast(found); + return (query == nullptr) ? mProxyQueries.GetHead() : query->GetNext(); } Server::DnsQueryType Server::GetQueryTypeAndName(const otDnssdQuery *aQuery, char (&aName)[Name::kMaxNameSize]) { - const QueryTransaction *query = static_cast(aQuery); - DnsQueryType type; + const ProxyQuery *query = static_cast(aQuery); + ProxyQueryInfo info; + DnsQueryType type; - OT_ASSERT(query->IsValid()); + ReadQueryName(*query, aName); + info.ReadFrom(*query); - query->GetQueryTypeAndName(type, aName); + type = kDnsQueryBrowse; - return type; -} - -void Server::Response::GetQueryTypeAndName(DnsQueryType &aType, DnsName &aName) const -{ - ReadQueryName(aName); - - aType = kDnsQueryBrowse; - - switch (mType) + switch (info.mType) { case kPtrQuery: break; @@ -1141,13 +1167,15 @@ void Server::Response::GetQueryTypeAndName(DnsQueryType &aType, DnsName &aName) case kSrvQuery: case kTxtQuery: case kSrvTxtQuery: - aType = kDnsQueryResolve; + type = kDnsQueryResolve; break; case kAaaaQuery: - aType = kDnsQueryResolveHost; + type = kDnsQueryResolveHost; break; } + + return type; } void Server::HandleTimer(void) @@ -1155,20 +1183,19 @@ void Server::HandleTimer(void) TimeMilli now = TimerMilli::GetNow(); TimeMilli nextExpire = now.GetDistantFuture(); - for (QueryTransaction &query : mQueryTransactions) + for (ProxyQuery &query : mProxyQueries) { - if (!query.IsValid()) - { - continue; - } + ProxyQueryInfo info; - if (query.mExpireTime <= now) + info.ReadFrom(query); + + if (info.mExpireTime <= now) { - query.Finalize(kErrorNone); + Finalize(query, Header::kResponseSuccess); } else { - nextExpire = Min(nextExpire, query.mExpireTime); + nextExpire = Min(nextExpire, info.mExpireTime); } } @@ -1197,19 +1224,16 @@ void Server::HandleTimer(void) } } -void Server::QueryTransaction::Finalize(Error aError) +void Server::Finalize(ProxyQuery &aQuery, ResponseCode aResponseCode) { - DnsName name; + Response response(GetInstance()); + ProxyQueryInfo info; - ReadQueryName(name); - Get().mQueryUnsubscribe.InvokeIfSet(name); + info.ReadFrom(aQuery); + RemoveQueryAndPrepareResponse(aQuery, info, response); - mHeader.SetResponseCode((aError == kErrorNone) ? Header::kResponseSuccess : Header::kResponseServerFailure); - Send(mMessageInfo); - - // Set the `mMessage` to null to indicate that - // `QueryTransaction` is unused. - mMessage = nullptr; + response.SetResponseCode(aResponseCode); + response.Send(info.mMessageInfo); } void Server::UpdateResponseCounters(ResponseCode aResponseCode) diff --git a/src/core/net/dnssd_server.hpp b/src/core/net/dnssd_server.hpp index d442d3cbc..de01d5534 100644 --- a/src/core/net/dnssd_server.hpp +++ b/src/core/net/dnssd_server.hpp @@ -39,6 +39,7 @@ #include "common/callback.hpp" #include "common/message.hpp" #include "common/non_copyable.hpp" +#include "common/owned_ptr.hpp" #include "common/timer.hpp" #include "net/dns_types.hpp" #include "net/ip6.hpp" @@ -292,15 +293,16 @@ public: private: static constexpr bool kBindUnspecifiedNetif = OPENTHREAD_CONFIG_DNSSD_SERVER_BIND_UNSPECIFIED_NETIF; - static constexpr uint8_t kProtocolLabelLength = 4; - static constexpr uint8_t kSubTypeLabelLength = 4; - static constexpr uint16_t kMaxConcurrentQueries = 32; + static constexpr uint32_t kQueryTimeout = OPENTHREAD_CONFIG_DNSSD_QUERY_TIMEOUT; static constexpr uint16_t kMaxConcurrentUpstreamQueries = 32; typedef Header::Response ResponseCode; typedef char DnsName[Name::kMaxNameSize]; typedef char DnsLabel[Name::kMaxLabelSize]; + typedef Message ProxyQuery; + typedef MessageQueue ProxyQueryList; + enum QueryType : uint8_t { kPtrQuery, @@ -326,17 +328,28 @@ private: QueryType mType; }; - class Response : public GetProvider, public Clearable + struct ProxyQueryInfo; + + struct NameOffsets : public Clearable + { + uint16_t mDomainName; + uint16_t mServiceName; + uint16_t mInstanceName; + uint16_t mHostName; + }; + + class Response : public InstanceLocator, private NonCopyable { public: - Response(void) { Clear(); } - Instance &GetInstance(void) const { return mMessage->GetInstance(); } + explicit Response(Instance &aInstance); + Error AllocateAndInitFrom(const Request &aRequest); + void InitFrom(ProxyQuery &aQuery, const ProxyQueryInfo &aInfo); void SetResponseCode(ResponseCode aResponseCode) { mHeader.SetResponseCode(aResponseCode); } ResponseCode AddQuestionsFrom(const Request &aRequest); Error ParseQueryName(void); void ReadQueryName(DnsName &aName) const; bool QueryNameMatches(const char *aName) const; - Error AppendQueryName(void) const; + Error AppendQueryName(void); Error AppendPtrRecord(const char *aInstanceLabel, uint32_t aTtl); Error AppendSrvRecord(const ServiceInstanceInfo &aInstanceInfo); Error AppendSrvRecord(const char *aHostName, @@ -349,10 +362,12 @@ private: Error AppendHostAddresses(const HostInfo &aHostInfo); Error AppendHostAddresses(const ServiceInstanceInfo &aInstanceInfo); Error AppendHostAddresses(const Ip6::Address *aAddrs, uint16_t aAddrsLength, uint32_t aTtl); - void UpdateRecordLength(ResourceRecord &aRecord, uint16_t aOffset) const; + void UpdateRecordLength(ResourceRecord &aRecord, uint16_t aOffset); void IncResourceRecordCount(void); void Send(const Ip6::MessageInfo &aMessageInfo); - void GetQueryTypeAndName(DnsQueryType &aType, DnsName &aName) const; + void Answer(const HostInfo &aHostInfo, const Ip6::MessageInfo &aMessageInfo); + void Answer(const ServiceInstanceInfo &aInstanceInfo, const Ip6::MessageInfo &aMessageInfo); + Error ExtractServiceInstanceLabel(const char *aInstanceName, DnsLabel &aLabel); #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE Error ResolveBySrp(void); bool QueryNameMatchesService(const Srp::Server::Service &aService) const; @@ -365,38 +380,37 @@ private: static const char *QueryTypeToString(QueryType aType); #endif - Message *mMessage; - Header mHeader; - QueryType mType; - Section mSection; - uint16_t mDomainOffset; - uint16_t mServiceOffset; - uint16_t mInstanceOffset; - uint16_t mHostOffset; + OwnedPtr mMessage; + Header mHeader; + QueryType mType; + Section mSection; + NameOffsets mOffsets; }; - struct QueryTransaction : public Response + struct ProxyQueryInfo { - bool IsValid(void) const { return mMessage != nullptr; } - Error ExtractServiceInstanceLabel(const char *aInstanceName, DnsLabel &aLabel); - bool CanAnswer(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo) const; - bool CanAnswer(const char *aHostFullName) const; - void Answer(const ServiceInstanceInfo &aInstanceInfo); - void Answer(const HostInfo &aHostInfo); - void Finalize(Error aError); + void ReadFrom(const ProxyQuery &aQuery); + void RemoveFrom(ProxyQuery &aQuery) const; + void UpdateIn(ProxyQuery &aQuery) const; + QueryType mType; Ip6::MessageInfo mMessageInfo; TimeMilli mExpireTime; + NameOffsets mOffsets; }; - static constexpr uint32_t kQueryTimeout = OPENTHREAD_CONFIG_DNSSD_QUERY_TIMEOUT; - bool IsRunning(void) const { return mSocket.IsBound(); } static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); void ProcessQuery(Request &aRequest); static uint8_t GetNameLength(const char *aName); + void ResolveByProxy(Response &aResponse, const Ip6::MessageInfo &aMessageInfo); + void RemoveQueryAndPrepareResponse(ProxyQuery &aQuery, const ProxyQueryInfo &aInfo, Response &aResponse); + void Finalize(ProxyQuery &aQuery, ResponseCode aResponseCode); + static void ReadQueryName(const Message &aQuery, DnsName &aName); + static bool QueryNameMatches(const Message &aQuery, const char *aName); + #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE static bool ShouldForwardToUpstream(const Request &aRequest); UpstreamQueryTransaction *AllocateUpstreamQueryTransaction(const Ip6::MessageInfo &aMessageInfo); @@ -404,9 +418,6 @@ private: Error ResolveByUpstream(const Request &aRequest); #endif - Error ResolveByQueryCallbacks(Response &aResponse, const Ip6::MessageInfo &aMessageInfo); - QueryTransaction *NewQuery(Response &aResponse, const Ip6::MessageInfo &aMessageInfo); - void HandleTimer(void); void ResetTimer(void); @@ -422,7 +433,7 @@ private: Ip6::Udp::Socket mSocket; - QueryTransaction mQueryTransactions[kMaxConcurrentQueries]; + ProxyQueryList mProxyQueries; Callback mQuerySubscribe; Callback mQueryUnsubscribe; diff --git a/tests/unit/test_dns_client.cpp b/tests/unit/test_dns_client.cpp index 456489ae7..e259ba9ae 100644 --- a/tests/unit/test_dns_client.cpp +++ b/tests/unit/test_dns_client.cpp @@ -238,14 +238,16 @@ void FinalizeTest(void) static const char kHostName[] = "elden"; static const char kHostFullName[] = "elden.default.service.arpa."; -static const char kService1Name[] = "_srv._udp"; -static const char kService1FullName[] = "_srv._udp.default.service.arpa."; -static const char kInstance1Label[] = "srv-instance"; +static const char kService1Name[] = "_srv._udp"; +static const char kService1FullName[] = "_srv._udp.default.service.arpa."; +static const char kInstance1Label[] = "srv-instance"; +static const char kInstance1FullName[] = "srv-instance._srv._udp.default.service.arpa."; static const char kService2Name[] = "_game._udp"; static const char kService2FullName[] = "_game._udp.default.service.arpa."; static const char kService2SubTypeFullName[] = "_best._sub._game._udp.default.service.arpa."; static const char kInstance2Label[] = "last-ninja"; +static const char kInstance2FullName[] = "last-ninja._game._udp.default.service.arpa."; void PrepareService1(Srp::Client::Service &aService) { @@ -908,12 +910,225 @@ void TestDnsClient(void) Log("End of TestDnsClient"); } +//---------------------------------------------------------------------------------------------------------------------- + +char sLastSubscribeName[Dns::Name::kMaxNameSize]; +char sLastUnsubscribeName[Dns::Name::kMaxNameSize]; + +void QuerySubscribe(void *aContext, const char *aFullName) +{ + uint16_t length = StringLength(aFullName, Dns::Name::kMaxNameSize); + + Log("QuerySubscribe(%s)", aFullName); + + VerifyOrQuit(aContext == sInstance); + VerifyOrQuit(length < Dns::Name::kMaxNameSize); + strcpy(sLastSubscribeName, aFullName); +} + +void QueryUnsubscribe(void *aContext, const char *aFullName) +{ + uint16_t length = StringLength(aFullName, Dns::Name::kMaxNameSize); + + Log("QueryUnsubscribe(%s)", aFullName); + + VerifyOrQuit(aContext == sInstance); + VerifyOrQuit(length < Dns::Name::kMaxNameSize); + strcpy(sLastUnsubscribeName, aFullName); +} + +void TestDnssdServerProxyCallback(void) +{ + Srp::Server *srpServer; + Srp::Client *srpClient; + Dns::Client *dnsClient; + Dns::ServiceDiscovery::Server *dnsServer; + otDnssdServiceInstanceInfo instanceInfo; + + Log("--------------------------------------------------------------------------------------------"); + Log("TestDnssdServerProxyCallback"); + + InitTest(); + + srpServer = &sInstance->Get(); + srpClient = &sInstance->Get(); + dnsClient = &sInstance->Get(); + dnsServer = &sInstance->Get(); + + //- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + // Start SRP server. + + SuccessOrQuit(srpServer->SetAddressMode(Srp::Server::kAddressModeUnicast)); + VerifyOrQuit(srpServer->GetState() == Srp::Server::kStateDisabled); + + srpServer->SetEnabled(true); + VerifyOrQuit(srpServer->GetState() != Srp::Server::kStateDisabled); + + AdvanceTime(10000); + VerifyOrQuit(srpServer->GetState() == Srp::Server::kStateRunning); + + //- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + // Start SRP client. + + srpClient->EnableAutoStartMode(nullptr, nullptr); + VerifyOrQuit(srpClient->IsAutoStartModeEnabled()); + + AdvanceTime(2000); + VerifyOrQuit(srpClient->IsRunning()); + + //- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + // Set the query subscribe/unsubscribe callbacks on server + + dnsServer->SetQueryCallbacks(QuerySubscribe, QueryUnsubscribe, sInstance); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - "); + + sLastSubscribeName[0] = '\0'; + sLastUnsubscribeName[0] = '\0'; + + sBrowseInfo.Reset(); + Log("Browse(%s)", kService1FullName); + SuccessOrQuit(dnsClient->Browse(kService1FullName, BrowseCallback, sInstance)); + AdvanceTime(10); + + VerifyOrQuit(strcmp(sLastSubscribeName, kService1FullName) == 0); + VerifyOrQuit(strcmp(sLastUnsubscribeName, "") == 0); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 0); + + Log("Invoke subscribe callback"); + + memset(&instanceInfo, 0, sizeof(instanceInfo)); + instanceInfo.mFullName = kInstance1FullName; + instanceInfo.mHostName = kHostFullName; + instanceInfo.mPort = 200; + + dnsServer->HandleDiscoveredServiceInstance(kService1FullName, instanceInfo); + + AdvanceTime(10); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 1); + SuccessOrQuit(sBrowseInfo.mError); + VerifyOrQuit(sBrowseInfo.mNumInstances == 1); + + VerifyOrQuit(strcmp(sLastUnsubscribeName, kService1FullName) == 0); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - "); + + sLastSubscribeName[0] = '\0'; + sLastUnsubscribeName[0] = '\0'; + + sBrowseInfo.Reset(); + Log("Browse(%s)", kService2FullName); + SuccessOrQuit(dnsClient->Browse(kService2FullName, BrowseCallback, sInstance)); + AdvanceTime(10); + + VerifyOrQuit(strcmp(sLastSubscribeName, kService2FullName) == 0); + VerifyOrQuit(strcmp(sLastUnsubscribeName, "") == 0); + + Log("Invoke subscribe callback for wrong name"); + + memset(&instanceInfo, 0, sizeof(instanceInfo)); + instanceInfo.mFullName = kInstance1FullName; + instanceInfo.mHostName = kHostFullName; + instanceInfo.mPort = 200; + + dnsServer->HandleDiscoveredServiceInstance(kService1FullName, instanceInfo); + + AdvanceTime(10); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 0); + + Log("Invoke subscribe callback for correct name"); + + memset(&instanceInfo, 0, sizeof(instanceInfo)); + instanceInfo.mFullName = kInstance2FullName; + instanceInfo.mHostName = kHostFullName; + instanceInfo.mPort = 200; + + dnsServer->HandleDiscoveredServiceInstance(kService2FullName, instanceInfo); + + AdvanceTime(10); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 1); + SuccessOrQuit(sBrowseInfo.mError); + VerifyOrQuit(sBrowseInfo.mNumInstances == 1); + + VerifyOrQuit(strcmp(sLastUnsubscribeName, kService2FullName) == 0); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - "); + + sLastSubscribeName[0] = '\0'; + sLastUnsubscribeName[0] = '\0'; + + sBrowseInfo.Reset(); + Log("Browse(%s)", kService2FullName); + SuccessOrQuit(dnsClient->Browse(kService2FullName, BrowseCallback, sInstance)); + AdvanceTime(10); + + VerifyOrQuit(strcmp(sLastSubscribeName, kService2FullName) == 0); + VerifyOrQuit(strcmp(sLastUnsubscribeName, "") == 0); + + Log("Do not invoke subscribe callback and let query to timeout"); + + // Query timeout is set to 6 seconds + + AdvanceTime(5000); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 0); + + AdvanceTime(2000); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 1); + SuccessOrQuit(sBrowseInfo.mError); + VerifyOrQuit(sBrowseInfo.mNumInstances == 0); + + VerifyOrQuit(strcmp(sLastUnsubscribeName, kService2FullName) == 0); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - "); + + sLastSubscribeName[0] = '\0'; + sLastUnsubscribeName[0] = '\0'; + + sBrowseInfo.Reset(); + Log("Browse(%s)", kService2FullName); + SuccessOrQuit(dnsClient->Browse(kService2FullName, BrowseCallback, sInstance)); + AdvanceTime(10); + + VerifyOrQuit(strcmp(sLastSubscribeName, kService2FullName) == 0); + VerifyOrQuit(strcmp(sLastUnsubscribeName, "") == 0); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 0); + + Log("Do not invoke subscribe callback and stop server"); + + dnsServer->Stop(); + + AdvanceTime(10); + + VerifyOrQuit(sBrowseInfo.mCallbackCount == 1); + VerifyOrQuit(sBrowseInfo.mError != kErrorNone); + + VerifyOrQuit(strcmp(sLastUnsubscribeName, kService2FullName) == 0); + + Log("- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - "); + + //- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + // Finalize OT instance and validate all heap allocations are freed. + + Log("Finalizing OT instance"); + FinalizeTest(); + + Log("End of TestDnssdServerProxyCallback"); +} + #endif // ENABLE_DNS_TEST int main(void) { #if ENABLE_DNS_TEST TestDnsClient(); + TestDnssdServerProxyCallback(); printf("All tests passed\n"); #else printf("DNS_CLIENT or DSNSSD_SERVER feature is not enabled\n");