diff --git a/src/core/common/message.cpp b/src/core/common/message.cpp index fd37cfa9d..cc1ca0c93 100644 --- a/src/core/common/message.cpp +++ b/src/core/common/message.cpp @@ -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(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(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(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(bufPtr - reinterpret_cast(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(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(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(aBuf) + bytesToCopy; - - curBuffer = curBuffer->GetNextBuffer(); - aOffset = 0; - } - - return bytesCopied; + return static_cast(bufPtr - reinterpret_cast(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; diff --git a/src/core/common/message.hpp b/src/core/common/message.hpp index 6480e6426..cc01d55ac 100644 --- a/src/core/common/message.hpp +++ b/src/core/common/message.hpp @@ -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(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(this)->GetFirstChunk(aOffset, aLength, static_cast(aChunk)); + } + + void GetNextChunk(uint16_t &aLength, WritableChunk &aChunk) + { + const_cast(this)->GetNextChunk(aLength, static_cast(aChunk)); + } }; /** diff --git a/tests/unit/test_message.cpp b/tests/unit/test_message.cpp index 0aadabc7a..6f1910ce5 100644 --- a/tests/unit/test_message.cpp +++ b/tests/unit/test_message.cpp @@ -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(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(testInitInstance()); VerifyOrQuit(instance != nullptr, "Null OpenThread instance\n"); - messagePool = &instance->Get(); + messagePool = &instance->Get(); - 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(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; }