diff --git a/src/core/common/message.cpp b/src/core/common/message.cpp index 3962e3e5b..d4341b805 100644 --- a/src/core/common/message.cpp +++ b/src/core/common/message.cpp @@ -434,6 +434,8 @@ void Message::GetFirstChunk(uint16_t aOffset, uint16_t &aLength, Chunk &aChunk) // its length. The `aLength` is also decreased by the chunk // length. + VerifyOrExit(aOffset < GetLength(), aChunk.mLength = 0); + if (aOffset + aLength >= GetLength()) { aLength = GetLength() - aOffset; @@ -512,8 +514,6 @@ 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) @@ -523,7 +523,6 @@ uint16_t Message::Read(uint16_t aOffset, uint16_t aLength, void *aBuf) const GetNextChunk(aLength, chunk); } -exit: return static_cast(bufPtr - reinterpret_cast(aBuf)); } @@ -544,23 +543,19 @@ void Message::Write(uint16_t aOffset, uint16_t aLength, const void *aBuf) } } -int Message::CopyTo(uint16_t aSourceOffset, uint16_t aDestinationOffset, uint16_t aLength, Message &aMessage) const +uint16_t Message::CopyTo(uint16_t aSourceOffset, uint16_t aDestinationOffset, uint16_t aLength, Message &aMessage) const { uint16_t bytesCopied = 0; - uint16_t bytesToCopy; - uint8_t buf[16]; + Chunk chunk; - while (aLength > 0) + GetFirstChunk(aSourceOffset, aLength, chunk); + + while (chunk.GetLength() > 0) { - bytesToCopy = (aLength < sizeof(buf)) ? aLength : sizeof(buf); - - Read(aSourceOffset, bytesToCopy, buf); - aMessage.Write(aDestinationOffset, bytesToCopy, buf); - - aSourceOffset += bytesToCopy; - aDestinationOffset += bytesToCopy; - aLength -= bytesToCopy; - bytesCopied += bytesToCopy; + aMessage.Write(aDestinationOffset, chunk.GetLength(), chunk.GetData()); + aDestinationOffset += chunk.GetLength(); + bytesCopied += chunk.GetLength(); + GetNextChunk(aLength, chunk); } return bytesCopied; diff --git a/src/core/common/message.hpp b/src/core/common/message.hpp index bdb4f268c..3492e91fe 100644 --- a/src/core/common/message.hpp +++ b/src/core/common/message.hpp @@ -540,7 +540,7 @@ public: * @returns The number of bytes copied. * */ - int CopyTo(uint16_t aSourceOffset, uint16_t aDestinationOffset, uint16_t aLength, Message &aMessage) const; + uint16_t CopyTo(uint16_t aSourceOffset, uint16_t aDestinationOffset, uint16_t aLength, Message &aMessage) const; /** * This method creates a copy of the message. diff --git a/tests/unit/test_message.cpp b/tests/unit/test_message.cpp index aa52f7c2b..f97494bca 100644 --- a/tests/unit/test_message.cpp +++ b/tests/unit/test_message.cpp @@ -40,12 +40,15 @@ void TestMessage(void) { enum : uint16_t { - kMaxSize = (kBufferSize * 3 + 24), + kMaxSize = (kBufferSize * 3 + 24), + kOffsetStep = 101, + kLengthStep = 21, }; Instance * instance; MessagePool *messagePool; Message * message; + Message * message2; uint8_t writeBuffer[kMaxSize]; uint8_t readBuffer[kMaxSize]; uint8_t zeroBuffer[kMaxSize]; @@ -68,7 +71,7 @@ void TestMessage(void) for (uint16_t offset = 0; offset < kMaxSize; offset++) { - for (uint16_t length = 0; length < kMaxSize - offset; length++) + for (uint16_t length = 0; length <= kMaxSize - offset; length++) { for (uint16_t i = 0; i < length; i++) { @@ -85,9 +88,62 @@ void TestMessage(void) VerifyOrQuit(memcmp(readBuffer, &writeBuffer[offset], length) == 0, "Message compare failed"); VerifyOrQuit(memcmp(&readBuffer[length], zeroBuffer, kMaxSize - length) == 0, "Message read after length"); } + + // Verify `Read()` behavior when requested read length goes beyond available bytes in the message. + + for (uint16_t length = kMaxSize - offset + 1; length <= kMaxSize + 1; length++) + { + uint16_t readLength; + + memset(readBuffer, 0, sizeof(readBuffer)); + readLength = message->Read(offset, length, readBuffer); + VerifyOrQuit(readLength <= length, "Message::Read() returned longer length"); + VerifyOrQuit(readLength == kMaxSize - offset, "Message::Read failed"); + VerifyOrQuit(memcmp(readBuffer, &writeBuffer[offset], readLength) == 0, "Message compare failed"); + VerifyOrQuit(memcmp(&readBuffer[readLength], zeroBuffer, kMaxSize - readLength) == 0, "read after length"); + } } VerifyOrQuit(message->GetLength() == kMaxSize, "Message::GetLength failed"); + + // Test `Message::CopyTo()` behavior. + + VerifyOrQuit((message2 = messagePool->New(Message::kTypeIp6, 0)) != nullptr, "Message::New failed"); + SuccessOrQuit(message2->SetLength(kMaxSize), "Message::SetLength failed"); + + for (uint16_t srcOffset = 0; srcOffset < kMaxSize; srcOffset += kOffsetStep) + { + for (uint16_t dstOffset = 0; dstOffset < kMaxSize; dstOffset += kOffsetStep) + { + for (uint16_t length = 0; length <= kMaxSize - dstOffset; length += kLengthStep) + { + uint16_t bytesCopied; + + message2->Write(0, kMaxSize, zeroBuffer); + + bytesCopied = message->CopyTo(srcOffset, dstOffset, length, *message2); + + if (srcOffset + length <= kMaxSize) + { + VerifyOrQuit(bytesCopied == length, "CopyTo() failed"); + } + else + { + VerifyOrQuit(bytesCopied == kMaxSize - srcOffset, "CopyTo() failed"); + } + + VerifyOrQuit(message2->Read(0, kMaxSize, readBuffer) == kMaxSize, "Message::Read failed"); + + VerifyOrQuit(memcmp(&readBuffer[0], zeroBuffer, dstOffset) == 0, "read before length"); + VerifyOrQuit(memcmp(&readBuffer[dstOffset], &writeBuffer[srcOffset], bytesCopied) == 0, + "Compare failed"); + VerifyOrQuit( + memcmp(&readBuffer[dstOffset + bytesCopied], zeroBuffer, kMaxSize - bytesCopied - dstOffset) == 0, + "read after length"); + } + } + } + message->Free(); testFreeInstance(instance);