[dns-server] handle multiple questions (#6155)

This commit enhances the DNS-SD server to handle multiple questions.

Specifically for DNS service resolving query with SRV+TXT questions,
this commit would reply with SRV+TXT+AAAA RRs, so that the service
instance can be fully resolved with one query.

More implementation details:
- This commit enhances NameCompressInfo so that question names can be
  compressed together with RR names.
  - Full compression is limited to only one instance / service / host
    name.
- For any RR type in the question, this commit will exclude it in
  additional RRs. For example, if the query contains a AAAA question,
  the response won't include any additional AAAA RR, but will only
  have AAAA answer(s). The purpose is to avoid duplicate RRs in answer
  section and additional section.
This commit is contained in:
Simon Lin
2021-02-19 09:38:25 -08:00
committed by GitHub
parent 5756b9db66
commit c261cf2dff
4 changed files with 503 additions and 243 deletions
+222 -157
View File
@@ -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<Srp::Server>().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<uint8_t>(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<uint8_t>(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<uint8_t>(StringLength(aName, Name::kMaxNameLength) -
StringLength(aCompressInfo.GetDomainName(), Name::kMaxNameLength));
uint8_t domainStart = static_cast<uint8_t>(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)
+94 -81
View File
@@ -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<NameCompressInfo>
{
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];
+126
View File
@@ -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
+61 -5
View File
@@ -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()