[message] add helper methods to get first/next data chunks (#5415)

The `Message` data is stored in a sequence of linked-listed `Buffer`
instances. When reading or writing data in the `Message`, the internal
implementation iterates over the `Buffer` list and updates the data in
chunks. This commit adds two helper methods `GetFirstChunk()` and
`GetNextChunk()` to get first/next data chunk (contiguous buffer)
corresponding to a given offset and length into the message. These
helper methods are then used to simplify `Read()`, `Write()`, and
`UpdateChecksum()` methods. The unit test `test_message` is also
updated to verify reading and writing data at different offsets and
lengths.
This commit is contained in:
Abtin Keshavarzian
2020-08-19 14:06:36 -07:00
committed by GitHub
parent 8fd46dde18
commit 6886ff1e4c
3 changed files with 167 additions and 185 deletions
+87 -168
View File
@@ -429,16 +429,14 @@ void Message::RemoveHeader(uint16_t aLength)
}
}
uint16_t Message::Read(uint16_t aOffset, uint16_t aLength, void *aBuf) const
void Message::GetFirstChunk(uint16_t aOffset, uint16_t &aLength, Chunk &aChunk) const
{
const Buffer *curBuffer;
uint16_t bytesCopied = 0;
uint16_t bytesToCopy;
if (aOffset >= GetLength())
{
ExitNow();
}
// This method gets the first message chunk (contiguous data
// buffer) corresponding to a given offset and length. On exit
// `aChunk` is updated such that `aChunk.GetData()` gives the
// pointer to the start of chunk and `aChunk.GetLength()` gives
// its length. The `aLength` is also decreased by the chunk
// length.
if (aOffset + aLength >= GetLength())
{
@@ -447,138 +445,109 @@ uint16_t Message::Read(uint16_t aOffset, uint16_t aLength, void *aBuf) const
aOffset += GetReserved();
// special case first buffer
aChunk.mBuffer = this;
// Special case for the first buffer
if (aOffset < kHeadBufferDataSize)
{
bytesToCopy = kHeadBufferDataSize - aOffset;
aChunk.mData = GetFirstData() + aOffset;
aChunk.mLength = kHeadBufferDataSize - aOffset;
ExitNow();
}
if (bytesToCopy > aLength)
aOffset -= kHeadBufferDataSize;
// Find the `Buffer` matching the offset
while (true)
{
aChunk.mBuffer = aChunk.mBuffer->GetNextBuffer();
OT_ASSERT(aChunk.mBuffer != nullptr);
if (aOffset < kBufferDataSize)
{
bytesToCopy = aLength;
aChunk.mData = aChunk.mBuffer->GetData() + aOffset;
aChunk.mLength = kBufferDataSize - aOffset;
ExitNow();
}
memcpy(aBuf, GetFirstData() + aOffset, bytesToCopy);
aLength -= bytesToCopy;
bytesCopied += bytesToCopy;
aBuf = static_cast<uint8_t *>(aBuf) + bytesToCopy;
aOffset = 0;
}
else
{
aOffset -= kHeadBufferDataSize;
}
// advance to offset
curBuffer = GetNextBuffer();
while (aOffset >= kBufferDataSize)
{
OT_ASSERT(curBuffer != nullptr);
curBuffer = curBuffer->GetNextBuffer();
aOffset -= kBufferDataSize;
}
// begin copy
while (aLength > 0)
exit:
if (aChunk.mLength > aLength)
{
OT_ASSERT(curBuffer != nullptr);
aChunk.mLength = aLength;
}
bytesToCopy = kBufferDataSize - aOffset;
aLength -= aChunk.mLength;
}
if (bytesToCopy > aLength)
{
bytesToCopy = aLength;
}
void Message::GetNextChunk(uint16_t &aLength, Chunk &aChunk) const
{
// This method gets the next message chunk. On input, the
// `aChunk` should be the previous chunk. On exit, it is
// updated to provide info about next chunk, and `aLength`
// is decreased by the chunk length. If there is no more
// chunk, `aChunk.GetLength()` would be zero.
memcpy(aBuf, curBuffer->GetData() + aOffset, bytesToCopy);
VerifyOrExit(aLength > 0, aChunk.mLength = 0);
aLength -= bytesToCopy;
bytesCopied += bytesToCopy;
aBuf = static_cast<uint8_t *>(aBuf) + bytesToCopy;
aChunk.mBuffer = aChunk.mBuffer->GetNextBuffer();
OT_ASSERT(aChunk.mBuffer != nullptr);
curBuffer = curBuffer->GetNextBuffer();
aOffset = 0;
aChunk.mData = aChunk.mBuffer->GetData();
aChunk.mLength = kBufferDataSize;
if (aChunk.mLength > aLength)
{
aChunk.mLength = aLength;
}
aLength -= aChunk.mLength;
exit:
return;
}
uint16_t Message::Read(uint16_t aOffset, uint16_t aLength, void *aBuf) const
{
uint8_t *bufPtr = reinterpret_cast<uint8_t *>(aBuf);
Chunk chunk;
VerifyOrExit(aOffset < GetLength(), OT_NOOP);
GetFirstChunk(aOffset, aLength, chunk);
while (chunk.GetLength() > 0)
{
memcpy(bufPtr, chunk.GetData(), chunk.GetLength());
bufPtr += chunk.GetLength();
GetNextChunk(aLength, chunk);
}
exit:
return bytesCopied;
return static_cast<uint16_t>(bufPtr - reinterpret_cast<uint8_t *>(aBuf));
}
int Message::Write(uint16_t aOffset, uint16_t aLength, const void *aBuf)
{
Buffer * curBuffer;
uint16_t bytesCopied = 0;
uint16_t bytesToCopy;
const uint8_t *bufPtr = reinterpret_cast<const uint8_t *>(aBuf);
WritableChunk chunk;
OT_ASSERT(aOffset + aLength <= GetLength());
if (aOffset + aLength >= GetLength())
GetFirstChunk(aOffset, aLength, chunk);
while (chunk.GetLength() > 0)
{
aLength = GetLength() - aOffset;
memcpy(chunk.GetData(), bufPtr, chunk.GetLength());
bufPtr += chunk.GetLength();
GetNextChunk(aLength, chunk);
}
aOffset += GetReserved();
// special case first buffer
if (aOffset < kHeadBufferDataSize)
{
bytesToCopy = kHeadBufferDataSize - aOffset;
if (bytesToCopy > aLength)
{
bytesToCopy = aLength;
}
memcpy(GetFirstData() + aOffset, aBuf, bytesToCopy);
aLength -= bytesToCopy;
bytesCopied += bytesToCopy;
aBuf = static_cast<const uint8_t *>(aBuf) + bytesToCopy;
aOffset = 0;
}
else
{
aOffset -= kHeadBufferDataSize;
}
// advance to offset
curBuffer = GetNextBuffer();
while (aOffset >= kBufferDataSize)
{
OT_ASSERT(curBuffer != nullptr);
curBuffer = curBuffer->GetNextBuffer();
aOffset -= kBufferDataSize;
}
// begin copy
while (aLength > 0)
{
OT_ASSERT(curBuffer != nullptr);
bytesToCopy = kBufferDataSize - aOffset;
if (bytesToCopy > aLength)
{
bytesToCopy = aLength;
}
memcpy(curBuffer->GetData() + aOffset, aBuf, bytesToCopy);
aLength -= bytesToCopy;
bytesCopied += bytesToCopy;
aBuf = static_cast<const uint8_t *>(aBuf) + bytesToCopy;
curBuffer = curBuffer->GetNextBuffer();
aOffset = 0;
}
return bytesCopied;
return static_cast<int>(bufPtr - reinterpret_cast<const uint8_t *>(aBuf));
}
int Message::CopyTo(uint16_t aSourceOffset, uint16_t aDestinationOffset, uint16_t aLength, Message &aMessage) const
@@ -686,66 +655,16 @@ uint16_t Message::UpdateChecksum(uint16_t aChecksum, const void *aBuf, uint16_t
uint16_t Message::UpdateChecksum(uint16_t aChecksum, uint16_t aOffset, uint16_t aLength) const
{
const Buffer *curBuffer;
uint16_t bytesCovered = 0;
uint16_t bytesToCover;
Chunk chunk;
OT_ASSERT(aOffset + aLength <= GetLength());
aOffset += GetReserved();
GetFirstChunk(aOffset, aLength, chunk);
// special case first buffer
if (aOffset < kHeadBufferDataSize)
while (chunk.GetLength() > 0)
{
bytesToCover = kHeadBufferDataSize - aOffset;
if (bytesToCover > aLength)
{
bytesToCover = aLength;
}
aChecksum = Message::UpdateChecksum(aChecksum, GetFirstData() + aOffset, bytesToCover);
aLength -= bytesToCover;
bytesCovered += bytesToCover;
aOffset = 0;
}
else
{
aOffset -= kHeadBufferDataSize;
}
// advance to offset
curBuffer = GetNextBuffer();
while (aOffset >= kBufferDataSize)
{
OT_ASSERT(curBuffer != nullptr);
curBuffer = curBuffer->GetNextBuffer();
aOffset -= kBufferDataSize;
}
// begin copy
while (aLength > 0)
{
OT_ASSERT(curBuffer != nullptr);
bytesToCover = kBufferDataSize - aOffset;
if (bytesToCover > aLength)
{
bytesToCover = aLength;
}
aChecksum = Message::UpdateChecksum(aChecksum, curBuffer->GetData() + aOffset, bytesToCover);
aLength -= bytesToCover;
bytesCovered += bytesToCover;
curBuffer = curBuffer->GetNextBuffer();
aOffset = 0;
aChecksum = Message::UpdateChecksum(aChecksum, chunk.GetData(), chunk.GetLength());
GetNextChunk(aLength, chunk);
}
return aChecksum;
+30
View File
@@ -361,6 +361,7 @@ public:
* This method returns the number of bytes in the message.
*
* @returns The number of bytes in the message.
*
*/
uint16_t GetLength(void) const { return GetMetadata().mLength; }
@@ -1010,6 +1011,35 @@ private:
*
*/
otError ResizeMessage(uint16_t aLength);
private:
struct Chunk
{
const uint8_t *GetData(void) const { return mData; }
uint16_t GetLength(void) const { return mLength; }
const uint8_t *mData; // Pointer to start of chunk data buffer.
uint16_t mLength; // Length of chunk data (in bytes).
const Buffer * mBuffer; // Buffer containing the chunk
};
struct WritableChunk : public Chunk
{
uint8_t *GetData(void) const { return const_cast<uint8_t *>(mData); }
};
void GetFirstChunk(uint16_t aOffset, uint16_t &aLength, Chunk &chunk) const;
void GetNextChunk(uint16_t &aLength, Chunk &aChunk) const;
void GetFirstChunk(uint16_t aOffset, uint16_t &aLength, WritableChunk &aChunk)
{
const_cast<const Message *>(this)->GetFirstChunk(aOffset, aLength, static_cast<Chunk &>(aChunk));
}
void GetNextChunk(uint16_t &aLength, WritableChunk &aChunk)
{
const_cast<const Message *>(this)->GetNextChunk(aLength, static_cast<Chunk &>(aChunk));
}
};
/**
+50 -17
View File
@@ -29,42 +29,75 @@
#include "common/debug.hpp"
#include "common/instance.hpp"
#include "common/message.hpp"
#include "common/random.hpp"
#include "test_platform.h"
#include "test_util.h"
#include "test_util.hpp"
namespace ot {
void TestMessage(void)
{
ot::Instance * instance;
ot::MessagePool *messagePool;
ot::Message * message;
uint8_t writeBuffer[1024];
uint8_t readBuffer[1024];
enum : uint16_t
{
kMaxSize = (kBufferSize * 3 + 24),
};
instance = static_cast<ot::Instance *>(testInitInstance());
Instance * instance;
MessagePool *messagePool;
Message * message;
uint8_t writeBuffer[kMaxSize];
uint8_t readBuffer[kMaxSize];
uint8_t zeroBuffer[kMaxSize];
memset(zeroBuffer, 0, sizeof(zeroBuffer));
instance = static_cast<Instance *>(testInitInstance());
VerifyOrQuit(instance != nullptr, "Null OpenThread instance\n");
messagePool = &instance->Get<ot::MessagePool>();
messagePool = &instance->Get<MessagePool>();
for (uint8_t &b : writeBuffer)
Random::NonCrypto::FillBuffer(writeBuffer, kMaxSize);
VerifyOrQuit((message = messagePool->New(Message::kTypeIp6, 0)) != nullptr, "Message::New failed");
SuccessOrQuit(message->SetLength(kMaxSize), "Message::SetLength failed");
VerifyOrQuit(message->Write(0, kMaxSize, writeBuffer) == kMaxSize, "Message::Write failed");
VerifyOrQuit(message->Read(0, kMaxSize, readBuffer) == kMaxSize, "Message::Read failed");
VerifyOrQuit(memcmp(writeBuffer, readBuffer, kMaxSize) == 0, "Message compare failed");
VerifyOrQuit(message->GetLength() == kMaxSize, "Message::GetLength failed");
for (uint16_t offset = 0; offset < kMaxSize; offset++)
{
b = static_cast<uint8_t>(random());
for (uint16_t length = 0; length < kMaxSize - offset; length++)
{
for (uint16_t i = 0; i < length; i++)
{
writeBuffer[offset + i]++;
}
VerifyOrQuit(message->Write(offset, length, &writeBuffer[offset]) == length, "Message::Write failed");
VerifyOrQuit(message->Read(0, kMaxSize, readBuffer) == kMaxSize, "Message::Read failed");
VerifyOrQuit(memcmp(writeBuffer, readBuffer, kMaxSize) == 0, "Message compare failed");
memset(readBuffer, 0, sizeof(readBuffer));
VerifyOrQuit(message->Read(offset, length, readBuffer) == length, "Message::Read failed");
VerifyOrQuit(memcmp(readBuffer, &writeBuffer[offset], length) == 0, "Message compare failed");
VerifyOrQuit(memcmp(&readBuffer[length], zeroBuffer, kMaxSize - length) == 0, "Message read after length");
}
}
VerifyOrQuit((message = messagePool->New(ot::Message::kTypeIp6, 0)) != nullptr, "Message::New failed");
SuccessOrQuit(message->SetLength(sizeof(writeBuffer)), "Message::SetLength failed");
VerifyOrQuit(message->Write(0, sizeof(writeBuffer), writeBuffer) == sizeof(writeBuffer), "Message::Write failed");
VerifyOrQuit(message->Read(0, sizeof(readBuffer), readBuffer) == sizeof(readBuffer), "Message::Read failed");
VerifyOrQuit(memcmp(writeBuffer, readBuffer, sizeof(writeBuffer)) == 0, "Message compare failed");
VerifyOrQuit(message->GetLength() == 1024, "Message::GetLength failed");
VerifyOrQuit(message->GetLength() == kMaxSize, "Message::GetLength failed");
message->Free();
testFreeInstance(instance);
}
} // namespace ot
int main(void)
{
TestMessage();
ot::TestMessage();
printf("All tests passed\n");
return 0;
}