diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index a2fc7f7..5a741b0 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -47,6 +47,8 @@ target_include_directories(netft_posix_transport PRIVATE ) target_link_options(netft_posix_transport PRIVATE "-Wl,--wrap=poll" + "-Wl,--wrap=recv" + "-Wl,--wrap=send" ) target_link_libraries(netft_posix_transport PRIVATE ${NETFT_GTEST_MAIN_TARGET}) @@ -140,3 +142,21 @@ target_link_libraries(netft_cli_test PRIVATE add_dependencies(netft_cli_test netft_cli) add_test(NAME netft_cli COMMAND netft_cli_test) + +add_test( + NAME netft_cli_help + COMMAND ${CMAKE_COMMAND} + -DPROGRAM=$ + -DEXPECTED_EXIT_CODE=0 + -DARGUMENTS=--help + -P ${CMAKE_CURRENT_SOURCE_DIR}/assert_exit_code.cmake +) + +add_test( + NAME netft_cli_invalid_command + COMMAND ${CMAKE_COMMAND} + -DPROGRAM=$ + -DEXPECTED_EXIT_CODE=2 + -DARGUMENTS=invalid-command + -P ${CMAKE_CURRENT_SOURCE_DIR}/assert_exit_code.cmake +) diff --git a/test/assert_exit_code.cmake b/test/assert_exit_code.cmake new file mode 100644 index 0000000..c350247 --- /dev/null +++ b/test/assert_exit_code.cmake @@ -0,0 +1,23 @@ +if(NOT DEFINED PROGRAM OR NOT DEFINED EXPECTED_EXIT_CODE) + message(FATAL_ERROR "PROGRAM and EXPECTED_EXIT_CODE are required") +endif() + +execute_process( + COMMAND "${PROGRAM}" ${ARGUMENTS} + TIMEOUT 30 + RESULT_VARIABLE actual_exit_code + OUTPUT_QUIET + ERROR_QUIET +) + +if(NOT "${actual_exit_code}" MATCHES "^-?[0-9]+$") + message(FATAL_ERROR + "Failed to execute ${PROGRAM}: ${actual_exit_code}" + ) +endif() + +if(NOT actual_exit_code EQUAL EXPECTED_EXIT_CODE) + message(FATAL_ERROR + "Expected exit code ${EXPECTED_EXIT_CODE}, got ${actual_exit_code}" + ) +endif() diff --git a/test/test_posix_transport.cpp b/test/test_posix_transport.cpp index 782c684..aa4a9a4 100644 --- a/test/test_posix_transport.cpp +++ b/test/test_posix_transport.cpp @@ -3,11 +3,14 @@ #include #include +#include +#include #include #include #include #include +#include #include #include #include @@ -16,42 +19,99 @@ namespace { using namespace std::chrono_literals; -enum class PollBehavior { RepeatedEintrThenTimeout, EintrPastDeadline, Error }; +enum class PollBehavior { + RepeatedEintrThenTimeout, + EintrPastDeadline, + Error, + InvalidDescriptor, + Readable, + Timeout +}; +enum class SendBehavior { Complete, Error, ShortWrite }; +enum class ReceiveBehavior { Payload, Interrupted, Error }; constexpr int kInterruptCount = 3; std::vector poll_timeouts; int poll_calls{}; PollBehavior poll_behavior{PollBehavior::RepeatedEintrThenTimeout}; +SendBehavior send_behavior{SendBehavior::Complete}; +ReceiveBehavior receive_behavior{ReceiveBehavior::Payload}; + +void reset_behaviors() { + poll_timeouts.clear(); + poll_calls = 0; + poll_behavior = PollBehavior::RepeatedEintrThenTimeout; + send_behavior = SendBehavior::Complete; + receive_behavior = ReceiveBehavior::Payload; +} } // namespace -extern "C" int __wrap_poll(pollfd *, nfds_t, const int timeout) { +extern "C" int __wrap_poll(pollfd *descriptors, const nfds_t count, const int timeout) { poll_timeouts.push_back(timeout); - if (poll_behavior == PollBehavior::Error) { + switch (poll_behavior) { + case PollBehavior::Error: errno = EBADF; return -1; - } - if (poll_behavior == PollBehavior::EintrPastDeadline) { + case PollBehavior::InvalidDescriptor: + if (count > 0) { + descriptors[0].revents = POLLNVAL; + } + return 1; + case PollBehavior::Readable: + if (count > 0) { + descriptors[0].revents = POLLIN; + } + return 1; + case PollBehavior::Timeout: + return 0; + case PollBehavior::EintrPastDeadline: if (poll_calls++ == 0) { std::this_thread::sleep_for(std::chrono::milliseconds{timeout} + 10ms); errno = EINTR; return -1; } return 0; + case PollBehavior::RepeatedEintrThenTimeout: + if (poll_calls++ < kInterruptCount) { + std::this_thread::sleep_for(10ms); + errno = EINTR; + return -1; + } + std::this_thread::sleep_for(std::chrono::milliseconds{timeout}); + return 0; } - if (poll_calls++ < kInterruptCount) { - std::this_thread::sleep_for(10ms); - errno = EINTR; + return 0; +} + +extern "C" ssize_t __wrap_send(int, const void *, const std::size_t length, int) { + if (send_behavior == SendBehavior::Error) { + errno = EIO; return -1; } + if (send_behavior == SendBehavior::ShortWrite) { + return static_cast(length - 1); + } + return static_cast(length); +} - std::this_thread::sleep_for(std::chrono::milliseconds{timeout}); - return 0; +extern "C" ssize_t __wrap_recv(int, void *data, const std::size_t length, int) { + if (receive_behavior == ReceiveBehavior::Interrupted) { + errno = EINTR; + return -1; + } + if (receive_behavior == ReceiveBehavior::Error) { + errno = EIO; + return -1; + } + constexpr std::size_t payload_size = 4; + const auto copied = std::min(length, payload_size); + std::fill_n(static_cast(data), copied, 0U); + return static_cast(copied); } TEST(PosixTransportTest, RepeatedEintrPreservesOriginalTimeout) { - poll_timeouts.clear(); - poll_calls = 0; + reset_behaviors(); poll_behavior = PollBehavior::RepeatedEintrThenTimeout; netft::detail::PosixTransport transport; @@ -67,8 +127,7 @@ TEST(PosixTransportTest, RepeatedEintrPreservesOriginalTimeout) { } TEST(PosixTransportTest, NonEintrPollErrorStillThrows) { - poll_timeouts.clear(); - poll_calls = 0; + reset_behaviors(); poll_behavior = PollBehavior::Error; netft::detail::PosixTransport transport; @@ -80,8 +139,7 @@ TEST(PosixTransportTest, NonEintrPollErrorStillThrows) { } TEST(PosixTransportTest, ExpiredDeadlineDoesNotRepoll) { - poll_timeouts.clear(); - poll_calls = 0; + reset_behaviors(); poll_behavior = PollBehavior::EintrPastDeadline; netft::detail::PosixTransport transport; @@ -91,3 +149,87 @@ TEST(PosixTransportTest, ExpiredDeadlineDoesNotRepoll) { EXPECT_EQ(transport.receive(buffer.data(), buffer.size(), 20ms), 0U); EXPECT_EQ(poll_timeouts.size(), 1U); } + +TEST(PosixTransportTest, SendBeforeConnectThrows) { + reset_behaviors(); + netft::detail::PosixTransport transport; + std::array request{}; + EXPECT_THROW(transport.send(request), std::runtime_error); +} + +TEST(PosixTransportTest, ReceiveBeforeConnectThrows) { + reset_behaviors(); + netft::detail::PosixTransport transport; + std::array buffer{}; + EXPECT_THROW(transport.receive(buffer.data(), buffer.size(), 20ms), std::runtime_error); +} + +TEST(PosixTransportTest, SendErrorThrows) { + reset_behaviors(); + send_behavior = SendBehavior::Error; + netft::detail::PosixTransport transport; + transport.connect("127.0.0.1", 49152); + std::array request{}; + EXPECT_THROW(transport.send(request), std::runtime_error); +} + +TEST(PosixTransportTest, ShortSendThrows) { + reset_behaviors(); + send_behavior = SendBehavior::ShortWrite; + netft::detail::PosixTransport transport; + transport.connect("127.0.0.1", 49152); + std::array request{}; + EXPECT_THROW(transport.send(request), std::runtime_error); +} + +TEST(PosixTransportTest, InvalidPollDescriptorThrows) { + reset_behaviors(); + poll_behavior = PollBehavior::InvalidDescriptor; + netft::detail::PosixTransport transport; + transport.connect("127.0.0.1", 49152); + std::array buffer{}; + EXPECT_THROW(transport.receive(buffer.data(), buffer.size(), 20ms), std::runtime_error); +} + +TEST(PosixTransportTest, InterruptedReceiveReturnsZero) { + reset_behaviors(); + poll_behavior = PollBehavior::Readable; + receive_behavior = ReceiveBehavior::Interrupted; + netft::detail::PosixTransport transport; + transport.connect("127.0.0.1", 49152); + std::array buffer{}; + EXPECT_EQ(transport.receive(buffer.data(), buffer.size(), 20ms), 0U); +} + +TEST(PosixTransportTest, ReceiveErrorThrows) { + reset_behaviors(); + poll_behavior = PollBehavior::Readable; + receive_behavior = ReceiveBehavior::Error; + netft::detail::PosixTransport transport; + transport.connect("127.0.0.1", 49152); + std::array buffer{}; + EXPECT_THROW(transport.receive(buffer.data(), buffer.size(), 20ms), std::runtime_error); +} + +TEST(PosixTransportTest, ReadableDatagramReturnsPayloadSize) { + reset_behaviors(); + poll_behavior = PollBehavior::Readable; + receive_behavior = ReceiveBehavior::Payload; + netft::detail::PosixTransport transport; + transport.connect("127.0.0.1", 49152); + std::array buffer{}; + EXPECT_EQ(transport.receive(buffer.data(), buffer.size(), 20ms), 4U); +} + +TEST(PosixTransportTest, LargeTimeoutSaturatesPollMilliseconds) { + reset_behaviors(); + poll_behavior = PollBehavior::Timeout; + netft::detail::PosixTransport transport; + transport.connect("127.0.0.1", 49152); + std::array buffer{}; + const auto timeout = std::chrono::duration{ + static_cast(std::numeric_limits::max()) / 1000.0 + 1.0}; + EXPECT_EQ(transport.receive(buffer.data(), buffer.size(), timeout), 0U); + ASSERT_EQ(poll_timeouts.size(), 1U); + EXPECT_EQ(poll_timeouts.front(), std::numeric_limits::max()); +}