From 037056b97e2aef88f7cb5cc4dafb42c9e107cbe9 Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Thu, 3 Aug 2023 09:33:31 -0700 Subject: [PATCH] [dnssd-server] simplifications and enhancements (#9334) This commit contains changes and enhancements in the DNS-SD server/resolver class `Dns::ServiceDiscovery::Server`. - It defines `Request` and `Response` structures, which contain all related information for a DNS query request and response. These structures simplify the code by encapsulating all the related information in one place. - The `Response` class provides methods for preparing the response, such as appending records, DNS names to the response, or checking the questions or updating the header. These methods replace the previous `static` methods, which required all the information to be passed as input parameters. - The `QueryTransacation` type is simplified by declaring it as a subclass of `Response`. It inherits all the helper methods from `Response`. Other previously defined `static` methods are now defined as methods of `QueryTransacation` (such as `CanAnswer()`). - `ResolveBySrp()` and its related methods now directly populate the DNS response code in the `Response` DNS header instead of returning it. - `ResolveBySrp()` is updated such that if a failure is encountered when preparing the answer section, no further processing is performed. This ensures that a previous failure `rcode` is not overwritten. When preparing the additional section, certain `rcode` failures are allowed, such as if the DNS name is not found. - The processing of `Timer` is simplified, using `FiretAtIfEarlier()` and determining next expire time from `HandleTimer()`. - This commit also adds an empty implementation of the `otPlatDns{}` functions (used with `OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE`) under the simulation platform. --- examples/platforms/simulation/CMakeLists.txt | 1 + examples/platforms/simulation/dns.c | 47 ++ src/core/net/dnssd_server.cpp | 838 +++++++++---------- src/core/net/dnssd_server.hpp | 227 ++--- 4 files changed, 515 insertions(+), 598 deletions(-) create mode 100644 examples/platforms/simulation/dns.c diff --git a/examples/platforms/simulation/CMakeLists.txt b/examples/platforms/simulation/CMakeLists.txt index 89841ff74..c99597d4a 100644 --- a/examples/platforms/simulation/CMakeLists.txt +++ b/examples/platforms/simulation/CMakeLists.txt @@ -61,6 +61,7 @@ add_library(openthread-simulation alarm.c crypto.c diag.c + dns.c dso_transport.c entropy.c flash.c diff --git a/examples/platforms/simulation/dns.c b/examples/platforms/simulation/dns.c new file mode 100644 index 000000000..0542af062 --- /dev/null +++ b/examples/platforms/simulation/dns.c @@ -0,0 +1,47 @@ +/* + * Copyright (c) 2023, The OpenThread Authors. + * All rights reserved. + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * 1. Redistributions of source code must retain the above copyright + * notice, this list of conditions and the following disclaimer. + * 2. Redistributions in binary form must reproduce the above copyright + * notice, this list of conditions and the following disclaimer in the + * documentation and/or other materials provided with the distribution. + * 3. Neither the name of the copyright holder nor the + * names of its contributors may be used to endorse or promote products + * derived from this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE + * ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE + * LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR + * CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF + * SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS + * INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN + * CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) + * ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE + * POSSIBILITY OF SUCH DAMAGE. + */ + +#include "platform-simulation.h" + +#include + +#if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE + +void otPlatDnsStartUpstreamQuery(otInstance *aInstance, otPlatDnsUpstreamQuery *aTxn, const otMessage *aQuery) +{ + OT_UNUSED_VARIABLE(aInstance); + OT_UNUSED_VARIABLE(aTxn); + OT_UNUSED_VARIABLE(aQuery); +} + +void otPlatDnsCancelUpstreamQuery(otInstance *aInstance, otPlatDnsUpstreamQuery *aTxn) +{ + otPlatDnsUpstreamQueryDone(aInstance, aTxn, NULL); +} + +#endif diff --git a/src/core/net/dnssd_server.cpp b/src/core/net/dnssd_server.cpp index 35d9a824a..13da75caf 100644 --- a/src/core/net/dnssd_server.cpp +++ b/src/core/net/dnssd_server.cpp @@ -63,9 +63,6 @@ const char *Server::kBlockedDomains[] = {"ipv4only.arpa."}; Server::Server(Instance &aInstance) : InstanceLocator(aInstance) , mSocket(aInstance) - , mQueryCallbackContext(nullptr) - , mQuerySubscribe(nullptr) - , mQueryUnsubscribe(nullptr) #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE , mEnableUpstreamQuery(false) #endif @@ -106,7 +103,7 @@ void Server::Stop(void) { if (query.IsValid()) { - FinalizeQuery(query, Header::kResponseServerFailure); + query.Finalize(Header::kResponseServerFailure); } } @@ -137,7 +134,7 @@ void Server::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessa void Server::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { - Header requestHeader; + Request request; #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE // We first let the `Srp::Server` process the received message. @@ -147,28 +144,29 @@ void Server::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessag VerifyOrExit(Get().HandleDnssdServerUdpReceive(aMessage, aMessageInfo) != kErrorNone); #endif - SuccessOrExit(aMessage.Read(aMessage.GetOffset(), requestHeader)); - VerifyOrExit(requestHeader.GetType() == Header::kTypeQuery); + request.mMessage = &aMessage; + request.mMessageInfo = &aMessageInfo; + SuccessOrExit(aMessage.Read(aMessage.GetOffset(), request.mHeader)); - ProcessQuery(requestHeader, aMessage, aMessageInfo); + VerifyOrExit(request.mHeader.GetType() == Header::kTypeQuery); + + ProcessQuery(request); exit: return; } -void Server::ProcessQuery(const Header &aRequestHeader, Message &aRequestMessage, const Ip6::MessageInfo &aMessageInfo) +void Server::ProcessQuery(const Request &aRequest) { - Error error = kErrorNone; - Message *responseMessage = nullptr; - Header responseHeader; - NameCompressInfo compressInfo(kDefaultDomainName); - Header::Response response = Header::kResponseSuccess; + Error error = kErrorNone; + Response response; bool shouldSendResponse = true; + Header::Response rcode = Header::kResponseSuccess; #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE - if (mEnableUpstreamQuery && ShouldForwardToUpstream(aRequestHeader, aRequestMessage)) + if (mEnableUpstreamQuery && ShouldForwardToUpstream(aRequest)) { - error = ResolveByUpstream(aRequestMessage, aMessageInfo); + error = ResolveByUpstream(aRequest); if (error == kErrorNone) { @@ -178,208 +176,203 @@ void Server::ProcessQuery(const Header &aRequestHeader, Message &aRequestMessage LogWarn("Failed to forward DNS query to upstream: %s", ErrorToString(error)); - error = kErrorNone; - response = Header::kResponseServerFailure; + error = kErrorNone; + rcode = Header::kResponseServerFailure; // Continue to allocate and prepare the response message // to send the `kResponseServerFailure` response code. } #endif - responseMessage = mSocket.NewMessage(); - VerifyOrExit(responseMessage != nullptr, error = kErrorNoBufs); + response.mMessage = mSocket.NewMessage(); + VerifyOrExit(response.mMessage != nullptr, error = kErrorNoBufs); - // Allocate space for DNS header - SuccessOrExit(error = responseMessage->SetLength(sizeof(Header))); + // Prepare DNS response header + response.mHeader.SetType(Header::kTypeResponse); + response.mHeader.SetMessageId(aRequest.mHeader.GetMessageId()); + response.mHeader.SetQueryType(aRequest.mHeader.GetQueryType()); - // Setup initial DNS response header - responseHeader.Clear(); - responseHeader.SetType(Header::kTypeResponse); - responseHeader.SetMessageId(aRequestHeader.GetMessageId()); - responseHeader.SetQueryType(aRequestHeader.GetQueryType()); - if (aRequestHeader.IsRecursionDesiredFlagSet()) + if (aRequest.mHeader.IsRecursionDesiredFlagSet()) { - responseHeader.SetRecursionDesiredFlag(); + 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)); + #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE // Forwarding the query to the upstream may have already set the // response error code. - VerifyOrExit(response == Header::kResponseSuccess); + VerifyOrExit(rcode == Header::kResponseSuccess); #endif // Validate the query - VerifyOrExit(aRequestHeader.GetQueryType() == Header::kQueryTypeStandard, - response = Header::kResponseNotImplemented); - VerifyOrExit(!aRequestHeader.IsTruncationFlagSet(), response = Header::kResponseFormatError); - VerifyOrExit(aRequestHeader.GetQuestionCount() > 0, response = Header::kResponseFormatError); + VerifyOrExit(aRequest.mHeader.GetQueryType() == Header::kQueryTypeStandard, + rcode = Header::kResponseNotImplemented); + VerifyOrExit(!aRequest.mHeader.IsTruncationFlagSet(), rcode = Header::kResponseFormatError); + VerifyOrExit(aRequest.mHeader.GetQuestionCount() > 0, rcode = Header::kResponseFormatError); if (mTestMode & kTestModeSingleQuestionOnly) { - VerifyOrExit(aRequestHeader.GetQuestionCount() == 1, response = Header::kResponseFormatError); + VerifyOrExit(aRequest.mHeader.GetQuestionCount() == 1, rcode = Header::kResponseFormatError); } - response = AddQuestions(aRequestHeader, aRequestMessage, responseHeader, *responseMessage, compressInfo); - VerifyOrExit(response == Header::kResponseSuccess); + SuccessOrExit(response.AddQuestionsFrom(aRequest)); #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE - // Answer the questions - response = ResolveBySrp(responseHeader, *responseMessage, compressInfo); -#endif - if (responseHeader.GetAnswerCount() == 0) + response.ResolveBySrp(); + + if (response.mHeader.GetAnswerCount() != 0) { - if (kErrorNone == ResolveByQueryCallbacks(responseHeader, *responseMessage, compressInfo, aMessageInfo)) - { - shouldSendResponse = false; - } - } -#if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE - else - { - ++mCounters.mResolvedBySrp; + mCounters.mResolvedBySrp++; + ExitNow(); } #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; + } + exit: if ((error == kErrorNone) && shouldSendResponse) { - SendResponse(responseHeader, response, *responseMessage, aMessageInfo, mSocket); + if (rcode != Header::kResponseSuccess) + { + response.mHeader.SetResponseCode(rcode); + } + + response.Send(*aRequest.mMessageInfo); } - FreeMessageOnError(responseMessage, error); + FreeMessageOnError(response.mMessage, error); } -void Server::SendResponse(Header aHeader, - Header::Response aResponseCode, - Message &aMessage, - const Ip6::MessageInfo &aMessageInfo, - Ip6::Udp::Socket &aSocket) +void Server::Response::Send(const Ip6::MessageInfo &aMessageInfo) { - Error error; + Error error; + Header::Response rcode = mHeader.GetResponseCode(); - if (aResponseCode == Header::kResponseServerFailure) + if (rcode == Header::kResponseServerFailure) { LogWarn("failed to handle DNS query due to server failure"); - aHeader.SetQuestionCount(0); - aHeader.SetAnswerCount(0); - aHeader.SetAdditionalRecordCount(0); - IgnoreError(aMessage.SetLength(sizeof(Header))); + mHeader.SetQuestionCount(0); + mHeader.SetAnswerCount(0); + mHeader.SetAdditionalRecordCount(0); + IgnoreError(mMessage->SetLength(sizeof(Header))); } - aHeader.SetResponseCode(aResponseCode); - aMessage.Write(0, aHeader); + mMessage->Write(0, mHeader); - error = aSocket.SendTo(aMessage, aMessageInfo); + error = Get().mSocket.SendTo(*mMessage, aMessageInfo); if (error != kErrorNone) { - // do not use `FreeMessageOnError()` to avoid null check on nonnull pointer - aMessage.Free(); + mMessage->Free(); LogWarn("failed to send DNS-SD reply: %s", ErrorToString(error)); } else { - LogInfo("send DNS-SD reply: %s, RCODE=%d", ErrorToString(error), aResponseCode); + LogInfo("send DNS-SD reply: %s, RCODE=%d", ErrorToString(error), rcode); } - UpdateResponseCounters(aResponseCode); + Get().UpdateResponseCounters(rcode); } -Header::Response Server::AddQuestions(const Header &aRequestHeader, - const Message &aRequestMessage, - Header &aResponseHeader, - Message &aResponseMessage, - NameCompressInfo &aCompressInfo) +Error Server::Response::AddQuestionsFrom(const Request &aRequest) { - Question question; uint16_t readOffset; - Header::Response response = Header::kResponseSuccess; - char name[Name::kMaxNameSize]; + Header::Response rcode = Header::kResponseSuccess; readOffset = sizeof(Header); - // Check and append the questions - for (uint16_t i = 0; i < aRequestHeader.GetQuestionCount(); i++) + for (uint16_t i = 0; i < aRequest.mHeader.GetQuestionCount(); i++) { - NameComponentsOffsetInfo nameComponentsOffsetInfo; - uint16_t qtype; + char name[Name::kMaxNameSize]; + NameComponentsOffsetInfo nameInfo; + Question question; - VerifyOrExit(kErrorNone == Name::ReadName(aRequestMessage, readOffset, name, sizeof(name)), - response = Header::kResponseFormatError); - VerifyOrExit(kErrorNone == aRequestMessage.Read(readOffset, question), response = Header::kResponseFormatError); + VerifyOrExit(Name::ReadName(*aRequest.mMessage, readOffset, name, sizeof(name)) == kErrorNone, + rcode = Header::kResponseFormatError); + VerifyOrExit(aRequest.mMessage->Read(readOffset, question) == kErrorNone, rcode = Header::kResponseFormatError); readOffset += sizeof(question); - qtype = question.GetType(); - - VerifyOrExit(qtype == ResourceRecord::kTypePtr || qtype == ResourceRecord::kTypeSrv || - qtype == ResourceRecord::kTypeTxt || qtype == ResourceRecord::kTypeAaaa, - response = Header::kResponseNotImplemented); - - VerifyOrExit(kErrorNone == FindNameComponents(name, aCompressInfo.GetDomainName(), nameComponentsOffsetInfo), - response = Header::kResponseNameError); - switch (question.GetType()) { case ResourceRecord::kTypePtr: - VerifyOrExit(nameComponentsOffsetInfo.IsServiceName(), response = Header::kResponseNameError); - break; case ResourceRecord::kTypeSrv: - VerifyOrExit(nameComponentsOffsetInfo.IsServiceInstanceName(), response = Header::kResponseNameError); - break; case ResourceRecord::kTypeTxt: - VerifyOrExit(nameComponentsOffsetInfo.IsServiceInstanceName(), response = Header::kResponseNameError); - break; case ResourceRecord::kTypeAaaa: - VerifyOrExit(nameComponentsOffsetInfo.IsHostName(), response = Header::kResponseNameError); break; + default: - ExitNow(response = Header::kResponseNotImplemented); + rcode = Header::kResponseNotImplemented; + ExitNow(); } - VerifyOrExit(AppendQuestion(name, question, aResponseMessage, aCompressInfo) == kErrorNone, - response = Header::kResponseServerFailure); + VerifyOrExit(FindNameComponents(name, kDefaultDomainName, nameInfo) == kErrorNone, + rcode = Header::kResponseNameError); + + switch (question.GetType()) + { + case ResourceRecord::kTypePtr: + VerifyOrExit(nameInfo.IsServiceName(), rcode = Header::kResponseNameError); + break; + case ResourceRecord::kTypeSrv: + VerifyOrExit(nameInfo.IsServiceInstanceName(), rcode = Header::kResponseNameError); + break; + case ResourceRecord::kTypeTxt: + VerifyOrExit(nameInfo.IsServiceInstanceName(), rcode = Header::kResponseNameError); + break; + case ResourceRecord::kTypeAaaa: + VerifyOrExit(nameInfo.IsHostName(), rcode = Header::kResponseNameError); + break; + default: + break; + } + + VerifyOrExit(AppendQuestion(name, question) == kErrorNone, rcode = Header::kResponseServerFailure); } - aResponseHeader.SetQuestionCount(aRequestHeader.GetQuestionCount()); + mHeader.SetQuestionCount(aRequest.mHeader.GetQuestionCount()); exit: - return response; + mHeader.SetResponseCode(rcode); + return (rcode == Header::kResponseSuccess) ? kErrorNone : kErrorFailed; } -Error Server::AppendQuestion(const char *aName, - const Question &aQuestion, - Message &aMessage, - NameCompressInfo &aCompressInfo) +Error Server::Response::AppendQuestion(const char *aName, const Question &aQuestion) { Error error = kErrorNone; switch (aQuestion.GetType()) { case ResourceRecord::kTypePtr: - SuccessOrExit(error = AppendServiceName(aMessage, aName, aCompressInfo)); + SuccessOrExit(error = AppendServiceName(aName)); break; case ResourceRecord::kTypeSrv: case ResourceRecord::kTypeTxt: - SuccessOrExit(error = AppendInstanceName(aMessage, aName, aCompressInfo)); + SuccessOrExit(error = AppendInstanceName(aName)); break; case ResourceRecord::kTypeAaaa: - SuccessOrExit(error = AppendHostName(aMessage, aName, aCompressInfo)); + SuccessOrExit(error = AppendHostName(aName)); break; default: OT_ASSERT(false); } - error = aMessage.Append(aQuestion); + error = mMessage->Append(aQuestion); exit: return error; } -Error Server::AppendPtrRecord(Message &aMessage, - const char *aServiceName, - const char *aInstanceName, - uint32_t aTtl, - NameCompressInfo &aCompressInfo) +Error Server::Response::AppendPtrRecord(const char *aServiceName, const char *aInstanceName, uint32_t aTtl) { Error error; PtrRecord ptrRecord; @@ -388,28 +381,28 @@ Error Server::AppendPtrRecord(Message &aMessage, ptrRecord.Init(); ptrRecord.SetTtl(aTtl); - SuccessOrExit(error = AppendServiceName(aMessage, aServiceName, aCompressInfo)); + SuccessOrExit(error = AppendServiceName(aServiceName)); - recordOffset = aMessage.GetLength(); - SuccessOrExit(error = aMessage.SetLength(recordOffset + sizeof(ptrRecord))); + recordOffset = mMessage->GetLength(); + SuccessOrExit(error = mMessage->SetLength(recordOffset + sizeof(ptrRecord))); - SuccessOrExit(error = AppendInstanceName(aMessage, aInstanceName, aCompressInfo)); + SuccessOrExit(error = AppendInstanceName(aInstanceName)); - ptrRecord.SetLength(aMessage.GetLength() - (recordOffset + sizeof(ResourceRecord))); - aMessage.Write(recordOffset, ptrRecord); + ptrRecord.SetLength(mMessage->GetLength() - (recordOffset + sizeof(ResourceRecord))); + mMessage->Write(recordOffset, ptrRecord); + + IncResourceRecordCount(); exit: return error; } -Error Server::AppendSrvRecord(Message &aMessage, - const char *aInstanceName, - const char *aHostName, - uint32_t aTtl, - uint16_t aPriority, - uint16_t aWeight, - uint16_t aPort, - NameCompressInfo &aCompressInfo) +Error Server::Response::AppendSrvRecord(const char *aInstanceName, + const char *aHostName, + uint32_t aTtl, + uint16_t aPriority, + uint16_t aWeight, + uint16_t aPort) { SrvRecord srvRecord; Error error = kErrorNone; @@ -421,25 +414,23 @@ Error Server::AppendSrvRecord(Message &aMessage, srvRecord.SetWeight(aWeight); srvRecord.SetPort(aPort); - SuccessOrExit(error = AppendInstanceName(aMessage, aInstanceName, aCompressInfo)); + SuccessOrExit(error = AppendInstanceName(aInstanceName)); - recordOffset = aMessage.GetLength(); - SuccessOrExit(error = aMessage.SetLength(recordOffset + sizeof(srvRecord))); + recordOffset = mMessage->GetLength(); + SuccessOrExit(error = mMessage->SetLength(recordOffset + sizeof(srvRecord))); - SuccessOrExit(error = AppendHostName(aMessage, aHostName, aCompressInfo)); + SuccessOrExit(error = AppendHostName(aHostName)); - srvRecord.SetLength(aMessage.GetLength() - (recordOffset + sizeof(ResourceRecord))); - aMessage.Write(recordOffset, srvRecord); + srvRecord.SetLength(mMessage->GetLength() - (recordOffset + sizeof(ResourceRecord))); + mMessage->Write(recordOffset, srvRecord); + + IncResourceRecordCount(); exit: return error; } -Error Server::AppendAaaaRecord(Message &aMessage, - const char *aHostName, - const Ip6::Address &aAddress, - uint32_t aTtl, - NameCompressInfo &aCompressInfo) +Error Server::Response::AppendAaaaRecord(const char *aHostName, const Ip6::Address &aAddress, uint32_t aTtl) { AaaaRecord aaaaRecord; Error error; @@ -448,17 +439,19 @@ Error Server::AppendAaaaRecord(Message &aMessage, aaaaRecord.SetTtl(aTtl); aaaaRecord.SetAddress(aAddress); - SuccessOrExit(error = AppendHostName(aMessage, aHostName, aCompressInfo)); - error = aMessage.Append(aaaaRecord); + SuccessOrExit(error = AppendHostName(aHostName)); + SuccessOrExit(error = mMessage->Append(aaaaRecord)); + + IncResourceRecordCount(); exit: return error; } -Error Server::AppendServiceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo) +Error Server::Response::AppendServiceName(const char *aName) { Error error; - uint16_t serviceCompressOffset = aCompressInfo.GetServiceNameOffset(aMessage, aName); + uint16_t serviceCompressOffset = mCompressInfo.GetServiceNameOffset(*mMessage, aName); const char *serviceName; // Check whether `aName` is a sub-type service name. @@ -468,7 +461,7 @@ Error Server::AppendServiceName(Message &aMessage, const char *aName, NameCompre { uint8_t subTypeLabelLength = static_cast(serviceName - aName) + sizeof(kDnssdSubTypeLabel) - 1; - SuccessOrExit(error = Name::AppendMultipleLabels(aName, subTypeLabelLength, aMessage)); + SuccessOrExit(error = Name::AppendMultipleLabels(aName, subTypeLabelLength, *mMessage)); // Skip over the "._sub." label to get to the root service name. serviceName += sizeof(kDnssdSubTypeLabel) - 1; @@ -480,26 +473,26 @@ Error Server::AppendServiceName(Message &aMessage, const char *aName, NameCompre if (serviceCompressOffset != NameCompressInfo::kUnknownOffset) { - error = Name::AppendPointerLabel(serviceCompressOffset, aMessage); + error = Name::AppendPointerLabel(serviceCompressOffset, *mMessage); } else { - uint8_t domainStart = static_cast(StringLength(serviceName, Name::kMaxNameSize - 1) - - StringLength(aCompressInfo.GetDomainName(), Name::kMaxNameSize - 1)); - uint16_t domainCompressOffset = aCompressInfo.GetDomainNameOffset(); + uint16_t domainStart = StringLength(serviceName, Name::kMaxNameSize - 1) - (sizeof(kDefaultDomainName) - 1); + uint16_t domainCompressOffset = mCompressInfo.GetDomainNameOffset(); - serviceCompressOffset = aMessage.GetLength(); - aCompressInfo.SetServiceNameOffset(serviceCompressOffset); + serviceCompressOffset = mMessage->GetLength(); + mCompressInfo.SetServiceNameOffset(serviceCompressOffset); if (domainCompressOffset == NameCompressInfo::kUnknownOffset) { - aCompressInfo.SetDomainNameOffset(serviceCompressOffset + domainStart); - error = Name::AppendName(serviceName, aMessage); + mCompressInfo.SetDomainNameOffset(serviceCompressOffset + domainStart); + error = Name::AppendName(serviceName, *mMessage); } else { - SuccessOrExit(error = Name::AppendMultipleLabels(serviceName, domainStart, aMessage)); - error = Name::AppendPointerLabel(domainCompressOffset, aMessage); + SuccessOrExit(error = + Name::AppendMultipleLabels(serviceName, static_cast(domainStart), *mMessage)); + error = Name::AppendPointerLabel(domainCompressOffset, *mMessage); } } @@ -507,39 +500,39 @@ exit: return error; } -Error Server::AppendInstanceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo) +Error Server::Response::AppendInstanceName(const char *aName) { Error error; - uint16_t instanceCompressOffset = aCompressInfo.GetInstanceNameOffset(aMessage, aName); + uint16_t instanceCompressOffset = mCompressInfo.GetInstanceNameOffset(*mMessage, aName); if (instanceCompressOffset != NameCompressInfo::kUnknownOffset) { - error = Name::AppendPointerLabel(instanceCompressOffset, aMessage); + error = Name::AppendPointerLabel(instanceCompressOffset, *mMessage); } else { NameComponentsOffsetInfo nameComponentsInfo; - IgnoreError(FindNameComponents(aName, aCompressInfo.GetDomainName(), nameComponentsInfo)); + IgnoreError(FindNameComponents(aName, kDefaultDomainName, nameComponentsInfo)); OT_ASSERT(nameComponentsInfo.IsServiceInstanceName()); - aCompressInfo.SetInstanceNameOffset(aMessage.GetLength()); + mCompressInfo.SetInstanceNameOffset(mMessage->GetLength()); // Append the instance name as one label - SuccessOrExit(error = Name::AppendLabel(aName, nameComponentsInfo.mServiceOffset - 1, aMessage)); + SuccessOrExit(error = Name::AppendLabel(aName, nameComponentsInfo.mServiceOffset - 1, *mMessage)); { const char *serviceName = aName + nameComponentsInfo.mServiceOffset; - uint16_t serviceCompressOffset = aCompressInfo.GetServiceNameOffset(aMessage, serviceName); + uint16_t serviceCompressOffset = mCompressInfo.GetServiceNameOffset(*mMessage, serviceName); if (serviceCompressOffset != NameCompressInfo::kUnknownOffset) { - error = Name::AppendPointerLabel(serviceCompressOffset, aMessage); + error = Name::AppendPointerLabel(serviceCompressOffset, *mMessage); } else { - aCompressInfo.SetServiceNameOffset(aMessage.GetLength()); - error = Name::AppendName(serviceName, aMessage); + mCompressInfo.SetServiceNameOffset(mMessage->GetLength()); + error = Name::AppendName(serviceName, *mMessage); } } } @@ -548,64 +541,64 @@ exit: return error; } -Error Server::AppendTxtRecord(Message &aMessage, - const char *aInstanceName, - const void *aTxtData, - uint16_t aTxtLength, - uint32_t aTtl, - NameCompressInfo &aCompressInfo) +Error Server::Response::AppendTxtRecord(const char *aInstanceName, + const void *aTxtData, + uint16_t aTxtLength, + uint32_t aTtl) { Error error = kErrorNone; TxtRecord txtRecord; const uint8_t kEmptyTxt = 0; - SuccessOrExit(error = AppendInstanceName(aMessage, aInstanceName, aCompressInfo)); + SuccessOrExit(error = AppendInstanceName(aInstanceName)); txtRecord.Init(); txtRecord.SetTtl(aTtl); txtRecord.SetLength(aTxtLength > 0 ? aTxtLength : sizeof(kEmptyTxt)); - SuccessOrExit(error = aMessage.Append(txtRecord)); + SuccessOrExit(error = mMessage->Append(txtRecord)); + if (aTxtLength > 0) { - error = aMessage.AppendBytes(aTxtData, aTxtLength); + SuccessOrExit(error = mMessage->AppendBytes(aTxtData, aTxtLength)); } else { - error = aMessage.Append(kEmptyTxt); + SuccessOrExit(error = mMessage->Append(kEmptyTxt)); } + IncResourceRecordCount(); + exit: return error; } -Error Server::AppendHostName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo) +Error Server::Response::AppendHostName(const char *aName) { Error error; - uint16_t hostCompressOffset = aCompressInfo.GetHostNameOffset(aMessage, aName); + uint16_t hostCompressOffset = mCompressInfo.GetHostNameOffset(*mMessage, aName); if (hostCompressOffset != NameCompressInfo::kUnknownOffset) { - error = Name::AppendPointerLabel(hostCompressOffset, aMessage); + error = Name::AppendPointerLabel(hostCompressOffset, *mMessage); } else { - uint8_t domainStart = static_cast(StringLength(aName, Name::kMaxNameLength) - - StringLength(aCompressInfo.GetDomainName(), Name::kMaxNameSize - 1)); - uint16_t domainCompressOffset = aCompressInfo.GetDomainNameOffset(); + uint16_t domainStart = StringLength(aName, Name::kMaxNameLength) - (sizeof(kDefaultDomainName) - 1); + uint16_t domainCompressOffset = mCompressInfo.GetDomainNameOffset(); - hostCompressOffset = aMessage.GetLength(); - aCompressInfo.SetHostNameOffset(hostCompressOffset); + hostCompressOffset = mMessage->GetLength(); + mCompressInfo.SetHostNameOffset(hostCompressOffset); if (domainCompressOffset == NameCompressInfo::kUnknownOffset) { - aCompressInfo.SetDomainNameOffset(hostCompressOffset + domainStart); - error = Name::AppendName(aName, aMessage); + mCompressInfo.SetDomainNameOffset(hostCompressOffset + domainStart); + error = Name::AppendName(aName, *mMessage); } else { - SuccessOrExit(error = Name::AppendMultipleLabels(aName, domainStart, aMessage)); - error = Name::AppendPointerLabel(domainCompressOffset, aMessage); + SuccessOrExit(error = Name::AppendMultipleLabels(aName, static_cast(domainStart), *mMessage)); + error = Name::AppendPointerLabel(domainCompressOffset, *mMessage); } } @@ -613,15 +606,15 @@ exit: return error; } -void Server::IncResourceRecordCount(Header &aHeader, bool aAdditional) +void Server::Response::IncResourceRecordCount(void) { - if (aAdditional) + if (mAdditional) { - aHeader.SetAdditionalRecordCount(aHeader.GetAdditionalRecordCount() + 1); + mHeader.SetAdditionalRecordCount(mHeader.GetAdditionalRecordCount() + 1); } else { - aHeader.SetAnswerCount(aHeader.GetAnswerCount() + 1); + mHeader.SetAnswerCount(mHeader.GetAnswerCount() + 1); } } @@ -710,41 +703,46 @@ exit: } #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE -Header::Response Server::ResolveBySrp(Header &aResponseHeader, - Message &aResponseMessage, - Server::NameCompressInfo &aCompressInfo) +void Server::Response::ResolveBySrp(void) { - Question question; - uint16_t readOffset = sizeof(Header); - Header::Response response = Header::kResponseSuccess; - char name[Name::kMaxNameSize]; + uint16_t readOffset = sizeof(Header); + char name[Name::kMaxNameSize]; + Question question; - for (uint16_t i = 0; i < aResponseHeader.GetQuestionCount(); i++) + mAdditional = false; + + for (uint16_t i = 0; i < mHeader.GetQuestionCount(); i++) { - IgnoreError(Name::ReadName(aResponseMessage, readOffset, name, sizeof(name))); - IgnoreError(aResponseMessage.Read(readOffset, question)); + // The names and questions in the request message are validated + // from `AddQuestionsFrom()`, so we `IgnoreError()` here. + + IgnoreError(Name::ReadName(*mMessage, readOffset, name, sizeof(name))); + IgnoreError(mMessage->Read(readOffset, question)); readOffset += sizeof(question); - response = ResolveQuestionBySrp(name, question, aResponseHeader, aResponseMessage, aCompressInfo, - /* aAdditional */ false); + ResolveQuestionBySrp(name, question); - LogInfo("ANSWER: TRANSACTION=0x%04x, QUESTION=[%s %d %d], RCODE=%d", aResponseHeader.GetMessageId(), name, - question.GetClass(), question.GetType(), response); + LogInfo("ANSWER: TRANSACTION=0x%04x, QUESTION=[%s %d %d], RCODE=%d", mHeader.GetMessageId(), name, + question.GetClass(), question.GetType(), mHeader.GetResponseCode()); + + VerifyOrExit(mHeader.GetResponseCode() == Header::kResponseSuccess); } // Answer the questions with additional RRs if required - if (aResponseHeader.GetAnswerCount() > 0) + if (mHeader.GetAnswerCount() > 0) { - VerifyOrExit(!(mTestMode & kTestModeEmptyAdditionalSection)); + mAdditional = true; + + VerifyOrExit(!(Get().mTestMode & kTestModeEmptyAdditionalSection)); readOffset = sizeof(Header); - for (uint16_t i = 0; i < aResponseHeader.GetQuestionCount(); i++) + for (uint16_t i = 0; i < mHeader.GetQuestionCount(); i++) { - IgnoreError(Name::ReadName(aResponseMessage, readOffset, name, sizeof(name))); - IgnoreError(aResponseMessage.Read(readOffset, question)); + IgnoreError(Name::ReadName(*mMessage, readOffset, name, sizeof(name))); + IgnoreError(mMessage->Read(readOffset, question)); readOffset += sizeof(question); - if ((question.GetType() == ResourceRecord::kTypePtr) && (aResponseHeader.GetAnswerCount() > 1)) + if ((question.GetType() == ResourceRecord::kTypePtr) && (mHeader.GetAnswerCount() > 1)) { // Skip adding additional records, when answering a // PTR query with more than one answer. This is the @@ -753,30 +751,25 @@ Header::Response Server::ResolveBySrp(Header &aResponseHeader, continue; } - VerifyOrExit(Header::kResponseServerFailure != ResolveQuestionBySrp(name, question, aResponseHeader, - aResponseMessage, aCompressInfo, - /* aAdditional */ true), - response = Header::kResponseServerFailure); + ResolveQuestionBySrp(name, question); - LogInfo("ADDITIONAL: TRANSACTION=0x%04x, QUESTION=[%s %d %d], RCODE=%d", aResponseHeader.GetMessageId(), - name, question.GetClass(), question.GetType(), response); + LogInfo("ADDITIONAL: TRANSACTION=0x%04x, QUESTION=[%s %d %d], RCODE=%d", mHeader.GetMessageId(), name, + question.GetClass(), question.GetType(), mHeader.GetResponseCode()); + + VerifyOrExit(mHeader.GetResponseCode() == Header::kResponseSuccess); } } + exit: - return response; + return; } -Header::Response Server::ResolveQuestionBySrp(const char *aName, - const Question &aQuestion, - Header &aResponseHeader, - Message &aResponseMessage, - NameCompressInfo &aCompressInfo, - bool aAdditional) +void Server::Response::ResolveQuestionBySrp(const char *aName, const Question &aQuestion) { - Error error = kErrorNone; - TimeMilli now = TimerMilli::GetNow(); - uint16_t qtype = aQuestion.GetType(); - Header::Response response = Header::kResponseNameError; + Error error = kErrorNone; + TimeMilli now = TimerMilli::GetNow(); + uint16_t qtype = aQuestion.GetType(); + Header::Response rcode = Header::kResponseNameError; for (const Srp::Server::Host &host : Get().GetHosts()) { @@ -819,41 +812,33 @@ Header::Response Server::ResolveQuestionBySrp(const char *aName, needAdditionalAaaaRecord = true; } - if (!aAdditional && ptrQueryMatched) + if (!mAdditional && ptrQueryMatched) { - SuccessOrExit( - error = AppendPtrRecord(aResponseMessage, aName, instanceName, instanceTtl, aCompressInfo)); - IncResourceRecordCount(aResponseHeader, aAdditional); - response = Header::kResponseSuccess; + SuccessOrExit(error = AppendPtrRecord(aName, instanceName, instanceTtl)); + rcode = Header::kResponseSuccess; } - if ((!aAdditional && srvQueryMatched) || - (aAdditional && ptrQueryMatched && - !HasQuestion(aResponseHeader, aResponseMessage, instanceName, ResourceRecord::kTypeSrv))) + if ((!mAdditional && srvQueryMatched) || + (mAdditional && ptrQueryMatched && !HasQuestion(instanceName, ResourceRecord::kTypeSrv))) { - SuccessOrExit(error = AppendSrvRecord(aResponseMessage, instanceName, hostName, instanceTtl, - service.GetPriority(), service.GetWeight(), service.GetPort(), - aCompressInfo)); - IncResourceRecordCount(aResponseHeader, aAdditional); - response = Header::kResponseSuccess; + SuccessOrExit(error = AppendSrvRecord(instanceName, hostName, instanceTtl, service.GetPriority(), + service.GetWeight(), service.GetPort())); + rcode = Header::kResponseSuccess; } - if ((!aAdditional && txtQueryMatched) || - (aAdditional && ptrQueryMatched && - !HasQuestion(aResponseHeader, aResponseMessage, instanceName, ResourceRecord::kTypeTxt))) + if ((!mAdditional && txtQueryMatched) || + (mAdditional && ptrQueryMatched && !HasQuestion(instanceName, ResourceRecord::kTypeTxt))) { - SuccessOrExit(error = AppendTxtRecord(aResponseMessage, instanceName, service.GetTxtData(), - service.GetTxtDataLength(), instanceTtl, aCompressInfo)); - IncResourceRecordCount(aResponseHeader, aAdditional); - response = Header::kResponseSuccess; + SuccessOrExit(error = AppendTxtRecord(instanceName, service.GetTxtData(), + service.GetTxtDataLength(), instanceTtl)); + rcode = Header::kResponseSuccess; } } } // Handle AAAA query - if ((!aAdditional && qtype == ResourceRecord::kTypeAaaa && host.Matches(aName)) || - (aAdditional && needAdditionalAaaaRecord && - !HasQuestion(aResponseHeader, aResponseMessage, hostName, ResourceRecord::kTypeAaaa))) + if ((!mAdditional && qtype == ResourceRecord::kTypeAaaa && host.Matches(aName)) || + (mAdditional && needAdditionalAaaaRecord && !HasQuestion(hostName, ResourceRecord::kTypeAaaa))) { uint8_t addrNum; const Ip6::Address *addrs = host.GetAddresses(addrNum); @@ -861,69 +846,81 @@ Header::Response Server::ResolveQuestionBySrp(const char *aName, for (uint8_t i = 0; i < addrNum; i++) { - SuccessOrExit(error = AppendAaaaRecord(aResponseMessage, hostName, addrs[i], hostTtl, aCompressInfo)); - IncResourceRecordCount(aResponseHeader, aAdditional); + SuccessOrExit(error = AppendAaaaRecord(hostName, addrs[i], hostTtl)); } - response = Header::kResponseSuccess; + rcode = Header::kResponseSuccess; } } exit: - return error == kErrorNone ? response : Header::kResponseServerFailure; + + // If there is an `error` (for example, appending to the message + // fails), we always set the response code in the header to + // `kResponseServerFailure`. Otherwise, we only set the response + // code if the entry is for the answer section, not the + // additional data section. + + if (error != kErrorNone) + { + mHeader.SetResponseCode(Header::kResponseServerFailure); + } + else if (!mAdditional) + { + mHeader.SetResponseCode(rcode); + } } #endif // OPENTHREAD_CONFIG_SRP_SERVER_ENABLE -Error Server::ResolveByQueryCallbacks(Header &aResponseHeader, - Message &aResponseMessage, - NameCompressInfo &aCompressInfo, - const Ip6::MessageInfo &aMessageInfo) +Error Server::ResolveByQueryCallbacks(Response &aResponse, const Ip6::MessageInfo &aMessageInfo) { + Error error = kErrorNone; QueryTransaction *query = nullptr; DnsQueryType queryType; char name[Name::kMaxNameSize]; - Error error = kErrorNone; + VerifyOrExit(mQuerySubscribe.IsSet(), error = kErrorFailed); - VerifyOrExit(mQuerySubscribe != nullptr, error = kErrorFailed); - - queryType = GetQueryTypeAndName(aResponseHeader, aResponseMessage, name); + aResponse.GetQueryTypeAndName(queryType, name); VerifyOrExit(queryType != kDnsQueryNone, error = kErrorNotImplemented); - query = NewQuery(aResponseHeader, aResponseMessage, aCompressInfo, aMessageInfo); + query = NewQuery(aResponse, aMessageInfo); VerifyOrExit(query != nullptr, error = kErrorNoBufs); - mQuerySubscribe(mQueryCallbackContext, name); + mQuerySubscribe.Invoke(name); exit: return error; } #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE -bool Server::ShouldForwardToUpstream(const Header &aRequestHeader, const Message &aRequestMessage) +bool Server::ShouldForwardToUpstream(const Request &aRequest) { - bool ret = true; + bool shouldForward = false; uint16_t readOffset; char name[Name::kMaxNameSize]; - VerifyOrExit(aRequestHeader.IsRecursionDesiredFlagSet(), ret = false); + VerifyOrExit(aRequest.mHeader.IsRecursionDesiredFlagSet()); readOffset = sizeof(Header); - for (uint16_t i = 0; i < aRequestHeader.GetQuestionCount(); i++) + for (uint16_t i = 0; i < aRequest.mHeader.GetQuestionCount(); i++) { - VerifyOrExit(kErrorNone == Name::ReadName(aRequestMessage, readOffset, name, sizeof(name)), ret = false); + SuccessOrExit(Name::ReadName(*aRequest.mMessage, readOffset, name, sizeof(name))); readOffset += sizeof(Question); - VerifyOrExit(!Name::IsSubDomainOf(name, kDefaultDomainName), ret = false); + VerifyOrExit(!Name::IsSubDomainOf(name, kDefaultDomainName)); + for (const char *blockedDomain : kBlockedDomains) { - VerifyOrExit(!Name::IsSameDomain(name, blockedDomain), ret = false); + VerifyOrExit(!Name::IsSameDomain(name, blockedDomain)); } } + shouldForward = true; + exit: - return ret; + return shouldForward; } void Server::OnUpstreamQueryDone(UpstreamQueryTransaction &aQueryTransaction, Message *aResponseMessage) @@ -936,8 +933,8 @@ void Server::OnUpstreamQueryDone(UpstreamQueryTransaction &aQueryTransaction, Me { error = mSocket.SendTo(*aResponseMessage, aQueryTransaction.GetMessageInfo()); } + ResetUpstreamQueryTransaction(aQueryTransaction, error); - ResetTimer(); exit: FreeMessageOnError(aResponseMessage, error); @@ -945,78 +942,74 @@ exit: Server::UpstreamQueryTransaction *Server::AllocateUpstreamQueryTransaction(const Ip6::MessageInfo &aMessageInfo) { - UpstreamQueryTransaction *ret = nullptr; + UpstreamQueryTransaction *newTxn = nullptr; for (UpstreamQueryTransaction &txn : mUpstreamQueryTransactions) { if (!txn.IsValid()) { - ret = &txn; - txn.Init(aMessageInfo); + newTxn = &txn; break; } } - if (ret != nullptr) - { - LogInfo("Upstream query transaction %d initialized.", static_cast(ret - mUpstreamQueryTransactions)); - mTimer.FireAtIfEarlier(ret->GetExpireTime()); - } + VerifyOrExit(newTxn != nullptr); - return ret; + newTxn->Init(aMessageInfo); + LogInfo("Upstream query transaction %d initialized.", static_cast(newTxn - mUpstreamQueryTransactions)); + mTimer.FireAtIfEarlier(newTxn->GetExpireTime()); + +exit: + return newTxn; } -Error Server::ResolveByUpstream(const Message &aRequestMessage, const Ip6::MessageInfo &aMessageInfo) +Error Server::ResolveByUpstream(const Request &aRequest) { Error error = kErrorNone; - UpstreamQueryTransaction *txn = nullptr; + UpstreamQueryTransaction *txn; - txn = AllocateUpstreamQueryTransaction(aMessageInfo); + txn = AllocateUpstreamQueryTransaction(*aRequest.mMessageInfo); VerifyOrExit(txn != nullptr, error = kErrorNoBufs); - otPlatDnsStartUpstreamQuery(&GetInstance(), txn, &aRequestMessage); + otPlatDnsStartUpstreamQuery(&GetInstance(), txn, aRequest.mMessage); exit: return error; } #endif // OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE -Server::QueryTransaction *Server::NewQuery(const Header &aResponseHeader, - Message &aResponseMessage, - const NameCompressInfo &aCompressInfo, - const Ip6::MessageInfo &aMessageInfo) +Server::QueryTransaction *Server::NewQuery(Response &aResponse, const Ip6::MessageInfo &aMessageInfo) { QueryTransaction *newQuery = nullptr; for (QueryTransaction &query : mQueryTransactions) { - if (query.IsValid()) + if (!query.IsValid()) { - continue; + newQuery = &query; + break; } - - query.Init(aResponseHeader, aResponseMessage, aCompressInfo, aMessageInfo, GetInstance()); - ExitNow(newQuery = &query); } + VerifyOrExit(newQuery != nullptr); + + *static_cast(newQuery) = aResponse; + newQuery->mMessageInfo = aMessageInfo; + newQuery->mExpireTime = TimerMilli::GetNow() + kQueryTimeout; + + mTimer.FireAtIfEarlier(newQuery->mExpireTime); + exit: - if (newQuery != nullptr) - { - ResetTimer(); - } - return newQuery; } -bool Server::CanAnswerQuery(const QueryTransaction &aQuery, - const char *aServiceFullName, - const otDnssdServiceInstanceInfo &aInstanceInfo) +bool Server::QueryTransaction::CanAnswer(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo) const { char name[Name::kMaxNameSize]; DnsQueryType sdType; bool canAnswer = false; - sdType = GetQueryTypeAndName(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), name); + GetQueryTypeAndName(sdType, name); switch (sdType) { @@ -1033,58 +1026,48 @@ bool Server::CanAnswerQuery(const QueryTransaction &aQuery, return canAnswer; } -bool Server::CanAnswerQuery(const Server::QueryTransaction &aQuery, const char *aHostFullName) +bool Server::QueryTransaction::CanAnswer(const char *aHostFullName) const { char name[Name::kMaxNameSize]; DnsQueryType sdType; - sdType = GetQueryTypeAndName(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), name); + GetQueryTypeAndName(sdType, name); + return (sdType == kDnsQueryResolveHost) && StringMatch(name, aHostFullName, kStringCaseInsensitiveMatch); } -void Server::AnswerQuery(QueryTransaction &aQuery, - const char *aServiceFullName, - const otDnssdServiceInstanceInfo &aInstanceInfo) +void Server::QueryTransaction::Answer(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo) { - Header &responseHeader = aQuery.GetResponseHeader(); - Message &responseMessage = aQuery.GetResponseMessage(); - Error error = kErrorNone; - NameCompressInfo &compressInfo = aQuery.GetNameCompressInfo(); + Error error = kErrorNone; - if (HasQuestion(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), aServiceFullName, - ResourceRecord::kTypePtr)) + mAdditional = false; + + if (HasQuestion(aServiceFullName, ResourceRecord::kTypePtr)) { - SuccessOrExit(error = AppendPtrRecord(responseMessage, aServiceFullName, aInstanceInfo.mFullName, - aInstanceInfo.mTtl, compressInfo)); - IncResourceRecordCount(responseHeader, false); + SuccessOrExit(error = AppendPtrRecord(aServiceFullName, aInstanceInfo.mFullName, aInstanceInfo.mTtl)); } for (uint8_t additional = 0; additional <= 1; additional++) { if (additional == 1) { - VerifyOrExit(!(mTestMode & kTestModeEmptyAdditionalSection)); + mAdditional = true; + VerifyOrExit(!(Get().mTestMode & kTestModeEmptyAdditionalSection)); } - if (HasQuestion(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), aInstanceInfo.mFullName, - ResourceRecord::kTypeSrv) == !additional) + if (HasQuestion(aInstanceInfo.mFullName, ResourceRecord::kTypeSrv) == !additional) { - SuccessOrExit(error = AppendSrvRecord(responseMessage, aInstanceInfo.mFullName, aInstanceInfo.mHostName, - aInstanceInfo.mTtl, aInstanceInfo.mPriority, aInstanceInfo.mWeight, - aInstanceInfo.mPort, compressInfo)); - IncResourceRecordCount(responseHeader, additional); + SuccessOrExit(error = AppendSrvRecord(aInstanceInfo.mFullName, aInstanceInfo.mHostName, aInstanceInfo.mTtl, + aInstanceInfo.mPriority, aInstanceInfo.mWeight, aInstanceInfo.mPort)); } - if (HasQuestion(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), aInstanceInfo.mFullName, - ResourceRecord::kTypeTxt) == !additional) + if (HasQuestion(aInstanceInfo.mFullName, ResourceRecord::kTypeTxt) == !additional) { - SuccessOrExit(error = AppendTxtRecord(responseMessage, aInstanceInfo.mFullName, aInstanceInfo.mTxtData, - aInstanceInfo.mTxtLength, aInstanceInfo.mTtl, compressInfo)); - IncResourceRecordCount(responseHeader, additional); + SuccessOrExit(error = AppendTxtRecord(aInstanceInfo.mFullName, aInstanceInfo.mTxtData, + aInstanceInfo.mTxtLength, aInstanceInfo.mTtl)); } - if (HasQuestion(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), aInstanceInfo.mHostName, - ResourceRecord::kTypeAaaa) == !additional) + if (HasQuestion(aInstanceInfo.mHostName, ResourceRecord::kTypeAaaa) == !additional) { for (uint8_t i = 0; i < aInstanceInfo.mAddressNum; i++) { @@ -1093,26 +1076,22 @@ void Server::AnswerQuery(QueryTransaction &aQuery, OT_ASSERT(!address.IsUnspecified() && !address.IsLinkLocal() && !address.IsMulticast() && !address.IsLoopback()); - SuccessOrExit(error = AppendAaaaRecord(responseMessage, aInstanceInfo.mHostName, address, - aInstanceInfo.mTtl, compressInfo)); - IncResourceRecordCount(responseHeader, additional); + SuccessOrExit(error = AppendAaaaRecord(aInstanceInfo.mHostName, address, aInstanceInfo.mTtl)); } } } exit: - FinalizeQuery(aQuery, error == kErrorNone ? Header::kResponseSuccess : Header::kResponseServerFailure); - ResetTimer(); + Finalize(error == kErrorNone ? Header::kResponseSuccess : Header::kResponseServerFailure); } -void Server::AnswerQuery(QueryTransaction &aQuery, const char *aHostFullName, const otDnssdHostInfo &aHostInfo) +void Server::QueryTransaction::Answer(const char *aHostFullName, const HostInfo &aHostInfo) { - Header &responseHeader = aQuery.GetResponseHeader(); - Message &responseMessage = aQuery.GetResponseMessage(); - Error error = kErrorNone; - NameCompressInfo &compressInfo = aQuery.GetNameCompressInfo(); + Error error = kErrorNone; - if (HasQuestion(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), aHostFullName, ResourceRecord::kTypeAaaa)) + mAdditional = false; + + if (HasQuestion(aHostFullName, ResourceRecord::kTypeAaaa)) { for (uint8_t i = 0; i < aHostInfo.mAddressNum; i++) { @@ -1121,30 +1100,23 @@ void Server::AnswerQuery(QueryTransaction &aQuery, const char *aHostFullName, co OT_ASSERT(!address.IsUnspecified() && !address.IsMulticast() && !address.IsLinkLocal() && !address.IsLoopback()); - SuccessOrExit(error = - AppendAaaaRecord(responseMessage, aHostFullName, address, aHostInfo.mTtl, compressInfo)); - IncResourceRecordCount(responseHeader, /* aAdditional */ false); + SuccessOrExit(error = AppendAaaaRecord(aHostFullName, address, aHostInfo.mTtl)); } } exit: - FinalizeQuery(aQuery, error == kErrorNone ? Header::kResponseSuccess : Header::kResponseServerFailure); - ResetTimer(); + Finalize(error == kErrorNone ? Header::kResponseSuccess : Header::kResponseServerFailure); } -void Server::SetQueryCallbacks(otDnssdQuerySubscribeCallback aSubscribe, - otDnssdQueryUnsubscribeCallback aUnsubscribe, - void *aContext) +void Server::SetQueryCallbacks(SubscribeCallback aSubscribe, UnsubscribeCallback aUnsubscribe, void *aContext) { OT_ASSERT((aSubscribe == nullptr) == (aUnsubscribe == nullptr)); - mQuerySubscribe = aSubscribe; - mQueryUnsubscribe = aUnsubscribe; - mQueryCallbackContext = aContext; + mQuerySubscribe.Set(aSubscribe, aContext); + mQueryUnsubscribe.Set(aUnsubscribe, aContext); } -void Server::HandleDiscoveredServiceInstance(const char *aServiceFullName, - const otDnssdServiceInstanceInfo &aInstanceInfo) +void Server::HandleDiscoveredServiceInstance(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo) { OT_ASSERT(StringEndsWith(aServiceFullName, Name::kLabelSeparatorChar)); OT_ASSERT(StringEndsWith(aInstanceInfo.mFullName, Name::kLabelSeparatorChar)); @@ -1152,22 +1124,22 @@ void Server::HandleDiscoveredServiceInstance(const char *a for (QueryTransaction &query : mQueryTransactions) { - if (query.IsValid() && CanAnswerQuery(query, aServiceFullName, aInstanceInfo)) + if (query.IsValid() && query.CanAnswer(aServiceFullName, aInstanceInfo)) { - AnswerQuery(query, aServiceFullName, aInstanceInfo); + query.Answer(aServiceFullName, aInstanceInfo); } } } -void Server::HandleDiscoveredHost(const char *aHostFullName, const otDnssdHostInfo &aHostInfo) +void Server::HandleDiscoveredHost(const char *aHostFullName, const HostInfo &aHostInfo) { OT_ASSERT(StringEndsWith(aHostFullName, Name::kLabelSeparatorChar)); for (QueryTransaction &query : mQueryTransactions) { - if (query.IsValid() && CanAnswerQuery(query, aHostFullName)) + if (query.IsValid() && query.CanAnswer(aHostFullName)) { - AnswerQuery(query, aHostFullName, aHostInfo); + query.Answer(aHostFullName, aHostInfo); } } } @@ -1198,69 +1170,71 @@ const otDnssdQuery *Server::GetNextQuery(const otDnssdQuery *aQuery) const Server::DnsQueryType Server::GetQueryTypeAndName(const otDnssdQuery *aQuery, char (&aName)[Name::kMaxNameSize]) { const QueryTransaction *query = static_cast(aQuery); + DnsQueryType type; OT_ASSERT(query->IsValid()); - return GetQueryTypeAndName(query->GetResponseHeader(), query->GetResponseMessage(), aName); + + query->GetQueryTypeAndName(type, aName); + + return type; } -Server::DnsQueryType Server::GetQueryTypeAndName(const Header &aHeader, - const Message &aMessage, - char (&aName)[Name::kMaxNameSize]) +void Server::Response::GetQueryTypeAndName(DnsQueryType &aType, char (&aName)[Name::kMaxNameSize]) const { - DnsQueryType sdType = kDnsQueryNone; + aType = kDnsQueryNone; - for (uint16_t i = 0, readOffset = sizeof(Header); i < aHeader.GetQuestionCount(); i++) + for (uint16_t i = 0, readOffset = sizeof(Header); i < mHeader.GetQuestionCount(); i++) { Question question; - IgnoreError(Name::ReadName(aMessage, readOffset, aName, sizeof(aName))); - IgnoreError(aMessage.Read(readOffset, question)); + IgnoreError(Name::ReadName(*mMessage, readOffset, aName, sizeof(aName))); + IgnoreError(mMessage->Read(readOffset, question)); readOffset += sizeof(question); switch (question.GetType()) { case ResourceRecord::kTypePtr: - ExitNow(sdType = kDnsQueryBrowse); + ExitNow(aType = kDnsQueryBrowse); case ResourceRecord::kTypeSrv: case ResourceRecord::kTypeTxt: - ExitNow(sdType = kDnsQueryResolve); + ExitNow(aType = kDnsQueryResolve); } } - for (uint16_t i = 0, readOffset = sizeof(Header); i < aHeader.GetQuestionCount(); i++) + for (uint16_t i = 0, readOffset = sizeof(Header); i < mHeader.GetQuestionCount(); i++) { Question question; - IgnoreError(Name::ReadName(aMessage, readOffset, aName, sizeof(aName))); - IgnoreError(aMessage.Read(readOffset, question)); + IgnoreError(Name::ReadName(*mMessage, readOffset, aName, sizeof(aName))); + IgnoreError(mMessage->Read(readOffset, question)); readOffset += sizeof(question); switch (question.GetType()) { case ResourceRecord::kTypeAaaa: case ResourceRecord::kTypeA: - ExitNow(sdType = kDnsQueryResolveHost); + ExitNow(aType = kDnsQueryResolveHost); } } exit: - return sdType; + return; } -bool Server::HasQuestion(const Header &aHeader, const Message &aMessage, const char *aName, uint16_t aQuestionType) +bool Server::Response::HasQuestion(const char *aName, uint16_t aQuestionType) const { bool found = false; - for (uint16_t i = 0, readOffset = sizeof(Header); i < aHeader.GetQuestionCount(); i++) + for (uint16_t i = 0, readOffset = sizeof(Header); i < mHeader.GetQuestionCount(); i++) { Question question; Error error; - error = Name::CompareName(aMessage, readOffset, aName); - IgnoreError(aMessage.Read(readOffset, question)); + error = Name::CompareName(*mMessage, readOffset, aName); + IgnoreError(mMessage->Read(readOffset, question)); readOffset += sizeof(question); - if (error == kErrorNone && aQuestionType == question.GetType()) + if ((error == kErrorNone) && (aQuestionType == question.GetType())) { ExitNow(found = true); } @@ -1272,21 +1246,23 @@ exit: void Server::HandleTimer(void) { - TimeMilli now = TimerMilli::GetNow(); + TimeMilli now = TimerMilli::GetNow(); + TimeMilli nextExpire = now.GetDistantFuture(); for (QueryTransaction &query : mQueryTransactions) { - TimeMilli expire; - if (!query.IsValid()) { continue; } - expire = query.GetStartTime() + kQueryTimeout; - if (expire <= now) + if (query.mExpireTime <= now) { - FinalizeQuery(query, Header::kResponseSuccess); + query.Finalize(Header::kResponseSuccess); + } + else + { + nextExpire = Min(nextExpire, query.mExpireTime); } } @@ -1302,87 +1278,37 @@ void Server::HandleTimer(void) { otPlatDnsCancelUpstreamQuery(&GetInstance(), &query); } + else + { + nextExpire = Min(nextExpire, query.GetExpireTime()); + } } #endif - ResetTimer(); -} - -void Server::ResetTimer(void) -{ - TimeMilli now = TimerMilli::GetNow(); - TimeMilli nextExpire = now.GetDistantFuture(); - - for (QueryTransaction &query : mQueryTransactions) + if (nextExpire != now.GetDistantFuture()) { - if (!query.IsValid()) - { - continue; - } - - nextExpire = Min(nextExpire, Max(now, query.GetStartTime() + kQueryTimeout)); - } - -#if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE - for (UpstreamQueryTransaction &query : mUpstreamQueryTransactions) - { - if (!query.IsValid()) - { - continue; - } - - nextExpire = Min(nextExpire, Max(now, query.GetExpireTime())); - } -#endif - - if (nextExpire < now.GetDistantFuture()) - { - mTimer.FireAt(nextExpire); - } - else - { - mTimer.Stop(); + mTimer.FireAtIfEarlier(nextExpire); } } -void Server::FinalizeQuery(QueryTransaction &aQuery, Header::Response aResponseCode) +void Server::QueryTransaction::Finalize(Header::Response aResponseCode) { char name[Name::kMaxNameSize]; DnsQueryType sdType; - OT_ASSERT(mQueryUnsubscribe != nullptr); - - sdType = GetQueryTypeAndName(aQuery.GetResponseHeader(), aQuery.GetResponseMessage(), name); + GetQueryTypeAndName(sdType, name); OT_ASSERT(sdType != kDnsQueryNone); OT_UNUSED_VARIABLE(sdType); - mQueryUnsubscribe(mQueryCallbackContext, name); - aQuery.Finalize(aResponseCode, mSocket); -} + Get().mQueryUnsubscribe.InvokeIfSet(name); -void Server::QueryTransaction::Init(const Header &aResponseHeader, - Message &aResponseMessage, - const NameCompressInfo &aCompressInfo, - const Ip6::MessageInfo &aMessageInfo, - Instance &aInstance) -{ - OT_ASSERT(mResponseMessage == nullptr); + mHeader.SetResponseCode(aResponseCode); + Send(mMessageInfo); - InstanceLocatorInit::Init(aInstance); - mResponseHeader = aResponseHeader; - mResponseMessage = &aResponseMessage; - mCompressInfo = aCompressInfo; - mMessageInfo = aMessageInfo; - mStartTime = TimerMilli::GetNow(); -} - -void Server::QueryTransaction::Finalize(Header::Response aResponseMessage, Ip6::Udp::Socket &aSocket) -{ - OT_ASSERT(mResponseMessage != nullptr); - - Get().SendResponse(mResponseHeader, aResponseMessage, *mResponseMessage, mMessageInfo, aSocket); - mResponseMessage = nullptr; + // Set the `mMessage` to null to indicate that + // `QueryTransaction` is unused. + mMessage = nullptr; } void Server::UpdateResponseCounters(Header::Response aResponseCode) diff --git a/src/core/net/dnssd_server.hpp b/src/core/net/dnssd_server.hpp index 5991e99d7..f5224ea3c 100644 --- a/src/core/net/dnssd_server.hpp +++ b/src/core/net/dnssd_server.hpp @@ -36,6 +36,7 @@ #include #include "common/as_core_type.hpp" +#include "common/callback.hpp" #include "common/message.hpp" #include "common/non_copyable.hpp" #include "common/timer.hpp" @@ -147,6 +148,12 @@ public: kDnsQueryResolveHost = OT_DNSSD_QUERY_TYPE_RESOLVE_HOST, ///< Service type resolve hostname. }; + typedef otDnssdServiceInstanceInfo ServiceInstanceInfo; ///< A discovered service instance for a DNS-SD query. + typedef otDnssdHostInfo HostInfo; ///< A discover host for a DNS-SD query. + + typedef otDnssdQuerySubscribeCallback SubscribeCallback; + typedef otDnssdQueryUnsubscribeCallback UnsubscribeCallback; + static constexpr uint16_t kPort = OPENTHREAD_CONFIG_DNSSD_SERVER_PORT; ///< The DNS-SD server port. /** @@ -180,9 +187,7 @@ public: * @param[in] aContext A pointer to the application-specific context. * */ - void SetQueryCallbacks(otDnssdQuerySubscribeCallback aSubscribe, - otDnssdQueryUnsubscribeCallback aUnsubscribe, - void *aContext); + void SetQueryCallbacks(SubscribeCallback aSubscribe, UnsubscribeCallback aUnsubscribe, void *aContext); /** * Notifies a discovered service instance. @@ -191,7 +196,7 @@ public: * @param[in] aInstanceInfo A reference to the discovered service instance information. * */ - void HandleDiscoveredServiceInstance(const char *aServiceFullName, const otDnssdServiceInstanceInfo &aInstanceInfo); + void HandleDiscoveredServiceInstance(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo); #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE /** @@ -231,7 +236,7 @@ public: * @param[in] aHostInfo A reference to the discovered host information. * */ - void HandleDiscoveredHost(const char *aHostFullName, const otDnssdHostInfo &aHostInfo); + void HandleDiscoveredHost(const char *aHostFullName, const HostInfo &aHostInfo); /** * Acquires the next query in the server. @@ -289,25 +294,14 @@ private: class NameCompressInfo : public Clearable { public: - explicit NameCompressInfo(void) = default; - - explicit NameCompressInfo(const char *aDomainName) - : mDomainName(aDomainName) - , mDomainNameOffset(kUnknownOffset) - , mServiceNameOffset(kUnknownOffset) - , mInstanceNameOffset(kUnknownOffset) - , mHostNameOffset(kUnknownOffset) - { - } - static constexpr uint16_t kUnknownOffset = 0; // Unknown offset value (used when offset is not yet set). + NameCompressInfo(void) { Clear(); } + uint16_t GetDomainNameOffset(void) const { return mDomainNameOffset; } void SetDomainNameOffset(uint16_t aOffset) { mDomainNameOffset = aOffset; } - const char *GetDomainName(void) const { return mDomainName; } - uint16_t GetServiceNameOffset(const Message &aMessage, const char *aServiceName) const { return MatchCompressedName(aMessage, mServiceNameOffset, aServiceName) @@ -357,11 +351,10 @@ private: return aOffset != kUnknownOffset && Name::CompareName(aMessage, aOffset, aName) == kErrorNone; } - const char *mDomainName; // The serialized domain name. - uint16_t mDomainNameOffset; // Offset of domain name serialization into the response message. - uint16_t mServiceNameOffset; // Offset of service name serialization into the response message. - uint16_t mInstanceNameOffset; // Offset of instance name serialization into the response message. - uint16_t mHostNameOffset; // Offset of host name serialization into the response message. + uint16_t mDomainNameOffset; // Offset of domain name serialization into the response message. + uint16_t mServiceNameOffset; // Offset of service name serialization into the response message. + uint16_t mInstanceNameOffset; // Offset of instance name serialization into the response message. + uint16_t mHostNameOffset; // Offset of host name serialization into the response message. }; static constexpr bool kBindUnspecifiedNetif = OPENTHREAD_CONFIG_DNSSD_SERVER_BIND_UNSPECIFIED_NETIF; @@ -390,140 +383,91 @@ private: bool IsHostName(void) const { return mProtocolOffset == kNotPresent && mDomainOffset != 0; } - uint8_t mDomainOffset; // The offset to the beginning of . - uint8_t mProtocolOffset; // The offset to the beginning of (i.e. _tcp or _udp) or `kNotPresent` if - // the name is not a service or instance. - uint8_t mServiceOffset; // The offset to the beginning of or `kNotPresent` if the name is not a - // service or instance. - uint8_t mSubTypeOffset; // The offset to the beginning of sub-type label or `kNotPresent` is not a sub-type. - uint8_t mInstanceOffset; // The offset to the beginning of or `kNotPresent` if the name is not a - // instance. + uint8_t mDomainOffset; // Offset to . + uint8_t mProtocolOffset; // Offset to (i.e. _tcp or _udp) or `kNotPresent` if not service name. + uint8_t mServiceOffset; // Offset to or `kNotPresent` if not service or instance. + uint8_t mSubTypeOffset; // Offset to sub-type label or `kNotPresent` is not a sub-type. + uint8_t mInstanceOffset; // Offset to or `kNotPresent` if the name is not a instance. }; - /** - * Contains the compress information for a dns packet. - * - */ - class QueryTransaction : public InstanceLocatorInit + struct Request { - public: - explicit QueryTransaction(void) - : mResponseMessage(nullptr) + const Message *mMessage; + const Ip6::MessageInfo *mMessageInfo; + Header mHeader; + }; + + struct Response : public GetProvider + { + Response(void) + : mMessage(nullptr) + , mAdditional(false) { } - void Init(const Header &aResponseHeader, - Message &aResponseMessage, - const NameCompressInfo &aCompressInfo, - const Ip6::MessageInfo &aMessageInfo, - Instance &aInstance); - bool IsValid(void) const { return mResponseMessage != nullptr; } - const Ip6::MessageInfo &GetMessageInfo(void) const { return mMessageInfo; } - const Header &GetResponseHeader(void) const { return mResponseHeader; } - Header &GetResponseHeader(void) { return mResponseHeader; } - const Message &GetResponseMessage(void) const { return *mResponseMessage; } - Message &GetResponseMessage(void) { return const_cast(*mResponseMessage); } - TimeMilli GetStartTime(void) const { return mStartTime; } - NameCompressInfo &GetNameCompressInfo(void) { return mCompressInfo; }; - void Finalize(Header::Response aResponseMessage, Ip6::Udp::Socket &aSocket); + Instance &GetInstance(void) const { return mMessage->GetInstance(); } - Header mResponseHeader; - Message *mResponseMessage; + Error AddQuestionsFrom(const Request &aRequest); + Error AppendQuestion(const char *aName, const Question &aQuestion); + Error AppendPtrRecord(const char *aServiceName, const char *aInstanceName, uint32_t aTtl); + Error AppendSrvRecord(const char *aInstanceName, + const char *aHostName, + uint32_t aTtl, + uint16_t aPriority, + uint16_t aWeight, + uint16_t aPort); + Error AppendTxtRecord(const char *aInstanceName, const void *aTxtData, uint16_t aTxtLength, uint32_t aTtl); + Error AppendAaaaRecord(const char *aHostName, const Ip6::Address &aAddress, uint32_t aTtl); + Error AppendServiceName(const char *aName); + Error AppendInstanceName(const char *aName); + Error AppendHostName(const char *aName); + void IncResourceRecordCount(void); + bool HasQuestion(const char *aName, uint16_t aQuestionType) const; + void Send(const Ip6::MessageInfo &aMessageInfo); + void GetQueryTypeAndName(DnsQueryType &aType, char (&aName)[Name::kMaxNameSize]) const; + +#if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE + void ResolveBySrp(void); + void ResolveQuestionBySrp(const char *aName, const Question &aQuestion); +#endif + + Message *mMessage; + Header mHeader; NameCompressInfo mCompressInfo; + bool mAdditional; // Whether or not appending new records in additional data section. + }; + + struct QueryTransaction : public Response + { + bool IsValid(void) const { return mMessage != nullptr; } + bool CanAnswer(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo) const; + bool CanAnswer(const char *aHostFullName) const; + void Answer(const char *aServiceFullName, const ServiceInstanceInfo &aInstanceInfo); + void Answer(const char *aHostFullName, const HostInfo &aHostInfo); + void Finalize(Header::Response aResponseCode); + Ip6::MessageInfo mMessageInfo; - TimeMilli mStartTime; + TimeMilli mExpireTime; }; 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(const Header &aRequestHeader, Message &aRequestMessage, const Ip6::MessageInfo &aMessageInfo); - static Header::Response AddQuestions(const Header &aRequestHeader, - const Message &aRequestMessage, - Header &aResponseHeader, - Message &aResponseMessage, - NameCompressInfo &aCompressInfo); - static Error AppendQuestion(const char *aName, - const Question &aQuestion, - Message &aMessage, - NameCompressInfo &aCompressInfo); - static Error AppendPtrRecord(Message &aMessage, - const char *aServiceName, - const char *aInstanceName, - uint32_t aTtl, - NameCompressInfo &aCompressInfo); - static Error AppendSrvRecord(Message &aMessage, - const char *aInstanceName, - const char *aHostName, - uint32_t aTtl, - uint16_t aPriority, - uint16_t aWeight, - uint16_t aPort, - NameCompressInfo &aCompressInfo); - static Error AppendTxtRecord(Message &aMessage, - const char *aInstanceName, - const void *aTxtData, - uint16_t aTxtLength, - uint32_t aTtl, - NameCompressInfo &aCompressInfo); - static Error AppendAaaaRecord(Message &aMessage, - const char *aHostName, - const Ip6::Address &aAddress, - uint32_t aTtl, - NameCompressInfo &aCompressInfo); - static Error AppendServiceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); - static Error AppendInstanceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); - static Error AppendHostName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); - static void IncResourceRecordCount(Header &aHeader, bool aAdditional); - static Error FindNameComponents(const char *aName, const char *aDomain, NameComponentsOffsetInfo &aInfo); - static Error FindPreviousLabel(const char *aName, uint8_t &aStart, uint8_t &aStop); - void SendResponse(Header aHeader, - Header::Response aResponseCode, - Message &aMessage, - const Ip6::MessageInfo &aMessageInfo, - Ip6::Udp::Socket &aSocket); -#if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE - Header::Response ResolveBySrp(Header &aResponseHeader, - Message &aResponseMessage, - Server::NameCompressInfo &aCompressInfo); - Header::Response ResolveQuestionBySrp(const char *aName, - const Question &aQuestion, - Header &aResponseHeader, - Message &aResponseMessage, - NameCompressInfo &aCompressInfo, - bool aAdditional); -#endif + 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(const Request &aRequest); + static Error FindNameComponents(const char *aName, const char *aDomain, NameComponentsOffsetInfo &aInfo); + static Error FindPreviousLabel(const char *aName, uint8_t &aStart, uint8_t &aStop); #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE - static bool ShouldForwardToUpstream(const Header &aRequestHeader, const Message &aRequestMessage); + static bool ShouldForwardToUpstream(const Request &aRequest); UpstreamQueryTransaction *AllocateUpstreamQueryTransaction(const Ip6::MessageInfo &aMessageInfo); void ResetUpstreamQueryTransaction(UpstreamQueryTransaction &aTxn, Error aError); - Error ResolveByUpstream(const Message &aRequestMessage, const Ip6::MessageInfo &aMessageInfo); + Error ResolveByUpstream(const Request &aRequest); #endif - Error ResolveByQueryCallbacks(Header &aResponseHeader, - Message &aResponseMessage, - NameCompressInfo &aCompressInfo, - const Ip6::MessageInfo &aMessageInfo); - QueryTransaction *NewQuery(const Header &aResponseHeader, - Message &aResponseMessage, - const NameCompressInfo &aCompressInfo, - const Ip6::MessageInfo &aMessageInfo); - static bool CanAnswerQuery(const QueryTransaction &aQuery, - const char *aServiceFullName, - const otDnssdServiceInstanceInfo &aInstanceInfo); - void AnswerQuery(QueryTransaction &aQuery, - const char *aServiceFullName, - const otDnssdServiceInstanceInfo &aInstanceInfo); - static bool CanAnswerQuery(const Server::QueryTransaction &aQuery, const char *aHostFullName); - void AnswerQuery(QueryTransaction &aQuery, const char *aHostFullName, const otDnssdHostInfo &aHostInfo); - void FinalizeQuery(QueryTransaction &aQuery, Header::Response aResponseCode); - static DnsQueryType GetQueryTypeAndName(const Header &aHeader, - const Message &aMessage, - char (&aName)[Name::kMaxNameSize]); - static bool HasQuestion(const Header &aHeader, const Message &aMessage, const char *aName, uint16_t aQuestionType); + Error ResolveByQueryCallbacks(Response &aResponse, const Ip6::MessageInfo &aMessageInfo); + QueryTransaction *NewQuery(Response &aResponse, const Ip6::MessageInfo &aMessageInfo); void HandleTimer(void); void ResetTimer(void); @@ -539,10 +483,9 @@ private: Ip6::Udp::Socket mSocket; - QueryTransaction mQueryTransactions[kMaxConcurrentQueries]; - void *mQueryCallbackContext; - otDnssdQuerySubscribeCallback mQuerySubscribe; - otDnssdQueryUnsubscribeCallback mQueryUnsubscribe; + QueryTransaction mQueryTransactions[kMaxConcurrentQueries]; + Callback mQuerySubscribe; + Callback mQueryUnsubscribe; static const char *kBlockedDomains[]; #if OPENTHREAD_CONFIG_DNS_UPSTREAM_QUERY_ENABLE