From a3a76f1ea3318060195aff762468d0af98a36608 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=C5=81ukasz=20Duda?= Date: Mon, 21 Dec 2020 16:29:11 +0100 Subject: [PATCH] [dns-client] randomize ID of outgoing query (#5981) --- src/core/net/dns_client.cpp | 27 +++++++++++++++++++++------ src/core/net/dns_client.hpp | 5 +++-- 2 files changed, 24 insertions(+), 8 deletions(-) diff --git a/src/core/net/dns_client.cpp b/src/core/net/dns_client.cpp index e922b3236..c59fc9b01 100644 --- a/src/core/net/dns_client.cpp +++ b/src/core/net/dns_client.cpp @@ -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(&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); diff --git a/src/core/net/dns_client.hpp b/src/core/net/dns_client.hpp index 4a491fe48..656f5c4f9 100644 --- a/src/core/net/dns_client.hpp +++ b/src/core/net/dns_client.hpp @@ -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; };