[dnssd-server] simplify resolving of query by proxy (#9353)

This commit enhances the DNSSD Server by updating how it stores and
tracks information when trying to resolve a received query using a
discovery proxy.

The previous approach used a fixed-size array to retain information
for pending proxy queries. The new approach uses the already
allocated `Response.mMessage` and appends the newly defined
`ProxyQueryInfo` struct to it. The `ProxyQueryInfo` struct specifies
all the information we need to track associated with a query.
`ProxyQuery` message entries are placed in a message queue waiting
for a callback from the proxy. Once a callback is received, the
`ProxyQuery` message is converted back to a `Response.mMessage` and
response is prepared and sent. The new approach removes the
restriction on the number of outstanding proxy queries and simplifies
the code.

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