[dns-client] randomize ID of outgoing query (#5981)

This commit is contained in:
Łukasz Duda
2020-12-21 07:29:11 -08:00
committed by GitHub
parent 42e1da5573
commit a3a76f1ea3
2 changed files with 24 additions and 8 deletions
+21 -6
View File
@@ -50,7 +50,6 @@ namespace Dns {
Client::Client(Instance &aInstance)
: mSocket(aInstance)
, mMessageId(0)
, mRetransmissionTimer(aInstance, Client::HandleRetransmissionTimer, this)
{
}
@@ -89,10 +88,13 @@ otError Client::Query(const QueryInfo &aQuery, ResponseHandler aHandler, void *a
Message * messageCopy = nullptr;
Header header;
QuestionAaaa question;
uint16_t messageId;
VerifyOrExit(aQuery.IsValid(), error = OT_ERROR_INVALID_ARGS);
header.SetMessageId(mMessageId++);
SuccessOrExit(error = GenerateUniqueRandomId(messageId));
header.SetMessageId(messageId);
header.SetType(Header::kTypeQuery);
header.SetQueryType(Header::kQueryTypeStandard);
@@ -199,6 +201,19 @@ exit:
}
}
otError Client::GenerateUniqueRandomId(uint16_t &aRandomId)
{
otError error;
do
{
SuccessOrExit(error = Random::Crypto::FillBuffer(reinterpret_cast<uint8_t *>(&aRandomId), sizeof(aRandomId)));
} while (FindQueryById(aRandomId) != nullptr);
exit:
return error;
}
otError Client::CompareQuestions(Message &aMessageResponse, Message &aMessageQuery, uint16_t &aOffset)
{
otError error = OT_ERROR_NONE;
@@ -228,7 +243,7 @@ exit:
return error;
}
Message *Client::FindRelatedQuery(const Header &aResponseHeader, QueryMetadata &aQueryMetadata)
Message *Client::FindQueryById(uint16_t aMessageId)
{
uint16_t messageId;
Message *message;
@@ -241,9 +256,8 @@ Message *Client::FindRelatedQuery(const Header &aResponseHeader, QueryMetadata &
OT_ASSERT(false);
}
if (HostSwap16(messageId) == aResponseHeader.GetMessageId())
if (HostSwap16(messageId) == aMessageId)
{
aQueryMetadata.ReadFrom(*message);
break;
}
}
@@ -346,7 +360,8 @@ void Client::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessag
aMessage.MoveOffset(sizeof(responseHeader));
offset = aMessage.GetOffset();
VerifyOrExit((message = FindRelatedQuery(responseHeader, queryMetadata)) != nullptr);
VerifyOrExit((message = FindQueryById(responseHeader.GetMessageId())) != nullptr);
queryMetadata.ReadFrom(*message);
VerifyOrExit(responseHeader.GetResponseCode() == Header::kResponseSuccess, error = OT_ERROR_FAILED);
+3 -2
View File
@@ -181,9 +181,11 @@ private:
otError SendMessage(Message &aMessage, const Ip6::MessageInfo &aMessageInfo);
void SendCopy(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo);
otError GenerateUniqueRandomId(uint16_t &aRandomId);
otError CompareQuestions(Message &aMessageResponse, Message &aMessageQuery, uint16_t &aOffset);
Message *FindRelatedQuery(const Header &aResponseHeader, QueryMetadata &aQueryMetadata);
Message *FindQueryById(uint16_t aMessageId);
void FinalizeDnsTransaction(Message & aQuery,
const QueryMetadata &aQueryMetadata,
const Ip6::Address * aAddress,
@@ -198,7 +200,6 @@ private:
Ip6::Udp::Socket mSocket;
uint16_t mMessageId;
MessageQueue mPendingQueries;
TimerMilli mRetransmissionTimer;
};