diff --git a/src/lib/synergy/ProtocolUtil.cpp b/src/lib/synergy/ProtocolUtil.cpp index 9771c2379..de18dda21 100644 --- a/src/lib/synergy/ProtocolUtil.cpp +++ b/src/lib/synergy/ProtocolUtil.cpp @@ -16,6 +16,7 @@ * along with this program. If not, see . */ +#include #include "synergy/ProtocolUtil.h" #include "io/IStream.h" #include "base/Log.h" @@ -47,24 +48,25 @@ ProtocolUtil::writef(synergy::IStream* stream, const char* fmt, ...) bool ProtocolUtil::readf(synergy::IStream* stream, const char* fmt, ...) { - assert(stream != NULL); - assert(fmt != NULL); - LOG((CLOG_DEBUG2 "readf(%s)", fmt)); + bool result = false; - bool result; - va_list args; - va_start(args, fmt); - try { - vreadf(stream, fmt, args); - result = true; + if (stream && fmt) { + LOG((CLOG_DEBUG2 "readf(%s)", fmt)); + va_list args; + va_start(args, fmt); + try { + vreadf(stream, fmt, args); + result = true; + } + catch (XIO&) { + result = false; + } + catch (const std::bad_alloc&) { + result = false; + } + va_end(args); } - catch (XIO&) { - result = false; - } - catch (std::bad_alloc & exception) { - result = false; - } - va_end(args); + return result; } @@ -111,15 +113,56 @@ ProtocolUtil::vreadf(synergy::IStream* stream, const char* fmt, va_list args) UInt32 len = eatLength(&fmt); switch (*fmt) { case 'i': { - readInt(stream, len, args); + void* destination = va_arg(args, void*); + switch (len) { + case 1: + // 1 byte integer + *static_cast(destination) = read1ByteInt(stream); + break; + case 2: + // 2 byte integer + *static_cast(destination) = read2BytesInt(stream); + break; + case 4: + // 4 byte integer + *static_cast(destination) = read4BytesInt(stream); + break; + default: + //the length is wrong + LOG((CLOG_ERR "read: length to be read is wrong: '%d' should be 1,2, or 4", len)); + assert(false); //assert for debugging + break; + } break; } + case 'I': { - readVectorInt(stream, len, args); + void* destination = va_arg(args, void*); + switch (len) { + case 1: + // 1 byte integer + readVector1ByteInt(stream, *static_cast*>(destination)); + break; + case 2: + // 2 byte integer + readVector2BytesInt(stream, *static_cast*>(destination)); + break; + case 4: + // 4 byte integer + readVector4BytesInt(stream, *static_cast*>(destination)); + break; + default: + //the length is wrong + LOG((CLOG_ERR "read: length to be read is wrong: '%d' should be 1,2, or 4", len)); + assert(false); //assert for debugging + break; + } break; } + case 's': { - readBytes(stream, len, args); + String* destination = va_arg(args, String*); + readBytes(stream, len, destination); break; } @@ -412,115 +455,76 @@ ProtocolUtil::read(synergy::IStream* stream, void* vbuffer, UInt32 count) } } -void ProtocolUtil::readInt(synergy::IStream * stream, UInt32 len, va_list args) { - // check for valid length - if (len == 4 || len == 2 || len == 1) { +UInt8 ProtocolUtil::read1ByteInt(synergy::IStream * stream) +{ + const UInt32 BufferSize = 1; + std::array buffer = {}; + read(stream, buffer.data(), BufferSize); - static const int buffer_size = 4; - // read the data - UInt8 buffer[buffer_size]; - //Read the buffer till the len or buffers_size, which ever is smaller - read(stream, buffer, len > buffer_size ? buffer_size : len); + UInt8 Result = buffer[0]; + LOG((CLOG_DEBUG2 "readf: read 1 byte integer: %d (0x%x)", Result, Result)); - // convert it - void* v = va_arg(args, void*); - switch (len) { - case 1: - // 1 byte integer - *static_cast(v) = buffer[0]; - LOG((CLOG_DEBUG2 "readf: read %d byte integer: %d (0x%x)", len, *static_cast(v), *static_cast(v))); - break; + return Result; +} - case 2: - // 2 byte integer - *static_cast(v) = - static_cast( - (static_cast(buffer[0]) << 8) | - static_cast(buffer[1])); - LOG((CLOG_DEBUG2 "readf: read %d byte integer: %d (0x%x)", len, *static_cast(v), *static_cast(v))); - break; +UInt16 ProtocolUtil::read2BytesInt(synergy::IStream * stream) +{ + const UInt32 BufferSize = 2; + std::array buffer = {}; + read(stream, buffer.data(), BufferSize); - case 4: - // 4 byte integer - *static_cast(v) = - (static_cast(buffer[0]) << 24) | - (static_cast(buffer[1]) << 16) | - (static_cast(buffer[2]) << 8) | - static_cast(buffer[3]); - LOG((CLOG_DEBUG2 "readf: read %d byte integer: %d (0x%x)", len, *static_cast(v), *static_cast(v))); - break; - } - } - else { - //the length is wrong - LOG((CLOG_ERR "read: length to be read is wrong: '%d' should be 1,2, or 4", len)); - assert(false); //assert for debugging + UInt16 Result = static_cast((static_cast(buffer[0]) << 8) | static_cast(buffer[1])); + LOG((CLOG_DEBUG2 "readf: read 2 byte integer: %d (0x%x)", Result, Result)); + + return Result; +} + +UInt32 ProtocolUtil::read4BytesInt(synergy::IStream * stream) +{ + const int BufferSize = 4; + std::array buffer = {}; + read(stream, buffer.data(), BufferSize); + + UInt32 Result = (static_cast(buffer[0]) << 24) | + (static_cast(buffer[1]) << 16) | + (static_cast(buffer[2]) << 8) | + (static_cast(buffer[3])); + + LOG((CLOG_DEBUG2 "readf: read 4 byte integer: %d (0x%x)", Result, Result)); + + return Result; +} + +void ProtocolUtil::readVector1ByteInt(synergy::IStream* stream, std::vector& destination) +{ + UInt32 size = read4BytesInt(stream); + for (UInt32 i = 0; i < size; ++i) { + destination.push_back(read1ByteInt(stream)); } } -void ProtocolUtil::readVectorInt(synergy::IStream * stream, UInt32 len, va_list args) { - // check for valid length - assert(len == 1 || len == 2 || len == 4); - - // read the vector length - UInt8 buffer[4]; - read(stream, buffer, 4); - UInt32 n = (static_cast(buffer[0]) << 24) | - (static_cast(buffer[1]) << 16) | - (static_cast(buffer[2]) << 8) | - static_cast(buffer[3]); - - // convert it - void* v = va_arg(args, void*); - switch (len) { - case 1: - // 1 byte integer - for (UInt32 i = 0; i < n; ++i) { - read(stream, buffer, 1); - static_cast*>(v)->push_back( - buffer[0]); - LOG((CLOG_DEBUG2 "readf: read %d byte integer[%d]: %d (0x%x)", len, i, static_cast*>(v)->back(), static_cast*>(v)->back())); - } - break; - - case 2: - // 2 byte integer - for (UInt32 i = 0; i < n; ++i) { - read(stream, buffer, 2); - static_cast*>(v)->push_back( - static_cast( - (static_cast(buffer[0]) << 8) | - static_cast(buffer[1]))); - LOG((CLOG_DEBUG2 "readf: read %d byte integer[%d]: %d (0x%x)", len, i, static_cast*>(v)->back(), static_cast*>(v)->back())); - } - break; - - case 4: - // 4 byte integer - for (UInt32 i = 0; i < n; ++i) { - read(stream, buffer, 4); - static_cast*>(v)->push_back( - (static_cast(buffer[0]) << 24) | - (static_cast(buffer[1]) << 16) | - (static_cast(buffer[2]) << 8) | - static_cast(buffer[3])); - LOG((CLOG_DEBUG2 "readf: read %d byte integer[%d]: %d (0x%x)", len, i, static_cast*>(v)->back(), static_cast*>(v)->back())); - } - break; +void ProtocolUtil::readVector2BytesInt(synergy::IStream* stream, std::vector& destination) +{ + UInt32 size = read4BytesInt(stream); + for (UInt32 i = 0; i < size; ++i) { + destination.push_back(read2BytesInt(stream)); } } -void ProtocolUtil::readBytes(synergy::IStream * stream, UInt32 len, va_list args) { +void ProtocolUtil::readVector4BytesInt(synergy::IStream* stream, std::vector& destination) +{ + UInt32 size = read4BytesInt(stream); + for (UInt32 i = 0; i < size; ++i) { + destination.push_back(read4BytesInt(stream)); + } +} + +void ProtocolUtil::readBytes(synergy::IStream * stream, UInt32 len, String* destination) { assert(len == 0); // read the string length UInt8 buffer[128]; - read(stream, buffer, 4); - len = (static_cast(buffer[0]) << 24) | - (static_cast(buffer[1]) << 16) | - (static_cast(buffer[2]) << 8) | - static_cast(buffer[3]); - + len = read4BytesInt(stream); // use a fixed size buffer if its big enough const bool useFixed = (len <= sizeof(buffer)); @@ -552,8 +556,10 @@ void ProtocolUtil::readBytes(synergy::IStream * stream, UInt32 len, va_list args LOG((CLOG_DEBUG2 "readf: read %d byte string", len)); // save the data - String* dst = va_arg(args, String*); - dst->assign((const char*)sBuffer, len); + + if (destination){ + destination->assign((const char*)sBuffer, len); + } // release the buffer if (!useFixed) { diff --git a/src/lib/synergy/ProtocolUtil.h b/src/lib/synergy/ProtocolUtil.h index 04fad66e2..a49da3096 100644 --- a/src/lib/synergy/ProtocolUtil.h +++ b/src/lib/synergy/ProtocolUtil.h @@ -86,17 +86,21 @@ private: /** * @brief Handles 1,2, or 4 byte Integers */ - static void readInt(synergy::IStream*, UInt32, va_list); + static UInt8 read1ByteInt(synergy::IStream * stream); + static UInt16 read2BytesInt(synergy::IStream * stream); + static UInt32 read4BytesInt(synergy::IStream * stream); /** * @brief Handles a Vector of integers */ - static void readVectorInt(synergy::IStream*, UInt32, va_list); + static void readVector1ByteInt(synergy::IStream*, std::vector&); + static void readVector2BytesInt(synergy::IStream*, std::vector&); + static void readVector4BytesInt(synergy::IStream*, std::vector&); /** * @brief Handles an array of bytes */ - static void readBytes(synergy::IStream*, UInt32, va_list); + static void readBytes(synergy::IStream*, UInt32, String*); }; //! Mismatched read exception diff --git a/src/test/unittests/synergy/ProtocolUtilTests.cpp b/src/test/unittests/synergy/ProtocolUtilTests.cpp new file mode 100644 index 000000000..5b7ec40cf --- /dev/null +++ b/src/test/unittests/synergy/ProtocolUtilTests.cpp @@ -0,0 +1,422 @@ +/* + * synergy -- mouse and keyboard sharing utility + * Copyright (C) 2014-2020 Symless Ltd. + * + * This package is free software; you can redistribute it and/or + * modify it under the terms of the GNU General Public License + * found in the file LICENSE that should have accompanied this file. + * + * This package is distributed in the hope that it will be useful, + * but WITHOUT ANY WARRANTY; without even the implied warranty of + * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + * GNU General Public License for more details. + * + * You should have received a copy of the GNU General Public License + * along with this program. If not, see . + */ + +#include +#include "test/global/gtest.h" +#include "test/mock/io/MockStream.h" +#include "synergy/ProtocolUtil.h" + +using ::testing::_; +using ::testing::Return; +using ::testing::DoAll; +using ::testing::SetArgPointee; +using ::testing::Eq; +using ::testing::StrEq; +using ::testing::TypedEq; + +ACTION_P2(SetValueToVoidPointerArg0, value, size) +{ + memcpy(arg0, value, size); +} + +ACTION(ThrowBadAlloc) +{ + throw std::bad_alloc(); +} + +class ProtocolUtilTests : public ::testing::Test +{ +public: + MockStream stream; + UInt8 ActualInt8 = 0; + UInt16 ActualInt16 = 0; + UInt32 ActualInt32 = 0; + std::string ActualString; +}; + +TEST_F(ProtocolUtilTests, readf__XIOEndOfStream_exception) +{ + ON_CALL(stream, read(_, _)).WillByDefault(Return(0)); + + EXPECT_FALSE(ProtocolUtil::readf(&stream, "%s", &ActualString)); + EXPECT_TRUE(ActualString.empty()); +} + +TEST_F(ProtocolUtilTests, readf_XIOReadMismatch_exception) +{ + EXPECT_CALL(stream, read(_, _)) + .WillOnce(DoAll(SetValueToVoidPointerArg0("b", 1), Return(1))); + + EXPECT_FALSE(ProtocolUtil::readf(&stream, "a%s", &ActualString)); + EXPECT_TRUE(ActualString.empty()); +} + +TEST_F(ProtocolUtilTests, readf_bad_alloc_exception) +{ + ON_CALL(stream, read(_, _)).WillByDefault(ThrowBadAlloc()); + + EXPECT_FALSE(ProtocolUtil::readf(&stream, "a%s", &ActualString)); + EXPECT_TRUE(ActualString.empty()); +} + +TEST_F(ProtocolUtilTests, readf_asserts) +{ + ASSERT_DEBUG_DEATH( + {ProtocolUtil::readf(&stream, "%x", &ActualString);}, + "invalid format specifier" + ); + + ASSERT_DEBUG_DEATH( + {ProtocolUtil::readf(&stream, "%5i", &ActualString);}, + "length to be read is wrong:" + ); + + ASSERT_DEBUG_DEATH( + {ProtocolUtil::readf(&stream, "%5I", &ActualString);}, + "" + ); +} + +TEST_F(ProtocolUtilTests, readf_params_validation) +{ + EXPECT_FALSE(ProtocolUtil::readf(NULL, "%x", NULL)); + EXPECT_FALSE(ProtocolUtil::readf(&stream, NULL, NULL)); +} + + +TEST_F(ProtocolUtilTests, readf_string) +{ + const UInt8 Length = 200; + const std::string Expected(Length, 'x'); + std::array StringSize = {0,0,0,Length}; + + EXPECT_CALL(stream, read(_, _)) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StringSize.data(), StringSize.size()), + Return(StringSize.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(Expected.c_str(), Expected.length()), + Return(Expected.length()) + ) + ); + + EXPECT_TRUE(ProtocolUtil::readf(&stream, "%s", &ActualString)); + EXPECT_EQ(Expected, ActualString); +} + +class ReadfIntTestFixture : public ::testing::TestWithParam< std::tuple > +{ +public: + MockStream stream; + UInt8 StreamData1Byte = 10; + std::array StreamData2Bytes = {0, 10}; + std::array StreamData4Bytes = {0, 0, 0, 10}; + + UInt8* getStreamData(int size) + { + UInt8* StreamData = nullptr; + switch(size){ + case 2: + StreamData = StreamData2Bytes.data(); + break; + case 4: + StreamData = StreamData4Bytes.data(); + break; + default: + StreamData = &StreamData1Byte; + break; + } + return StreamData; + } +}; + +TEST_P(ReadfIntTestFixture, readf_int) +{ + int Actual = 0; + const int Expected = 10; + const char* Format = std::get<0>(GetParam()); + int StreamDataSize = std::get<1>(GetParam()); + UInt8* StreamData = getStreamData(StreamDataSize); + + ON_CALL(stream, read(_, _)) + .WillByDefault( + DoAll( + SetValueToVoidPointerArg0(StreamData, StreamDataSize), + Return(StreamDataSize) + ) + ); + + EXPECT_TRUE(ProtocolUtil::readf(&stream, Format, &Actual)); + EXPECT_EQ(Expected, Actual); +} + +INSTANTIATE_TEST_CASE_P( + ReadfIntTests, + ReadfIntTestFixture, + ::testing::Values( + std::make_tuple("%1i", 1), + std::make_tuple("%2i", 2), + std::make_tuple("%4i", 4))); + +class ReadfIntVectorTestFixture : public ReadfIntTestFixture +{ +}; + +TEST_P(ReadfIntVectorTestFixture, readf_int_vector) +{ + std::vector Actual1Byte = {}; + std::vector Actual2Bytes = {}; + std::vector Actual4Bytes = {}; + + const std::vector Expected1Byte = {10,10}; + const std::vector Expected2Bytes = {10,10}; + const std::vector Expected4Bytes = {10,10}; + std::array StreamVectorSize = {0,0,0,2}; + + const char* Format = std::get<0>(GetParam()); + int StreamDataSize = std::get<1>(GetParam()); + UInt8* StreamData = getStreamData(StreamDataSize); + + MockStream stream; + EXPECT_CALL(stream, read(_, _)) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StreamVectorSize.data(), StreamVectorSize.size()), + Return(StreamVectorSize.size()) + ) + ) + .WillRepeatedly( + DoAll( + SetValueToVoidPointerArg0(StreamData, StreamDataSize), + Return(StreamDataSize) + )); + + switch(StreamDataSize){ + case 2: + EXPECT_TRUE(ProtocolUtil::readf(&stream, Format, &Actual2Bytes)); + EXPECT_EQ(Expected2Bytes, Actual2Bytes); + break; + case 4: + EXPECT_TRUE(ProtocolUtil::readf(&stream, Format, &Actual4Bytes)); + EXPECT_EQ(Expected4Bytes, Actual4Bytes); + break; + default: + EXPECT_TRUE(ProtocolUtil::readf(&stream, Format, &Actual1Byte)); + EXPECT_EQ(Expected1Byte, Actual1Byte); + break; + } +} + +INSTANTIATE_TEST_CASE_P( + ReadfIntVectorTests, + ReadfIntVectorTestFixture, + ::testing::Values( + std::make_tuple("%1I", 1), + std::make_tuple("%2I", 2), + std::make_tuple("%4I", 4))); + +class ReadfIntAndStringTest : public ReadfIntTestFixture +{ +public: + UInt8 ActualInt8 = 0; + UInt16 ActualInt16 = 0; + UInt32 ActualInt32 = 32; + std::string ActualString; +}; + +TEST_P(ReadfIntAndStringTest, readf_int_and_string) +{ + const int ExpectedInt = 10; + const UInt8 StringLength = 200; + const std::string ExpectedString(StringLength, 'x'); + std::array StringSize = {0,0,0,StringLength}; + + const char* Format = std::get<0>(GetParam()); + int StreamDataSize = std::get<1>(GetParam()); + UInt8* StreamData = getStreamData(StreamDataSize); + + EXPECT_CALL(stream, read(_, _)) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StreamData, StreamDataSize), + Return(StreamDataSize) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StringSize.data(), StringSize.size()), + Return(StringSize.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(ExpectedString.c_str(), ExpectedString.length()), + Return(ExpectedString.length()) + ) + ); + + switch(StreamDataSize){ + case 2: + EXPECT_TRUE(ProtocolUtil::readf(&stream, Format, &ActualInt16, &ActualString)); + EXPECT_EQ(ExpectedInt, ActualInt16); + break; + case 4: + EXPECT_TRUE(ProtocolUtil::readf(&stream, Format, &ActualInt32, &ActualString)); + EXPECT_EQ(ExpectedInt, ActualInt32); + break; + default: + EXPECT_TRUE(ProtocolUtil::readf(&stream, Format, &ActualInt8, &ActualString)); + EXPECT_EQ(ExpectedInt, ActualInt8); + break; + } + EXPECT_EQ(ExpectedString, ActualString); +} + +INSTANTIATE_TEST_CASE_P( + IntAndStringTest, + ReadfIntAndStringTest, + ::testing::Values( + std::make_tuple("%1i%s", 1), + std::make_tuple("%2i%s", 2), + std::make_tuple("%4i%s", 4))); + + +TEST_F(ProtocolUtilTests, readf_string_and_int4bytes) +{ + const UInt8 ExpectedInt = 10; + std::array StreamIntData = {0,0,0,ExpectedInt}; + + const std::string ExpectedStr(32768, 'x'); + std::array Size = {0,0,128,0}; + + EXPECT_CALL(stream, read(_, _)) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(Size.data(), Size.size()), + Return(Size.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(ExpectedStr.c_str(), ExpectedStr.length()), + Return(ExpectedStr.length()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StreamIntData.data(), StreamIntData.size()), + Return(StreamIntData.size()) + ) + ); + + EXPECT_TRUE(ProtocolUtil::readf(&stream, "%s%4i", &ActualString, &ActualInt32)); + EXPECT_EQ(ExpectedStr, ActualString); + EXPECT_EQ(ExpectedInt, ActualInt32); +} + +TEST_F(ProtocolUtilTests, readf_string_and_vector_int4bytes) +{ + std::vector Actual = {}; + const std::vector Expected4Bytes = {10,10}; + std::array StreamVectorSize = {0,0,0,2}; + std::array StreamData4Bytes = {0, 0, 0, 10}; + + const std::string ExpString(32768, 'x'); + std::array SizeString = {0,0,128,0}; + + EXPECT_CALL(stream, read(_, _)) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(SizeString.data(), SizeString.size()), + Return(SizeString.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(ExpString.c_str(), ExpString.length()), + Return(ExpString.length()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StreamVectorSize.data(), StreamVectorSize.size()), + Return(StreamVectorSize.size()) + ) + ) + .WillRepeatedly( + DoAll( + SetValueToVoidPointerArg0(StreamData4Bytes.data(), StreamData4Bytes.size()), + Return(StreamData4Bytes.size()) + ) + ); + + EXPECT_TRUE(ProtocolUtil::readf(&stream, "%s%4I", &ActualString, &Actual)); + EXPECT_EQ(ExpString, ActualString); + EXPECT_EQ(Expected4Bytes, Actual); +} + +TEST_F(ProtocolUtilTests, readf_vector_int4bytes_and_string) +{ + std::vector Actual4Bytes = {}; + const std::vector Expected4Bytes = {10,10}; + std::array StreamVectorSize = {0,0,0,2}; + std::array StreamData4Bytes = {0, 0, 0, 10}; + + const std::string ExpectedString(32768, 'x'); + std::array StringSize = {0,0,128,0}; + + EXPECT_CALL(stream, read(_, _)) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StreamVectorSize.data(), StreamVectorSize.size()), + Return(StreamVectorSize.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StreamData4Bytes.data(), StreamData4Bytes.size()), + Return(StreamData4Bytes.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StreamData4Bytes.data(), StreamData4Bytes.size()), + Return(StreamData4Bytes.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(StringSize.data(), StringSize.size()), + Return(StringSize.size()) + ) + ) + .WillOnce( + DoAll( + SetValueToVoidPointerArg0(ExpectedString.c_str(), ExpectedString.length()), + Return(ExpectedString.length()) + ) + ); + + EXPECT_TRUE(ProtocolUtil::readf(&stream, "%4I%s", &Actual4Bytes, &ActualString)); + EXPECT_EQ(ExpectedString, ActualString); + EXPECT_EQ(Expected4Bytes, Actual4Bytes); +} +