[dns-name] allow 'Dns::Name' to be given as string or from message (#6117)

This commit allows `Dns::Name` instances to be created which can be
empty or given as a string (a dot '.' separated sequence of labels) or
from a message at a given offset (i.e., name is already encoded in a
message). `CompareName` and `FindRecord()` helper methods are updated
to allow the new `Dns::Name` type to be used when comparing/ searching
for a name.
This commit is contained in:
Abtin Keshavarzian
2021-01-28 14:27:17 -08:00
committed by GitHub
parent 2a3f13e1c7
commit c03458433e
3 changed files with 193 additions and 15 deletions
+9 -1
View File
@@ -404,6 +404,14 @@ exit:
return error;
}
otError Name::CompareName(const Message &aMessage, uint16_t &aOffset, const Name &aName)
{
return aName.IsFromCString()
? CompareName(aMessage, aOffset, aName.mString)
: (aName.IsFromMessage() ? CompareName(aMessage, aOffset, *aName.mMessage, aName.mOffset)
: ParseName(aMessage, aOffset));
}
otError Name::LabelIterator::GetNextLabel(void)
{
otError error;
@@ -549,7 +557,7 @@ exit:
return error;
}
otError ResourceRecord::FindRecord(const Message &aMessage, uint16_t &aOffset, uint16_t &aNumRecords, const char *aName)
otError ResourceRecord::FindRecord(const Message &aMessage, uint16_t &aOffset, uint16_t &aNumRecords, const Name &aName)
{
otError error;
+168 -6
View File
@@ -480,10 +480,10 @@ public:
} OT_TOOL_PACKED_END;
/**
* This class implement helper methods for encoding/decoding of DNS Names.
* This class represents a DNS name and implements helper methods for encoding/decoding of DNS Names.
*
*/
class Name
class Name : public Clearable<Name>
{
public:
enum : uint8_t
@@ -493,6 +493,139 @@ public:
kMaxEncodedLength = 255, ///< Max length of an encoded name.
};
/**
* This enumeration represents the name type.
*
*/
enum Type : uint8_t
{
kTypeEmpty, ///< The name is empty (not specified).
kTypeCString, ///< The name is given as a C string (dot '.' separated sequence of labels).
kTypeMessage, ///< The name is specified from a message at a given offset (encoded in the message).
};
/**
* This constructor initializes the `Name` object as empty (not specified).
*
*/
Name(void)
: Name(nullptr, nullptr, 0)
{
}
/**
* This constructor initializes the `Name` object with a given string.
*
* @param[in] aString A C string specifying the name (dot '.' separated sequence of labels').
*
*/
explicit Name(const char *aString)
: Name(aString, nullptr, 0)
{
}
/**
* This constructor initializes the `Name` object from a message at a given offset.
*
* @param[in] aMessage The message containing the encoded name. `aMessage.GetOffset()` MUST point to the start of
* the DNS header in the message (used to parse compressed name).
* @param[in] aOffset The offset in @p aMessage pointing to the start of the name.
*
*/
Name(const Message &aMessage, uint16_t aOffset)
: Name(nullptr, &aMessage, aOffset)
{
}
/**
* This method indicates whether the name is empty (not specified).
*
* @returns TRUE if the name is empty, FALSE otherwise.
*
*/
bool IsEmpty(void) const { return (mString == nullptr) && (mMessage == nullptr); }
/**
* This method indicates whether the name is specified from a C string.
*
* @returns TRUE if the name is specified from a string, FALSE otherwise.
*
*/
bool IsFromCString(void) const { return mString != nullptr; }
/**
* This method indicates whether the name is specified from a message.
*
* @returns TRUE if the name is specified from a message, FALSE otherwise.
*
*/
bool IsFromMessage(void) const { return mMessage != nullptr; }
/**
* This method gets the type of `Name` object indicating whether it is empty, specified by a C string or from a
* message
*
* @returns The name type.
*
*/
Type GetFromType(void) const
{
return IsFromCString() ? kTypeCString : (IsFromMessage() ? kTypeMessage : kTypeEmpty);
}
/**
* This method sets the name from a given C string.
*
* @param[in] aString A C string specifying the name (dot '.' separated sequence of labels).
*
*/
void Set(const char *aString)
{
mString = aString;
mMessage = nullptr;
}
/**
* This method sets the name from a message at a given offset.
*
* @param[in] aMessage The message containing the encoded name. `aMessage.GetOffset()` MUST point to the start of
* the DNS header in the message (used to parse compressed name).
* @param[in] aOffset The offset in @p aMessage pointing to the start of the name.
*
*/
void SetFromMessage(const Message &aMessage, uint16_t aOffset)
{
mString = nullptr;
mMessage = &aMessage;
mOffset = aOffset;
}
/**
* This method gets the name as a C string.
*
* This method MUST be used only when the type is `kTypeString`. Otherwise its behavior is undefined.
*
* @returns A pointer to the C string.
*
*/
const char *GetAsCString(void) const { return mString; }
/**
* This method gets the name message and offset.
*
* This method MUST be used only when the type is `kTypeMessage`. Otherwise its behavior is undefined.
*
* @param[out] aOffset A reference to a variable to output the offset of the start of the name in the message.
*
* @returns A reference to the message containing the name.
*
*/
const Message &GetAsMessage(uint16_t &aOffset) const
{
aOffset = mOffset;
return *mMessage;
}
/**
* This static method encodes and appends a single name label to a message.
*
@@ -735,6 +868,26 @@ public:
*/
static otError CompareName(const Message &aMessage, uint16_t &aOffset, const Message &aMessage2, uint16_t aOffset2);
/**
* This static method parses and compares a full name from a message with a given name.
*
* If @p aName is empty (not specified), then any name in @p aMessage is considered a match to it.
*
* @param[in] aMessage The message to read the name from and compare. `aMessage.GetOffset()` MUST point
* to the start of DNS header (this is used to handle compressed names).
* @param[inout] aOffset On input, the offset in @p aMessage pointing to the start of the name field.
* On exit (when parsed successfully independent of whether the read name matches
* or not), @p aOffset is updated to point to the byte after the end of the name
* field.
* @param[in] aName A reference to a name to compare with.
*
* @retval OT_ERROR_NONE The name from @p aMessage matches @p aName. @p aOffset is updated.
* @retval OT_ERROR_NOT_FOUND The name from @p aMessage does not match @p aName. @p aOffset is updated.
* @retval OT_ERROR_PARSE Name in @p aMessage could not be parsed (invalid format).
*
*/
static otError CompareName(const Message &aMessage, uint16_t &aOffset, const Name &aName);
private:
enum : char
{
@@ -788,9 +941,18 @@ private:
uint16_t mNameEndOffset; // Offset in `mMessage` to the byte after the end of domain name field.
};
Name(void) = default;
static otError AppendLabel(const char *aLabel, uint8_t aLabelLength, Message &aMessage);
Name(const char *aString, const Message *aMessage, uint16_t aOffset)
: mString(aString)
, mMessage(aMessage)
, mOffset(aOffset)
{
}
const char * mString; // String containing the name or `nullptr` if name is not from string.
const Message *mMessage; // Message containing the encoded name, or `nullptr` if `Name` is not from message.
uint16_t mOffset; // Offset in `mMessage` to the start of name (used when name is from `mMessage`).
};
/**
@@ -964,14 +1126,14 @@ public:
* @param[inout] aNumRecords On input, the maximum number of records to check (starting from @p aOffset).
* On exit and if a matching record is found, @p aNumRecords is updated to give the
* number of remaining records after @p aOffset (excluding the matching record).
* @param[in] aName The record name to match against (MUST be a null terminated string).
* @param[in] aName The record name to match against.
*
* @retval OT_ERROR_NONE A matching record was found. @p aOffset, @p aNumRecords are updated.
* @retval OT_ERROR_NOT_FOUND A matching record could not be found. @p aOffset and @p aNumRecords are updated.
* @retval OT_ERROR_PARSE Could not parse records from @p aMessage (e.g., ran out of bytes in @p aMessage).
*
*/
static otError FindRecord(const Message &aMessage, uint16_t &aOffset, uint16_t &aNumRecords, const char *aName);
static otError FindRecord(const Message &aMessage, uint16_t &aOffset, uint16_t &aNumRecords, const Name &aName);
/**
* This template static method tries to read a resource record of a given type from a message. If the record type
+16 -8
View File
@@ -862,7 +862,8 @@ void TestHeaderAndResourceRecords(void)
{
uint16_t prevNumRecords = numRecords;
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, kServiceName), "FindRecord failed");
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(kServiceName)),
"FindRecord failed");
VerifyOrQuit(numRecords == prevNumRecords - 1, "Incorrect num records");
SuccessOrQuit(Dns::ResourceRecord::ReadRecord(*message, offset, ptrRecord), "ReadRecord() failed");
VerifyOrQuit(ptrRecord.GetTtl() == kTtl, "Read PTR is incorrect");
@@ -873,7 +874,8 @@ void TestHeaderAndResourceRecords(void)
}
VerifyOrQuit(offset == additionalSectionOffset, "offset is incorrect after answer section parse");
VerifyOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, kServiceName) == OT_ERROR_NOT_FOUND,
VerifyOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(kServiceName)) ==
OT_ERROR_NOT_FOUND,
"FindRecord did not fail with no records");
// Use `ReadRecord()` with a non-matching record type. Verify that it correct skips over the record.
@@ -883,7 +885,8 @@ void TestHeaderAndResourceRecords(void)
while (numRecords > 0)
{
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, kServiceName), "FindRecord failed");
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(kServiceName)),
"FindRecord failed");
VerifyOrQuit(Dns::ResourceRecord::ReadRecord(*message, offset, srvRecord) == OT_ERROR_NOT_FOUND,
"ReadRecord() did not fail with non-matching type");
}
@@ -894,7 +897,8 @@ void TestHeaderAndResourceRecords(void)
offset = answerSectionOffset;
numRecords = kAnswerCount;
VerifyOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, kInstance1Name) == OT_ERROR_NOT_FOUND,
VerifyOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(kInstance1Name)) ==
OT_ERROR_NOT_FOUND,
"FindRecord did not fail with non-matching name");
VerifyOrQuit(numRecords == 0, "Incorrect num records");
VerifyOrQuit(offset == additionalSectionOffset, "offset is incorrect after answer section parse");
@@ -950,18 +954,21 @@ void TestHeaderAndResourceRecords(void)
offset = additionalSectionOffset;
numRecords = kAdditionalCount;
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, instanceName), "FindRecord failed");
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(instanceName)),
"FindRecord failed");
SuccessOrQuit(Dns::ResourceRecord::ReadRecord(*message, offset, srvRecord), "ReadRecord() failed");
SuccessOrQuit(Dns::Name::ParseName(*message, offset), "ParseName() failed");
printf(" \"%s\" SRV %u %d %d %d %d\n", instanceName, srvRecord.GetTtl(), srvRecord.GetLength(),
srvRecord.GetPort(), srvRecord.GetWeight(), srvRecord.GetPriority());
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, instanceName), "FindRecord failed");
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(instanceName)),
"FindRecord failed");
SuccessOrQuit(Dns::ResourceRecord::ReadRecord(*message, offset, txtRecord), "ReadRecord() failed");
offset += txtRecord.GetLength();
printf(" \"%s\" TXT %u %d\n", instanceName, txtRecord.GetTtl(), txtRecord.GetLength());
VerifyOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, instanceName) == OT_ERROR_NOT_FOUND,
VerifyOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(instanceName)) ==
OT_ERROR_NOT_FOUND,
"FindRecord() did not fail with no more records");
VerifyOrQuit(offset == message->GetLength(), "offset is incorrect after additional section parse");
@@ -969,7 +976,8 @@ void TestHeaderAndResourceRecords(void)
offset = additionalSectionOffset;
numRecords = kAdditionalCount;
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, kHostName), "FindRecord() failed");
SuccessOrQuit(Dns::ResourceRecord::FindRecord(*message, offset, numRecords, Dns::Name(kHostName)),
"FindRecord() failed");
SuccessOrQuit(Dns::ResourceRecord::ReadRecord(*message, offset, record), "ReadRecord() failed");
VerifyOrQuit(record.GetType() == Dns::ResourceRecord::kTypeAaaa, "Read record has incorrect type");
offset += record.GetLength();