diff --git a/src/core/net/dnssd_server.cpp b/src/core/net/dnssd_server.cpp index bf94f3ed5..c066ddb5f 100644 --- a/src/core/net/dnssd_server.cpp +++ b/src/core/net/dnssd_server.cpp @@ -113,141 +113,128 @@ exit: void Server::ProcessQuery(Message &aMessage, Message &aResponse, const Header &aRequestHeader) { Header responseHeader; - uint16_t readOffset, nameSerializeOffset; + uint16_t readOffset; Question question; - uint16_t qtype; char name[Dns::Name::kMaxNameSize]; - otError error = OT_ERROR_NONE; - NameCompressInfo compressInfo; + NameCompressInfo compressInfo(kDefaultDomainName); + Header::Response response = Header::Response::kResponseSuccess; + otError error = OT_ERROR_NONE; + uint8_t resolveAdditional = kResolveAdditionalAll; // Setup initial DNS response header responseHeader.Clear(); - responseHeader.SetResponseCode(Header::kResponseSuccess); responseHeader.SetType(Header::kTypeResponse); responseHeader.SetMessageId(aRequestHeader.GetMessageId()); // Validate the query VerifyOrExit(aRequestHeader.GetQueryType() == Header::kQueryTypeStandard, - responseHeader.SetResponseCode(Header::kResponseNotImplemented)); - VerifyOrExit(!aRequestHeader.IsTruncationFlagSet(), responseHeader.SetResponseCode(Header::kResponseFormatError)); - VerifyOrExit(aRequestHeader.GetQuestionCount() == 1, - responseHeader.SetResponseCode(Header::kResponseNotImplemented)); + response = Header::kResponseNotImplemented); + VerifyOrExit(!aRequestHeader.IsTruncationFlagSet(), response = Header::kResponseFormatError); + VerifyOrExit(aRequestHeader.GetQuestionCount() > 0, response = Header::kResponseFormatError); - // Read the query name and question readOffset = sizeof(Header); - VerifyOrExit(OT_ERROR_NONE == Dns::Name::ReadName(aMessage, readOffset, name, sizeof(name)), - responseHeader.SetResponseCode(Header::kResponseFormatError)); - VerifyOrExit(OT_ERROR_NONE == aMessage.Read(readOffset, question), - responseHeader.SetResponseCode(Header::kResponseFormatError)); - // Add the question to the response, and save the serialize offset of name - nameSerializeOffset = aResponse.GetLength(); - SuccessOrExit(error = AddQuestionToResponse(name, question, aResponse, responseHeader)); - - // Further validate the query - qtype = question.GetType(); - VerifyOrExit(qtype == ResourceRecord::kTypePtr || qtype == ResourceRecord::kTypeSrv || - qtype == ResourceRecord::kTypeTxt || qtype == ResourceRecord::kTypeAaaa, - responseHeader.SetResponseCode(Header::kResponseNotImplemented)); - - VerifyOrExit(question.GetClass() == ResourceRecord::kClassInternet || - question.GetClass() == ResourceRecord::kClassAny, - responseHeader.SetResponseCode(Header::kResponseNotImplemented)); - - // Prepare the information for name compression - VerifyOrExit(OT_ERROR_NONE == PrepareCompressInfo(qtype, name, nameSerializeOffset, compressInfo), - responseHeader.SetResponseCode(Header::kResponseNameError)); - - // Resolve the question - SuccessOrExit(error = ResolveQuestion(name, question, responseHeader, aResponse, compressInfo)); - - otLogInfoDns("[server] TRANSACTION=0x%04x, QUESTION=[%s %d %d], RCODE=%d, ANSWER=%d, ADDITIONAL=%d", - aRequestHeader.GetMessageId(), name, question.GetClass(), question.GetType(), - responseHeader.GetResponseCode(), responseHeader.GetQuestionCount(), - responseHeader.GetAdditionalRecordCount()); -exit: - if (error != OT_ERROR_NONE) + // Check and append the questions + for (uint16_t i = 0; i < aRequestHeader.GetQuestionCount(); i++) { - otLogWarnDns("[server] failed to handle DNS query: %s", otThreadErrorToString(error)); + NameComponentsOffsetInfo nameComponentsOffsetInfo; + VerifyOrExit(OT_ERROR_NONE == Dns::Name::ReadName(aMessage, readOffset, name, sizeof(name)), + response = Header::kResponseFormatError); + VerifyOrExit(OT_ERROR_NONE == aMessage.Read(readOffset, question), response = Header::kResponseFormatError); + readOffset += sizeof(question); + + uint16_t qtype = question.GetType(); + + VerifyOrExit(qtype == ResourceRecord::kTypePtr || qtype == ResourceRecord::kTypeSrv || + qtype == ResourceRecord::kTypeTxt || qtype == ResourceRecord::kTypeAaaa, + response = Header::kResponseNotImplemented); + + VerifyOrExit(OT_ERROR_NONE == FindNameComponents(name, compressInfo.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); + resolveAdditional &= ~kResolveAdditionalSrv; + break; + case ResourceRecord::kTypeTxt: + VerifyOrExit(nameComponentsOffsetInfo.IsServiceInstanceName(), response = Header::kResponseNameError); + resolveAdditional &= ~kResolveAdditionalTxt; + break; + case ResourceRecord::kTypeAaaa: + VerifyOrExit(nameComponentsOffsetInfo.IsHostName(), response = Header::kResponseNameError); + resolveAdditional &= ~kResolveAdditionalAaaa; + break; + default: + ExitNow(response = Header::kResponseNotImplemented); + } + + SuccessOrExit(error = AppendQuestion(name, question, aResponse, compressInfo)); + } + + responseHeader.SetQuestionCount(aRequestHeader.GetQuestionCount()); + + // Answer the questions + readOffset = sizeof(Header); + for (uint16_t i = 0; i < aRequestHeader.GetQuestionCount(); i++) + { + uint8_t resolveKind = kResolveAnswer; + + IgnoreError(Dns::Name::ReadName(aMessage, readOffset, name, sizeof(name))); + IgnoreError(aMessage.Read(readOffset, question)); + readOffset += sizeof(question); + + response = ResolveQuestion(name, question, responseHeader, aResponse, resolveKind, compressInfo); + + otLogInfoDns("[server] ANSWER: TRANSACTION=0x%04x, QUESTION=[%s %d %d], RCODE=%d", + aRequestHeader.GetMessageId(), name, question.GetClass(), question.GetType(), response); + } + + // Answer the questions with additional RRs if required + VerifyOrExit(resolveAdditional != kResolveNone); + + readOffset = sizeof(Header); + for (uint16_t i = 0; i < aRequestHeader.GetQuestionCount(); i++) + { + IgnoreError(Dns::Name::ReadName(aMessage, readOffset, name, sizeof(name))); + IgnoreError(aMessage.Read(readOffset, question)); + readOffset += sizeof(question); + + VerifyOrExit(Header::kResponseServerFailure != + ResolveQuestion(name, question, responseHeader, aResponse, resolveAdditional, compressInfo), + response = Header::kResponseServerFailure); + + otLogInfoDns("[server] ADDITIONAL: TRANSACTION=0x%04x, QUESTION=[%s %d %d], RCODE=%d", + aRequestHeader.GetMessageId(), name, question.GetClass(), question.GetType(), response); + } + +exit: + response = (error == OT_ERROR_NONE) ? response : Header::Response::kResponseServerFailure; + + if (response == Header::Response::kResponseServerFailure) + { + otLogWarnDns("[server] failed to handle DNS query due to server failure"); responseHeader.SetQuestionCount(0); responseHeader.SetAnswerCount(0); responseHeader.SetAdditionalRecordCount(0); - responseHeader.SetResponseCode(Header::kResponseServerFailure); IgnoreError(aResponse.SetLength(sizeof(Header))); } + responseHeader.SetResponseCode(response); aResponse.Write(0, responseHeader); } -otError Server::AddQuestionToResponse(const char * aName, - const Question &aQuestion, - Message & aResponse, - Header & aResponseHeader) -{ - otError error = OT_ERROR_NONE; - - SuccessOrExit(error = Dns::Name::AppendName(aName, aResponse)); - SuccessOrExit(error = aResponse.Append(aQuestion)); - - aResponseHeader.SetQuestionCount(1); - -exit: - return error; -} - -otError Server::PrepareCompressInfo(uint16_t aQueryType, - const char * aName, - uint16_t aNameSerializeOffset, - Server::NameCompressInfo &aCompressInfo) -{ - const char * domain = kDefaultDomainName; - NameComponentsOffsetInfo nameComponentsInfo; - otError error = OT_ERROR_NONE; - -#if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE - domain = Get().GetDomain(); -#endif - - SuccessOrExit(error = FindNameComponents(aName, domain, nameComponentsInfo)); - - switch (aQueryType) - { - case ResourceRecord::kTypePtr: - VerifyOrExit(nameComponentsInfo.IsServiceName(), error = OT_ERROR_INVALID_ARGS); - aCompressInfo.SetServiceNameOffset(aNameSerializeOffset, aName); - break; - - case ResourceRecord::kTypeSrv: - case ResourceRecord::kTypeTxt: - VerifyOrExit(nameComponentsInfo.IsServiceInstanceName(), error = OT_ERROR_INVALID_ARGS); - aCompressInfo.SetInstanceNameOffset(aNameSerializeOffset, aName); - aCompressInfo.SetServiceNameOffset(aNameSerializeOffset + nameComponentsInfo.mServiceOffset, - aName + nameComponentsInfo.mServiceOffset); - break; - - case ResourceRecord::kTypeAaaa: - VerifyOrExit(nameComponentsInfo.IsHostName(), error = OT_ERROR_INVALID_ARGS); - aCompressInfo.SetHostNameOffset(aNameSerializeOffset, aName); - break; - - default: - OT_ASSERT(false); - } - - OT_ASSERT(nameComponentsInfo.mDomainOffset != NameComponentsOffsetInfo::kNotPresent); - aCompressInfo.SetDomainNameOffset(aNameSerializeOffset + nameComponentsInfo.mDomainOffset, - aName + nameComponentsInfo.mDomainOffset); - -exit: - return error; -} - -otError Server::ResolveQuestion(const char * aName, - const Question & aQuestion, - Header & aResponseHeader, - Message & aResponseMessage, - NameCompressInfo &aCompressInfo) +Header::Response Server::ResolveQuestion(const char * aName, + const Question & aQuestion, + Header & aResponseHeader, + Message & aResponseMessage, + uint8_t aResolveKind, + NameCompressInfo &aCompressInfo) { OT_UNUSED_VARIABLE(aName); OT_UNUSED_VARIABLE(aQuestion); @@ -255,24 +242,41 @@ otError Server::ResolveQuestion(const char * aName, OT_UNUSED_VARIABLE(aResponseMessage); OT_UNUSED_VARIABLE(aCompressInfo); - otError error = OT_ERROR_NONE; + Header::Response response = Header::kResponseNameError; #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE - SuccessOrExit(error = ResolveQuestionBySrp(aName, aQuestion, aResponseHeader, aResponseMessage, - /* aAdditional */ false, aCompressInfo)); + response = ResolveQuestionBySrp(aName, aQuestion, aResponseHeader, aResponseMessage, aResolveKind, aCompressInfo); +#endif - if (aResponseHeader.GetAnswerCount() > 0) + return response; +} + +otError Server::AppendQuestion(const char * aName, + const Question & aQuestion, + Message & aMessage, + NameCompressInfo &aCompressInfo) +{ + otError error = OT_ERROR_NONE; + + switch (aQuestion.GetType()) { - SuccessOrExit(error = ResolveQuestionBySrp(aName, aQuestion, aResponseHeader, aResponseMessage, - /* aAdditional */ true, aCompressInfo)); - } - else - { - aResponseHeader.SetResponseCode(Header::kResponseNameError); + case ResourceRecord::kTypePtr: + SuccessOrExit(error = AppendServiceName(aMessage, aName, aCompressInfo)); + break; + case ResourceRecord::kTypeSrv: + case ResourceRecord::kTypeTxt: + SuccessOrExit(error = AppendInstanceName(aMessage, aName, aCompressInfo)); + break; + case ResourceRecord::kTypeAaaa: + SuccessOrExit(error = AppendHostName(aMessage, aName, aCompressInfo)); + break; + default: + OT_ASSERT(false); } + error = aMessage.Append(aQuestion); + exit: -#endif return error; } @@ -358,31 +362,74 @@ exit: otError Server::AppendServiceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo) { - OT_UNUSED_VARIABLE(aName); - OT_ASSERT(strcmp(aCompressInfo.GetServiceName(), aName) == 0); + otError error; + uint16_t serviceCompressOffset = aCompressInfo.GetServiceNameOffset(aName); - return Dns::Name::AppendPointerLabel(aCompressInfo.GetServiceNameOffset(), aMessage); + if (serviceCompressOffset != NameCompressInfo::kUnknownOffset) + { + error = Dns::Name::AppendPointerLabel(serviceCompressOffset, aMessage); + } + else + { + uint8_t domainStart = static_cast(StringLength(aName, Name::kMaxNameSize - 1) - + StringLength(aCompressInfo.GetDomainName(), Name::kMaxNameSize - 1)); + uint16_t domainCompressOffset = aCompressInfo.GetDomainNameOffset(); + + serviceCompressOffset = aMessage.GetLength(); + aCompressInfo.SetServiceNameOffset(serviceCompressOffset, aName); + + if (domainCompressOffset == NameCompressInfo::kUnknownOffset) + { + aCompressInfo.SetDomainNameOffset(serviceCompressOffset + domainStart); + error = Dns::Name::AppendName(aName, aMessage); + } + else + { + SuccessOrExit(error = Dns::Name::AppendMultipleLabels(aName, domainStart, aMessage)); + error = Dns::Name::AppendPointerLabel(domainCompressOffset, aMessage); + } + } + +exit: + return error; } otError Server::AppendInstanceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo) { otError error; - uint16_t nameOffset = aCompressInfo.GetInstanceNameOffset(aName); + uint16_t instanceCompressOffset = aCompressInfo.GetInstanceNameOffset(aName); - if (nameOffset != NameCompressInfo::kUnknownOffset) + if (instanceCompressOffset != NameCompressInfo::kUnknownOffset) { - error = Dns::Name::AppendPointerLabel(nameOffset, aMessage); + error = Dns::Name::AppendPointerLabel(instanceCompressOffset, aMessage); } else { - uint8_t serviceStart = static_cast(StringLength(aName, Name::kMaxNameLength) - - StringLength(aCompressInfo.GetServiceName(), Name::kMaxNameLength)); + NameComponentsOffsetInfo nameComponentsInfo; + + IgnoreError(FindNameComponents(aName, aCompressInfo.GetDomainName(), nameComponentsInfo)); + OT_ASSERT(nameComponentsInfo.IsServiceInstanceName()); aCompressInfo.SetInstanceNameOffset(aMessage.GetLength(), aName); - SuccessOrExit(error = Dns::Name::AppendLabel(aName, serviceStart - 1, aMessage)); - error = Dns::Name::AppendPointerLabel(aCompressInfo.GetServiceNameOffset(), aMessage); + // Append the instance name as one label + SuccessOrExit(error = Dns::Name::AppendLabel(aName, nameComponentsInfo.mServiceOffset - 1, aMessage)); + + { + const char *serviceName = aName + nameComponentsInfo.mServiceOffset; + uint16_t serviceCompressOffset = aCompressInfo.GetServiceNameOffset(serviceName); + + if (serviceCompressOffset != NameCompressInfo::kUnknownOffset) + { + error = Dns::Name::AppendPointerLabel(serviceCompressOffset, aMessage); + } + else + { + aCompressInfo.SetServiceNameOffset(aMessage.GetLength(), serviceName); + error = Dns::Name::AppendName(serviceName, aMessage); + } + } } exit: @@ -400,13 +447,23 @@ otError Server::AppendHostName(Message &aMessage, const char *aName, NameCompres } else { - uint8_t domainStart = static_cast(StringLength(aName, Name::kMaxNameLength) - - StringLength(aCompressInfo.GetDomainName(), Name::kMaxNameLength)); + uint8_t domainStart = static_cast(StringLength(aName, Name::kMaxNameLength) - + StringLength(aCompressInfo.GetDomainName(), Name::kMaxNameSize - 1)); + uint16_t domainCompressOffset = aCompressInfo.GetDomainNameOffset(); - aCompressInfo.SetHostNameOffset(aMessage.GetLength(), aName); + hostCompressOffset = aMessage.GetLength(); + aCompressInfo.SetHostNameOffset(hostCompressOffset, aName); - SuccessOrExit(error = Dns::Name::AppendMultipleLabels(aName, domainStart, aMessage)); - error = Dns::Name::AppendPointerLabel(aCompressInfo.GetDomainNameOffset(), aMessage); + if (domainCompressOffset == NameCompressInfo::kUnknownOffset) + { + aCompressInfo.SetDomainNameOffset(hostCompressOffset + domainStart); + error = Dns::Name::AppendName(aName, aMessage); + } + else + { + SuccessOrExit(error = Dns::Name::AppendMultipleLabels(aName, domainStart, aMessage)); + error = Dns::Name::AppendPointerLabel(domainCompressOffset, aMessage); + } } exit: @@ -498,17 +555,18 @@ exit: } #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE -otError Server::ResolveQuestionBySrp(const char * aName, - const Question & aQuestion, - Header & aResponseHeader, - Message & aResponseMessage, - bool aAdditional, - NameCompressInfo &aCompressInfo) +Header::Response Server::ResolveQuestionBySrp(const char * aName, + const Question & aQuestion, + Header & aResponseHeader, + Message & aResponseMessage, + uint8_t aResolveKind, + NameCompressInfo &aCompressInfo) { - otError error = OT_ERROR_NONE; - const Srp::Server::Host *host = nullptr; - TimeMilli now = TimerMilli::GetNow(); - uint16_t qtype = aQuestion.GetType(); + otError error = OT_ERROR_NONE; + const Srp::Server::Host *host = nullptr; + TimeMilli now = TimerMilli::GetNow(); + uint16_t qtype = aQuestion.GetType(); + Header::Response response = Header::kResponseNameError; while ((host = GetNextSrpHost(host)) != nullptr) { @@ -535,33 +593,38 @@ otError Server::ResolveQuestionBySrp(const char * aName, needAdditionalAaaaRecord = true; } - if (!aAdditional && ptrQueryMatched) + if (aResolveKind == kResolveAnswer && ptrQueryMatched) { SuccessOrExit( error = AppendPtrRecord(aResponseMessage, aName, instanceName, instanceTtl, aCompressInfo)); - IncResourceRecordCount(aResponseHeader, aAdditional); + IncResourceRecordCount(aResponseHeader, aResolveKind != kResolveAnswer); + response = Header::Response::kResponseSuccess; } - if ((!aAdditional && srvQueryMatched) || (aAdditional && ptrQueryMatched)) + if ((aResolveKind == kResolveAnswer && srvQueryMatched) || + ((aResolveKind & kResolveAdditionalSrv) && ptrQueryMatched)) { SuccessOrExit(error = AppendSrvRecord(aResponseMessage, instanceName, hostName, instanceTtl, service->GetPriority(), service->GetWeight(), service->GetPort(), aCompressInfo)); - IncResourceRecordCount(aResponseHeader, aAdditional); + IncResourceRecordCount(aResponseHeader, aResolveKind != kResolveAnswer); + response = Header::Response::kResponseSuccess; } - if ((!aAdditional && txtQueryMatched) || (aAdditional && ptrQueryMatched)) + if ((aResolveKind == kResolveAnswer && txtQueryMatched) || + ((aResolveKind & kResolveAdditionalTxt) && ptrQueryMatched)) { SuccessOrExit( error = AppendTxtRecord(aResponseMessage, instanceName, *service, instanceTtl, aCompressInfo)); - IncResourceRecordCount(aResponseHeader, aAdditional); + IncResourceRecordCount(aResponseHeader, aResolveKind != kResolveAnswer); + response = Header::Response::kResponseSuccess; } } } // Handle AAAA query - if ((!aAdditional && qtype == ResourceRecord::kTypeAaaa && host->Matches(aName)) || - (aAdditional && needAdditionalAaaaRecord)) + if ((aResolveKind == kResolveAnswer && qtype == ResourceRecord::kTypeAaaa && host->Matches(aName)) || + ((aResolveKind & kResolveAdditionalAaaa) && needAdditionalAaaaRecord)) { uint8_t addrNum; const Ip6::Address *addrs = host->GetAddresses(addrNum); @@ -570,13 +633,15 @@ otError 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); + IncResourceRecordCount(aResponseHeader, aResolveKind != kResolveAnswer); } + + response = Header::Response::kResponseSuccess; } } exit: - return error; + return error == OT_ERROR_NONE ? response : Header::Response::kResponseServerFailure; } const Srp::Server::Host *Server::GetNextSrpHost(const Srp::Server::Host *aHost) diff --git a/src/core/net/dnssd_server.hpp b/src/core/net/dnssd_server.hpp index 191c557fa..b026ad003 100644 --- a/src/core/net/dnssd_server.hpp +++ b/src/core/net/dnssd_server.hpp @@ -89,6 +89,16 @@ private: kProtocolLabelLength = 4, }; + enum : uint8_t + { + kResolveNone = 0, + kResolveAnswer = 1u << 0, + kResolveAdditionalSrv = 1u << 1, + kResolveAdditionalTxt = 1u << 2, + kResolveAdditionalAaaa = 1u << 3, + kResolveAdditionalAll = kResolveAdditionalSrv | kResolveAdditionalTxt | kResolveAdditionalAaaa, + }; + class NameCompressInfo : public Clearable { public: @@ -97,38 +107,43 @@ private: kUnknownOffset = 0, // Unknown offset value (used when offset is not yet set). }; - explicit NameCompressInfo(void) { Clear(); } - - uint16_t GetDomainNameOffset(void) const + explicit NameCompressInfo(const char *aDomainName) + : mDomainName(aDomainName) + , mServiceName(nullptr) + , mInstanceName(nullptr) + , mHostName(nullptr) + , mDomainNameOffset(kUnknownOffset) + , mServiceNameOffset(kUnknownOffset) + , mInstanceNameOffset(kUnknownOffset) + , mHostNameOffset(kUnknownOffset) { - OT_ASSERT(mDomainNameOffset != kUnknownOffset); - - return mDomainNameOffset; } - void SetDomainNameOffset(uint16_t aOffset, const char *aName) - { - OT_ASSERT(mDomainName == nullptr); + uint16_t GetDomainNameOffset(void) const { return mDomainNameOffset; } - mDomainName = aName; - mDomainNameOffset = aOffset; - } + void SetDomainNameOffset(uint16_t aOffset) { mDomainNameOffset = aOffset; } const char *GetDomainName(void) const { return mDomainName; } - uint16_t GetServiceNameOffset(void) const + uint16_t GetServiceNameOffset(const char *aServiceName) const { - OT_ASSERT(mServiceNameOffset != kUnknownOffset); + uint16_t offset = mServiceNameOffset; - return mServiceNameOffset; + if (offset != kUnknownOffset && strcmp(aServiceName, mServiceName) != 0) + { + offset = kUnknownOffset; + } + + return offset; }; void SetServiceNameOffset(uint16_t aOffset, const char *aName) { - OT_ASSERT(mServiceName == nullptr); - - mServiceName = aName; - mServiceNameOffset = aOffset; + if (mServiceName == nullptr) + { + mServiceName = aName; + mServiceNameOffset = aOffset; + } } const char *GetServiceName() const { return mServiceName; } @@ -176,14 +191,14 @@ private: } private: - const char *mDomainName; // The serialized domain name (must NOT be nullptr) - const char *mServiceName; // The serialized service name (must NOT be nullptr for PTR/SRV/TXT queries) - const char *mInstanceName; // The serialized instance name or nullptr (only support one instance name) - const char *mHostName; // The serialized host name or nullptr (only support one host 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. + const char *const mDomainName; // The serialized domain name. + const char * mServiceName; // The serialized service name (only support one service name). + const char * mInstanceName; // The serialized instance name or nullptr (only support one instance name). + const char * mHostName; // The serialized host name or nullptr (only support one host 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. }; // This structure represents the splitting information of a full name. @@ -217,62 +232,60 @@ private: // instance. }; - 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(Message &aMessage, Message &aResponse, const Header &aRequestHeader); - static otError AddQuestionToResponse(const char * aName, - const Question &aQuestion, - Message & aResponse, - Header & aResponseHeader); - otError PrepareCompressInfo(uint16_t aQueryType, - const char * aName, - uint16_t aNameSerializeOffset, - Server::NameCompressInfo &aCompressInfo); - otError ResolveQuestion(const char * aName, - const Question & aQuestion, - Header & aResponseHeader, - Message & aResponseMessage, - NameCompressInfo &aCompressInfo); - static otError AppendPtrRecord(Message & aMessage, - const char * aServiceName, - const char * aInstanceName, - uint32_t aTtl, - NameCompressInfo &aCompressInfo); - static otError AppendSrvRecord(Message & aMessage, - const char * aInstanceName, - const char * aHostName, - uint32_t aTtl, - uint16_t aPriority, - uint16_t aWeight, - uint16_t aPort, - NameCompressInfo &aCompressInfo); - static otError AppendAaaaRecord(Message & aMessage, - const char * aHostName, - const Ip6::Address &aAddress, - uint32_t aTtl, - NameCompressInfo & aCompressInfo); - static otError AppendServiceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); - static otError AppendInstanceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); - static otError AppendHostName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); - static void IncResourceRecordCount(Header &aHeader, bool aAdditional); - static otError FindNameComponents(const char *aName, const char *aDomain, NameComponentsOffsetInfo &aInfo); - static otError FindPreviousLabel(const char *aName, uint8_t &aStart, uint8_t &aStop); + 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(Message &aMessage, Message &aResponse, const Header &aRequestHeader); + Header::Response ResolveQuestion(const char * aName, + const Question & aQuestion, + Header & aResponseHeader, + Message & aResponseMessage, + uint8_t aResolveKind, + NameCompressInfo &aCompressInfo); + static otError AppendQuestion(const char * aName, + const Question & aQuestion, + Message & aMessage, + NameCompressInfo &aCompressInfo); + static otError AppendPtrRecord(Message & aMessage, + const char * aServiceName, + const char * aInstanceName, + uint32_t aTtl, + NameCompressInfo &aCompressInfo); + static otError AppendSrvRecord(Message & aMessage, + const char * aInstanceName, + const char * aHostName, + uint32_t aTtl, + uint16_t aPriority, + uint16_t aWeight, + uint16_t aPort, + NameCompressInfo &aCompressInfo); + static otError AppendAaaaRecord(Message & aMessage, + const char * aHostName, + const Ip6::Address &aAddress, + uint32_t aTtl, + NameCompressInfo & aCompressInfo); + static otError AppendServiceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); + static otError AppendInstanceName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); + static otError AppendHostName(Message &aMessage, const char *aName, NameCompressInfo &aCompressInfo); + static void IncResourceRecordCount(Header &aHeader, bool aAdditional); + static otError FindNameComponents(const char *aName, const char *aDomain, NameComponentsOffsetInfo &aInfo); + static otError FindPreviousLabel(const char *aName, uint8_t &aStart, uint8_t &aStop); #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE - otError ResolveQuestionBySrp(const char * aName, - const Question & aQuestion, - Header & aResponseHeader, - Message & aResponseMessage, - bool aAdditional, - NameCompressInfo &aCompressInfo); - const Srp::Server::Host * GetNextSrpHost(const Srp::Server::Host *aHost); - const Srp::Server::Service *GetNextSrpService(const Srp::Server::Host &aHost, const Srp::Server::Service *aService); - static otError AppendTxtRecord(Message & aMessage, - const char * aInstanceName, - const Srp::Server::Service &aService, - uint32_t aTtl, - NameCompressInfo & aCompressInfo); + Header::Response ResolveQuestionBySrp(const char * aName, + const Question & aQuestion, + Header & aResponseHeader, + Message & aResponseMessage, + uint8_t aResolveKind, + NameCompressInfo &aCompressInfo); + const Srp::Server::Host * GetNextSrpHost(const Srp::Server::Host *aHost); + static const Srp::Server::Service *GetNextSrpService(const Srp::Server::Host & aHost, + const Srp::Server::Service *aService); + static otError AppendTxtRecord(Message & aMessage, + const char * aInstanceName, + const Srp::Server::Service &aService, + uint32_t aTtl, + NameCompressInfo & aCompressInfo); #endif static const char kDnssdProtocolUdp[4]; diff --git a/tests/scripts/thread-cert/node.py b/tests/scripts/thread-cert/node.py index 351034c04..9cec245db 100755 --- a/tests/scripts/thread-cert/node.py +++ b/tests/scripts/thread-cert/node.py @@ -2429,6 +2429,132 @@ class NodeImpl: return list(zip(ip, ttl)) + def dns_resolve_service(self, instance, service, server=None, port=53): + """ + Resolves the service instance and returns the instance information as a dict. + + Example return value: + { + 'port': 12345, + 'priority': 0, + 'weight': 0, + 'host': 'ins1._ipps._tcp.default.service.arpa.', + 'address': '2001::1', + 'txt_data': b'\x00', + 'srv_ttl': 7100, + 'txt_ttl': 7100, + 'aaaa_ttl': 7100, + } + """ + cmd = f'dns service {instance} {service}' + if server is not None: + cmd += f' {server} {port}' + + self.send_command(cmd) + self.simulator.go(10) + output = self._expect_command_output(cmd) + + # Example output: + # DNS service resolution response for ins2 for service _ipps._tcp.default.service.arpa. + # Port:22222, Priority:2, Weight:2, TTL:7155 + # Host:host2.default.service.arpa. + # HostAddress:0:0:0:0:0:0:0:0 TTL:0 + # TXT-Data:(len:1) [00] TTL:7155 + # Done + + m = re.match( + r'.*Port:(\d+), Priority:(\d+), Weight:(\d+), TTL:(\d+)\s+Host:(.*?)\s+HostAddress:(\S+) TTL:(\d+)\s+TXT-Data:\(len:\d+\) \[(.*?)\] TTL:(\d+)', + '\r'.join(output)) + if m: + port, priority, weight, srv_ttl, hostname, address, aaaa_ttl, txt_data, txt_ttl = m.groups() + return { + 'port': int(port), + 'priority': int(priority), + 'weight': int(weight), + 'host': hostname, + 'address': address, + 'txt_data': self.__parse_hex_string(txt_data), + 'srv_ttl': int(srv_ttl), + 'txt_ttl': int(txt_ttl), + 'aaaa_ttl': int(aaaa_ttl), + } + else: + raise Exception('dns resolve service failed: %s.%s' % (instance, service)) + + @staticmethod + def __parse_hex_string(hexstr: str) -> bytes: + assert (len(hexstr) % 2 == 0) + return bytes(int(hexstr[i:i + 2], 16) for i in range(0, len(hexstr), 2)) + + def dns_browse(self, service_name, server=None, port=53): + """ + Browse the service and returns the instances. + + Example return value: + { + 'ins1': { + 'port': 12345, + 'priority': 1, + 'weight': 1, + 'host': 'ins1._ipps._tcp.default.service.arpa.', + 'address': '2001::1', + 'txt_data': b'\x00', + 'srv_ttl': 7100, + 'txt_ttl': 7100, + 'aaaa_ttl': 7100, + }, + 'ins2': { + 'port': 12345, + 'priority': 2, + 'weight': 2, + 'host': 'ins2._ipps._tcp.default.service.arpa.', + 'address': '2001::2', + 'txt_data': b'\x00', + 'srv_ttl': 7100, + 'txt_ttl': 7100, + 'aaaa_ttl': 7100, + } + } + """ + cmd = f'dns browse {service_name}' + if server is not None: + cmd += f' {server} {port}' + + self.send_command(cmd) + self.simulator.go(10) + output = '\n'.join(self._expect_command_output(cmd)) + + # Example output: + # ins2 + # Port:22222, Priority:2, Weight:2, TTL:7175 + # Host:host2.default.service.arpa. + # HostAddress:fd00:db8:0:0:3205:28dd:5b87:6a63 TTL:7175 + # TXT-Data:(len:1) [00] TTL:7175 + # ins1 + # Port:11111, Priority:1, Weight:1, TTL:7170 + # Host:host1.default.service.arpa. + # HostAddress:fd00:db8:0:0:39f4:d9:eb4f:778 TTL:7170 + # TXT-Data:(len:1) [00] TTL:7170 + # Done + + result = {} + for ins, port, priority, weight, srv_ttl, hostname, address, aaaa_ttl, txt_data, txt_ttl in re.findall( + r'(.*?)\s+Port:(\d+), Priority:(\d+), Weight:(\d+), TTL:(\d+)\s*Host:(\S+)\s+HostAddress:(\S+) TTL:(\d+)\s+TXT-Data:\(len:\d+\) \[(.*?)\] TTL:(\d+)', + output): + result[ins] = { + 'port': int(port), + 'priority': int(priority), + 'weight': int(weight), + 'host': hostname, + 'address': address, + 'txt_data': self.__parse_hex_string(txt_data), + 'srv_ttl': int(srv_ttl), + 'txt_ttl': int(txt_ttl), + 'aaaa_ttl': int(aaaa_ttl), + } + + return result + class Node(NodeImpl, OtCli): pass diff --git a/tests/scripts/thread-cert/test_dnssd.py b/tests/scripts/thread-cert/test_dnssd.py index 9158fc0a9..e51f3c73b 100755 --- a/tests/scripts/thread-cert/test_dnssd.py +++ b/tests/scripts/thread-cert/test_dnssd.py @@ -27,6 +27,7 @@ # POSSIBILITY OF SUCH DAMAGE. # import ipaddress +import typing import unittest import thread_cert @@ -88,13 +89,68 @@ class TestDnssd(thread_cert.TestCase): self._config_srp_client_services(CLIENT2, 'ins2', 'host2', 22222, 2, 2, client2_addrs) # Test AAAA query using DNS client - response = self.nodes[CLIENT1].dns_resolve(f"host1.{DOMAIN}", self.nodes[SERVER].get_mleid(), 53) - self.assertIn(ipaddress.IPv6Address(response[0][0]), map(ipaddress.IPv6Address, client1_addrs)) + answers = self.nodes[CLIENT1].dns_resolve(f"host1.{DOMAIN}", self.nodes[SERVER].get_mleid(), 53) + self.assertEqual(set(ipaddress.IPv6Address(ip) for ip, _ in answers), + set(map(ipaddress.IPv6Address, client1_addrs))) - response = self.nodes[CLIENT1].dns_resolve(f"host2.{DOMAIN}", self.nodes[SERVER].get_mleid(), 53) - self.assertIn(ipaddress.IPv6Address(response[0][0]), map(ipaddress.IPv6Address, client2_addrs)) + answers = self.nodes[CLIENT1].dns_resolve(f"host2.{DOMAIN}", self.nodes[SERVER].get_mleid(), 53) + self.assertEqual(set(ipaddress.IPv6Address(ip) for ip, _ in answers), + set(map(ipaddress.IPv6Address, client2_addrs))) - # TODO: test other query types using DNS-SD client + service_instances = self.nodes[CLIENT1].dns_browse(f'{SERVICE}.{DOMAIN}', self.nodes[SERVER].get_mleid(), 53) + self.assertEqual({'ins1', 'ins2'}, set(service_instances.keys()), service_instances) + + instance1_verify_info = { + 'port': 11111, + 'priority': 1, + 'weight': 1, + 'host': 'host1.default.service.arpa.', + 'address': client1_addrs, + 'txt_data': b'\x00', + 'srv_ttl': lambda x: x > 0, + 'txt_ttl': lambda x: x > 0, + 'aaaa_ttl': lambda x: x > 0, + } + + instance2_verify_info = { + 'port': 22222, + 'priority': 2, + 'weight': 2, + 'host': 'host2.default.service.arpa.', + 'address': client2_addrs, + 'txt_data': b'\x00', + 'srv_ttl': lambda x: x > 0, + 'txt_ttl': lambda x: x > 0, + 'aaaa_ttl': lambda x: x > 0, + } + + self._assert_service_instance_equal(service_instances['ins1'], instance1_verify_info) + self._assert_service_instance_equal(service_instances['ins2'], instance2_verify_info) + + service_instance = self.nodes[CLIENT1].dns_resolve_service('ins1', f'{SERVICE}.{DOMAIN}', + self.nodes[SERVER].get_mleid(), 53) + self._assert_service_instance_equal(service_instance, instance1_verify_info) + + service_instance = self.nodes[CLIENT1].dns_resolve_service('ins2', f'{SERVICE}.{DOMAIN}', + self.nodes[SERVER].get_mleid(), 53) + self._assert_service_instance_equal(service_instance, instance2_verify_info) + + def _assert_service_instance_equal(self, instance, info): + for f in ('port', 'priority', 'weight', 'host', 'txt_data'): + self.assertEqual(instance[f], info[f], instance) + + verify_addresses = info['address'] + if not isinstance(verify_addresses, typing.Collection): + verify_addresses = [verify_addresses] + self.assertIn(ipaddress.IPv6Address(instance['address']), map(ipaddress.IPv6Address, verify_addresses), + instance) + + for ttl_f in ('srv_ttl', 'txt_ttl', 'aaaa_ttl'): + check_ttl = info[ttl_f] + if not callable(check_ttl): + check_ttl = lambda x: x == check_ttl + + self.assertTrue(check_ttl(instance[ttl_f]), instance) def _config_srp_client_services(self, client, instancename, hostname, port, priority, weight, addrs): self.nodes[client].netdata_show()