From 2911018ad8e3abf9749649e6410379a57d46945a Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Tue, 25 Aug 2020 07:48:27 -0700 Subject: [PATCH] [message] update and fix 'CopyTo()' (#5454) This commit updates `Message::CopyTo()` method to use data chunks (this copies data directly from the source message to the destination avoiding the need for an extra buffer copy). It also addresses an issue where if the requested copy length is more than the available bytes in the source message, incorrect bytes may be copied into the destination message, and the returned value (i.e. number of bytes copied) can be invalid. This commit also updates the unit test `test_message` covering behavior of `CopyTo()` for different offsets and copy lengths. It also checks the behavior of `Read()` when the read length goes beyond the end of the message (trying to read more bytes than available in the message). --- src/core/common/message.cpp | 27 +++++++---------- src/core/common/message.hpp | 2 +- tests/unit/test_message.cpp | 60 +++++++++++++++++++++++++++++++++++-- 3 files changed, 70 insertions(+), 19 deletions(-) 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);