diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index b758946..a370cc9 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -3,6 +3,8 @@ add_library(internal ${CMAKE_CURRENT_SOURCE_DIR}/url.cpp ${CMAKE_CURRENT_SOURCE_DIR}/platform/native_socket.cpp ${CMAKE_CURRENT_SOURCE_DIR}/platform/network_runtime.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/platform/socket_mode.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/platform/socket_io.cpp ${CMAKE_CURRENT_SOURCE_DIR}/platform/socket_connect.cpp ${CMAKE_CURRENT_SOURCE_DIR}/net/resolver.cpp ${CMAKE_CURRENT_SOURCE_DIR}/net/tcp_connection.cpp diff --git a/src/net/tcp_connection.cpp b/src/net/tcp_connection.cpp index faeebc2..02dda56 100644 --- a/src/net/tcp_connection.cpp +++ b/src/net/tcp_connection.cpp @@ -2,10 +2,30 @@ #include "platform/network_runtime.hpp" #include "platform/socket_connect.hpp" +#include "platform/socket_io.hpp" +#include "platform/socket_mode.hpp" #include namespace cpp_request::detail::net { +namespace { + +std::chrono::milliseconds remaining_timeout( + std::chrono::steady_clock::time_point deadline) noexcept { + const auto now = std::chrono::steady_clock::now(); + if (now >= deadline) { + return std::chrono::milliseconds{0}; + } + + auto remaining = std::chrono::duration_cast( + deadline - now); + if (remaining.count() <= 0) { + remaining = std::chrono::milliseconds{1}; + } + return remaining; +} + +} // namespace Result TcpConnection::connect( const std::vector& endpoints, @@ -52,8 +72,7 @@ Result TcpConnection::connect( } created_socket = true; - const auto remaining = std::chrono::duration_cast( - deadline - std::chrono::steady_clock::now()); + const auto remaining = remaining_timeout(deadline); const auto attempt = platform::connect_with_timeout( socket.get(), @@ -82,4 +101,188 @@ Result TcpConnection::connect( return Error{ErrorCode::ConnectFailed, last_native_error}; } +Result TcpConnection::write_all( + std::string_view data, + std::chrono::milliseconds timeout) { + if (!connected()) { + return Error{ErrorCode::ConnectionClosed}; + } + + if (data.empty()) { + return std::size_t{0}; + } + + if (timeout.count() <= 0) { + close(); + return Error{ErrorCode::WriteTimeout}; + } + + int native_code = 0; + if (!platform::set_socket_nonblocking(socket_.get(), true, native_code)) { + close(); + return Error{ErrorCode::WriteFailed, native_code}; + } + + const auto deadline = std::chrono::steady_clock::now() + timeout; + std::size_t total_written = 0; + + while (total_written < data.size()) { + if (std::chrono::steady_clock::now() >= deadline) { + close(); + return Error{ErrorCode::WriteTimeout}; + } + + const auto attempt = platform::send_some( + socket_.get(), + data.data() + total_written, + data.size() - total_written); + + if (attempt.transferred > 0) { + total_written += static_cast(attempt.transferred); + continue; + } + + if (attempt.transferred == 0) { + close(); + return Error{ErrorCode::ConnectionClosed}; + } + + native_code = attempt.native_code; + if (platform::socket_error_is_interrupted(native_code)) { + continue; + } + if (platform::socket_error_is_connection_closed(native_code)) { + close(); + return Error{ErrorCode::ConnectionClosed, native_code}; + } + if (!platform::socket_error_is_would_block(native_code)) { + close(); + return Error{ErrorCode::WriteFailed, native_code}; + } + + const auto remaining = remaining_timeout(deadline); + if (remaining.count() <= 0) { + close(); + return Error{ErrorCode::WriteTimeout}; + } + + const auto wait = platform::wait_socket_writable(socket_.get(), remaining); + if (wait.status == platform::SocketWaitStatus::Ready) { + continue; + } + if (wait.status == platform::SocketWaitStatus::Failed + && platform::socket_error_is_interrupted(wait.native_code)) { + continue; + } + + close(); + if (wait.status == platform::SocketWaitStatus::TimedOut) { + return Error{ErrorCode::WriteTimeout, wait.native_code}; + } + if (platform::socket_error_is_connection_closed(wait.native_code)) { + return Error{ErrorCode::ConnectionClosed, wait.native_code}; + } + return Error{ErrorCode::WriteFailed, wait.native_code}; + } + + if (!platform::set_socket_nonblocking(socket_.get(), false, native_code)) { + close(); + return Error{ErrorCode::WriteFailed, native_code}; + } + + return total_written; +} + +Result TcpConnection::read_some( + char* buffer, + std::size_t capacity, + std::chrono::milliseconds timeout) { + if (!connected()) { + return Error{ErrorCode::ConnectionClosed}; + } + + if (capacity == 0) { + return std::size_t{0}; + } + + if (buffer == nullptr) { + return Error{ErrorCode::ReadFailed}; + } + + if (timeout.count() <= 0) { + close(); + return Error{ErrorCode::ReadTimeout}; + } + + int native_code = 0; + if (!platform::set_socket_nonblocking(socket_.get(), true, native_code)) { + close(); + return Error{ErrorCode::ReadFailed, native_code}; + } + + const auto deadline = std::chrono::steady_clock::now() + timeout; + + for (;;) { + if (std::chrono::steady_clock::now() >= deadline) { + close(); + return Error{ErrorCode::ReadTimeout}; + } + + const auto attempt = platform::receive_some( + socket_.get(), + buffer, + capacity); + + if (attempt.transferred > 0) { + if (!platform::set_socket_nonblocking(socket_.get(), false, native_code)) { + close(); + return Error{ErrorCode::ReadFailed, native_code}; + } + return static_cast(attempt.transferred); + } + + if (attempt.transferred == 0) { + close(); + return Error{ErrorCode::ConnectionClosed}; + } + + native_code = attempt.native_code; + if (platform::socket_error_is_interrupted(native_code)) { + continue; + } + if (platform::socket_error_is_connection_closed(native_code)) { + close(); + return Error{ErrorCode::ConnectionClosed, native_code}; + } + if (!platform::socket_error_is_would_block(native_code)) { + close(); + return Error{ErrorCode::ReadFailed, native_code}; + } + + const auto remaining = remaining_timeout(deadline); + if (remaining.count() <= 0) { + close(); + return Error{ErrorCode::ReadTimeout}; + } + + const auto wait = platform::wait_socket_readable(socket_.get(), remaining); + if (wait.status == platform::SocketWaitStatus::Ready) { + continue; + } + if (wait.status == platform::SocketWaitStatus::Failed + && platform::socket_error_is_interrupted(wait.native_code)) { + continue; + } + + close(); + if (wait.status == platform::SocketWaitStatus::TimedOut) { + return Error{ErrorCode::ReadTimeout, wait.native_code}; + } + if (platform::socket_error_is_connection_closed(wait.native_code)) { + return Error{ErrorCode::ConnectionClosed, wait.native_code}; + } + return Error{ErrorCode::ReadFailed, wait.native_code}; + } +} + } // namespace cpp_request::detail::net diff --git a/src/net/tcp_connection.hpp b/src/net/tcp_connection.hpp index 23141ee..fc4d196 100644 --- a/src/net/tcp_connection.hpp +++ b/src/net/tcp_connection.hpp @@ -6,6 +6,8 @@ #include #include +#include +#include #include #include @@ -24,6 +26,15 @@ class TcpConnection final { const std::vector& endpoints, std::chrono::milliseconds timeout); + [[nodiscard]] Result write_all( + std::string_view data, + std::chrono::milliseconds timeout); + + [[nodiscard]] Result read_some( + char* buffer, + std::size_t capacity, + std::chrono::milliseconds timeout); + [[nodiscard]] bool connected() const noexcept { return socket_.valid(); } diff --git a/src/platform/socket_connect.cpp b/src/platform/socket_connect.cpp index 82bd63e..ba89ed6 100644 --- a/src/platform/socket_connect.cpp +++ b/src/platform/socket_connect.cpp @@ -1,49 +1,23 @@ #include "platform/socket_connect.hpp" +#include "platform/socket_io.hpp" +#include "platform/socket_mode.hpp" + #ifdef _WIN32 #include #else #include -#include -#include #include #endif -#include -#include - namespace cpp_request::detail::platform { namespace { #ifdef _WIN32 -bool set_nonblocking(NativeSocketHandle socket, bool enabled, int& native_code) noexcept { - u_long mode = enabled ? 1UL : 0UL; - if (::ioctlsocket(socket, FIONBIO, &mode) == 0) { - return true; - } - native_code = ::WSAGetLastError(); - return false; -} - bool connect_in_progress(int code) noexcept { return code == WSAEWOULDBLOCK || code == WSAEINPROGRESS || code == WSAEALREADY; } #else -bool set_nonblocking(NativeSocketHandle socket, bool enabled, int& native_code) noexcept { - const int flags = ::fcntl(socket, F_GETFL, 0); - if (flags == -1) { - native_code = errno; - return false; - } - - const int next = enabled ? (flags | O_NONBLOCK) : (flags & ~O_NONBLOCK); - if (::fcntl(socket, F_SETFL, next) == 0) { - return true; - } - native_code = errno; - return false; -} - bool connect_in_progress(int code) noexcept { return code == EINPROGRESS || code == EWOULDBLOCK || code == EALREADY; } @@ -80,7 +54,23 @@ NativeSocket create_tcp_socket( native_code = last_socket_error(); return {}; } - return NativeSocket{handle}; + + NativeSocket socket{handle}; + +#if !defined(_WIN32) && defined(SO_NOSIGPIPE) && !defined(MSG_NOSIGNAL) + const int enabled = 1; + if (::setsockopt( + handle, + SOL_SOCKET, + SO_NOSIGPIPE, + &enabled, + static_cast(sizeof(enabled))) != 0) { + native_code = last_socket_error(); + return {}; + } +#endif + + return socket; } ConnectAttemptResult connect_with_timeout( @@ -93,7 +83,7 @@ ConnectAttemptResult connect_with_timeout( } int native_code = 0; - if (!set_nonblocking(socket, true, native_code)) { + if (!set_socket_nonblocking(socket, true, native_code)) { return {ConnectStatus::Failed, native_code}; } @@ -110,7 +100,7 @@ ConnectAttemptResult connect_with_timeout( #endif if (connect_result == 0) { - if (!set_nonblocking(socket, false, native_code)) { + if (!set_socket_nonblocking(socket, false, native_code)) { return {ConnectStatus::Failed, native_code}; } return {ConnectStatus::Connected, 0}; @@ -121,46 +111,17 @@ ConnectAttemptResult connect_with_timeout( return {ConnectStatus::Failed, native_code}; } -#ifdef _WIN32 - fd_set write_set; - fd_set except_set; - FD_ZERO(&write_set); - FD_ZERO(&except_set); - FD_SET(socket, &write_set); - FD_SET(socket, &except_set); - - const auto total_ms = timeout.count(); - timeval tv{}; - tv.tv_sec = static_cast(total_ms / 1000); - tv.tv_usec = static_cast((total_ms % 1000) * 1000); - - const int wait_result = ::select(0, nullptr, &write_set, &except_set, &tv); - if (wait_result == 0) { - return {ConnectStatus::TimedOut, 0}; + const auto wait = wait_socket_writable(socket, timeout); + if (wait.status == SocketWaitStatus::TimedOut) { + return {ConnectStatus::TimedOut, wait.native_code}; } - if (wait_result == SOCKET_ERROR) { - return {ConnectStatus::Failed, ::WSAGetLastError()}; + if (wait.status == SocketWaitStatus::Failed) { + return {ConnectStatus::Failed, wait.native_code}; } -#else - pollfd descriptor{}; - descriptor.fd = socket; - descriptor.events = POLLOUT; - - const auto bounded = std::min( - timeout.count(), - std::numeric_limits::max()); - const int wait_result = ::poll(&descriptor, 1, static_cast(bounded)); - if (wait_result == 0) { - return {ConnectStatus::TimedOut, 0}; - } - if (wait_result < 0) { - return {ConnectStatus::Failed, errno}; - } -#endif auto result = socket_error_result(socket); if (result.status == ConnectStatus::Connected) { - if (!set_nonblocking(socket, false, native_code)) { + if (!set_socket_nonblocking(socket, false, native_code)) { return {ConnectStatus::Failed, native_code}; } } diff --git a/src/platform/socket_io.cpp b/src/platform/socket_io.cpp new file mode 100644 index 0000000..6bd09b0 --- /dev/null +++ b/src/platform/socket_io.cpp @@ -0,0 +1,184 @@ +#include "platform/socket_io.hpp" + +#ifdef _WIN32 +#include +#else +#include +#include +#include +#endif + +#include +#include + +namespace cpp_request::detail::platform { +namespace { + +SocketWaitResult wait_socket( + NativeSocketHandle socket, + std::chrono::milliseconds timeout, + bool writable) noexcept { + if (timeout.count() <= 0) { + return {SocketWaitStatus::TimedOut, 0}; + } + +#ifdef _WIN32 + fd_set ready_set; + fd_set except_set; + FD_ZERO(&ready_set); + FD_ZERO(&except_set); + FD_SET(socket, &ready_set); + FD_SET(socket, &except_set); + + const long long total_ms = std::max(1, timeout.count()); + const long long seconds = std::min( + total_ms / 1000, + (std::numeric_limits::max)()); + + timeval tv{}; + tv.tv_sec = static_cast(seconds); + tv.tv_usec = seconds == (std::numeric_limits::max)() + ? 0L + : static_cast((total_ms % 1000) * 1000); + + const int wait_result = writable + ? ::select(0, nullptr, &ready_set, &except_set, &tv) + : ::select(0, &ready_set, nullptr, &except_set, &tv); + + if (wait_result == 0) { + return {SocketWaitStatus::TimedOut, 0}; + } + if (wait_result == SOCKET_ERROR) { + return {SocketWaitStatus::Failed, ::WSAGetLastError()}; + } + return {SocketWaitStatus::Ready, 0}; +#else + pollfd descriptor{}; + descriptor.fd = socket; + descriptor.events = writable ? POLLOUT : POLLIN; + + const long long total_ms = std::max(1, timeout.count()); + const long long bounded = std::min( + total_ms, + (std::numeric_limits::max)()); + + const int wait_result = ::poll( + &descriptor, + 1, + static_cast(bounded)); + + if (wait_result == 0) { + return {SocketWaitStatus::TimedOut, 0}; + } + if (wait_result < 0) { + return {SocketWaitStatus::Failed, errno}; + } + return {SocketWaitStatus::Ready, 0}; +#endif +} + +std::size_t bounded_io_size(std::size_t size) noexcept { + return std::min( + size, + static_cast((std::numeric_limits::max)())); +} + +} // namespace + +SocketWaitResult wait_socket_readable( + NativeSocketHandle socket, + std::chrono::milliseconds timeout) noexcept { + return wait_socket(socket, timeout, false); +} + +SocketWaitResult wait_socket_writable( + NativeSocketHandle socket, + std::chrono::milliseconds timeout) noexcept { + return wait_socket(socket, timeout, true); +} + +SocketIoAttempt send_some( + NativeSocketHandle socket, + const char* data, + std::size_t size) noexcept { + const std::size_t bounded = bounded_io_size(size); + +#ifdef _WIN32 + const int result = ::send( + socket, + data, + static_cast(bounded), + 0); + if (result == SOCKET_ERROR) { + return {-1, ::WSAGetLastError()}; + } + return {result, 0}; +#else + int flags = 0; +#ifdef MSG_NOSIGNAL + flags |= MSG_NOSIGNAL; +#endif + const ssize_t result = ::send(socket, data, bounded, flags); + if (result < 0) { + return {-1, errno}; + } + return {static_cast(result), 0}; +#endif +} + +SocketIoAttempt receive_some( + NativeSocketHandle socket, + char* data, + std::size_t size) noexcept { + const std::size_t bounded = bounded_io_size(size); + +#ifdef _WIN32 + const int result = ::recv( + socket, + data, + static_cast(bounded), + 0); + if (result == SOCKET_ERROR) { + return {-1, ::WSAGetLastError()}; + } + return {result, 0}; +#else + const ssize_t result = ::recv(socket, data, bounded, 0); + if (result < 0) { + return {-1, errno}; + } + return {static_cast(result), 0}; +#endif +} + +bool socket_error_is_would_block(int code) noexcept { +#ifdef _WIN32 + return code == WSAEWOULDBLOCK; +#else + return code == EAGAIN || code == EWOULDBLOCK; +#endif +} + +bool socket_error_is_interrupted(int code) noexcept { +#ifdef _WIN32 + return code == WSAEINTR; +#else + return code == EINTR; +#endif +} + +bool socket_error_is_connection_closed(int code) noexcept { +#ifdef _WIN32 + return code == WSAECONNRESET + || code == WSAECONNABORTED + || code == WSAENOTCONN + || code == WSAESHUTDOWN; +#else + return code == EPIPE + || code == ECONNRESET + || code == ENOTCONN + || code == ESHUTDOWN; +#endif +} + +} // namespace cpp_request::detail::platform diff --git a/src/platform/socket_io.hpp b/src/platform/socket_io.hpp new file mode 100644 index 0000000..71b214a --- /dev/null +++ b/src/platform/socket_io.hpp @@ -0,0 +1,48 @@ +#pragma once + +#include "platform/native_socket.hpp" + +#include +#include + +namespace cpp_request::detail::platform { + +enum class SocketWaitStatus { + Ready, + TimedOut, + Failed +}; + +struct SocketWaitResult { + SocketWaitStatus status{SocketWaitStatus::Failed}; + int native_code{0}; +}; + +struct SocketIoAttempt { + std::ptrdiff_t transferred{-1}; + int native_code{0}; +}; + +[[nodiscard]] SocketWaitResult wait_socket_readable( + NativeSocketHandle socket, + std::chrono::milliseconds timeout) noexcept; + +[[nodiscard]] SocketWaitResult wait_socket_writable( + NativeSocketHandle socket, + std::chrono::milliseconds timeout) noexcept; + +[[nodiscard]] SocketIoAttempt send_some( + NativeSocketHandle socket, + const char* data, + std::size_t size) noexcept; + +[[nodiscard]] SocketIoAttempt receive_some( + NativeSocketHandle socket, + char* data, + std::size_t size) noexcept; + +[[nodiscard]] bool socket_error_is_would_block(int code) noexcept; +[[nodiscard]] bool socket_error_is_interrupted(int code) noexcept; +[[nodiscard]] bool socket_error_is_connection_closed(int code) noexcept; + +} // namespace cpp_request::detail::platform diff --git a/src/platform/socket_mode.cpp b/src/platform/socket_mode.cpp new file mode 100644 index 0000000..5cb945d --- /dev/null +++ b/src/platform/socket_mode.cpp @@ -0,0 +1,42 @@ +#include "platform/socket_mode.hpp" + +#ifdef _WIN32 +#include +#else +#include +#include +#endif + +namespace cpp_request::detail::platform { + +bool set_socket_nonblocking( + NativeSocketHandle socket, + bool enabled, + int& native_code) noexcept { + native_code = 0; + +#ifdef _WIN32 + u_long mode = enabled ? 1UL : 0UL; + if (::ioctlsocket(socket, FIONBIO, &mode) == 0) { + return true; + } + native_code = ::WSAGetLastError(); + return false; +#else + const int flags = ::fcntl(socket, F_GETFL, 0); + if (flags == -1) { + native_code = errno; + return false; + } + + const int next = enabled ? (flags | O_NONBLOCK) : (flags & ~O_NONBLOCK); + if (::fcntl(socket, F_SETFL, next) == 0) { + return true; + } + + native_code = errno; + return false; +#endif +} + +} // namespace cpp_request::detail::platform diff --git a/src/platform/socket_mode.hpp b/src/platform/socket_mode.hpp new file mode 100644 index 0000000..ec471e5 --- /dev/null +++ b/src/platform/socket_mode.hpp @@ -0,0 +1,12 @@ +#pragma once + +#include "platform/native_socket.hpp" + +namespace cpp_request::detail::platform { + +[[nodiscard]] bool set_socket_nonblocking( + NativeSocketHandle socket, + bool enabled, + int& native_code) noexcept; + +} // namespace cpp_request::detail::platform diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index e0be62b..35aa12f 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -1,10 +1,12 @@ include(${CMAKE_SOURCE_DIR}/cmake/packages/google-test.cmake) +find_package(Threads REQUIRED) add_executable(cpp_request_tests ${CMAKE_CURRENT_SOURCE_DIR}/url_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/native_socket_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/resolver_test.cpp ${CMAKE_CURRENT_SOURCE_DIR}/tcp_connection_test.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/tcp_io_test.cpp ) target_include_directories(cpp_request_tests @@ -16,6 +18,7 @@ target_link_libraries(cpp_request_tests PRIVATE cpp_request::internal GTest::gtest_main + Threads::Threads ) if(WIN32) diff --git a/tests/tcp_io_test.cpp b/tests/tcp_io_test.cpp new file mode 100644 index 0000000..8749a31 --- /dev/null +++ b/tests/tcp_io_test.cpp @@ -0,0 +1,297 @@ +#include + +#include "net/tcp_connection.hpp" +#include "platform/native_socket.hpp" +#include "platform/network_runtime.hpp" + +#include +#include +#include +#include +#include +#include + +#ifdef _WIN32 +#include +#else +#include +#include +#endif + +namespace { + +using cpp_request::ErrorCode; +using cpp_request::detail::net::Endpoint; +using cpp_request::detail::net::TcpConnection; +using cpp_request::detail::platform::NativeSocket; +using cpp_request::detail::platform::NativeSocketHandle; +using cpp_request::detail::platform::kInvalidSocket; + +struct LoopbackListener { + NativeSocket socket; + Endpoint endpoint; +}; + +LoopbackListener make_ipv4_listener() { + const auto runtime = cpp_request::detail::platform::ensure_network_runtime(); + if (!runtime.ok) { + return {}; + } + + NativeSocket listener{::socket(AF_INET, SOCK_STREAM, IPPROTO_TCP)}; + if (!listener) { + return {}; + } + + sockaddr_in address{}; + address.sin_family = AF_INET; + address.sin_port = 0; + address.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + + if (::bind( + listener.get(), + reinterpret_cast(&address), + static_cast(sizeof(address))) != 0) { + return {}; + } + + if (::listen(listener.get(), 1) != 0) { + return {}; + } + +#ifdef _WIN32 + int address_size = sizeof(address); +#else + socklen_t address_size = sizeof(address); +#endif + if (::getsockname( + listener.get(), + reinterpret_cast(&address), + &address_size) != 0) { + return {}; + } + + return { + std::move(listener), + Endpoint{ + reinterpret_cast(&address), + sizeof(address), + SOCK_STREAM, + IPPROTO_TCP}}; +} + +NativeSocket accept_one(NativeSocketHandle listener) { + return NativeSocket{::accept(listener, nullptr, nullptr)}; +} + +int receive_native( + NativeSocketHandle socket, + char* buffer, + std::size_t capacity) { +#ifdef _WIN32 + return ::recv( + socket, + buffer, + static_cast(capacity), + 0); +#else + return static_cast(::recv(socket, buffer, capacity, 0)); +#endif +} + +int send_native( + NativeSocketHandle socket, + const char* data, + std::size_t size) { +#ifdef _WIN32 + return ::send(socket, data, static_cast(size), 0); +#else + return static_cast(::send(socket, data, size, 0)); +#endif +} + +TEST(TcpIoTest, WriteAllTransfersCompletePayload) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto connection_result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(connection_result); + + NativeSocket peer = accept_one(listener.socket.get()); + ASSERT_TRUE(peer); + + std::string payload(2 * 1024 * 1024, 'x'); + std::string received; + received.reserve(payload.size()); + + std::thread reader([&] { + std::array buffer{}; + while (received.size() < payload.size()) { + const int count = receive_native( + peer.get(), + buffer.data(), + buffer.size()); + if (count <= 0) { + break; + } + received.append(buffer.data(), static_cast(count)); + } + }); + + auto write_result = connection_result.value().write_all( + payload, + std::chrono::seconds{3}); + + reader.join(); + + ASSERT_TRUE(write_result); + EXPECT_EQ(write_result.value(), payload.size()); + EXPECT_EQ(received, payload); + EXPECT_TRUE(connection_result.value().connected()); +} + +TEST(TcpIoTest, ReadSomeReturnsAvailableBytes) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto connection_result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(connection_result); + + NativeSocket peer = accept_one(listener.socket.get()); + ASSERT_TRUE(peer); + + const std::string message = "hello from loopback"; + ASSERT_EQ( + send_native(peer.get(), message.data(), message.size()), + static_cast(message.size())); + + std::array buffer{}; + auto read_result = connection_result.value().read_some( + buffer.data(), + buffer.size(), + std::chrono::seconds{2}); + + ASSERT_TRUE(read_result); + EXPECT_EQ( + std::string_view(buffer.data(), read_result.value()), + message); + EXPECT_TRUE(connection_result.value().connected()); +} + +TEST(TcpIoTest, ReadTimeoutClosesConnection) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto connection_result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(connection_result); + + NativeSocket peer = accept_one(listener.socket.get()); + ASSERT_TRUE(peer); + + std::array buffer{}; + auto read_result = connection_result.value().read_some( + buffer.data(), + buffer.size(), + std::chrono::milliseconds{75}); + + ASSERT_FALSE(read_result); + EXPECT_EQ(read_result.error().code, ErrorCode::ReadTimeout); + EXPECT_FALSE(connection_result.value().connected()); +} + +TEST(TcpIoTest, PeerCloseIsReportedAsConnectionClosed) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto connection_result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(connection_result); + + NativeSocket peer = accept_one(listener.socket.get()); + ASSERT_TRUE(peer); + peer.close(); + + std::array buffer{}; + auto read_result = connection_result.value().read_some( + buffer.data(), + buffer.size(), + std::chrono::seconds{2}); + + ASSERT_FALSE(read_result); + EXPECT_EQ(read_result.error().code, ErrorCode::ConnectionClosed); + EXPECT_FALSE(connection_result.value().connected()); +} + +TEST(TcpIoTest, WriteTimeoutClosesConnection) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto connection_result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(connection_result); + + NativeSocket peer = accept_one(listener.socket.get()); + ASSERT_TRUE(peer); + + const std::string payload = "timeout"; + auto write_result = connection_result.value().write_all( + payload, + std::chrono::milliseconds{0}); + + ASSERT_FALSE(write_result); + EXPECT_EQ(write_result.error().code, ErrorCode::WriteTimeout); + EXPECT_FALSE(connection_result.value().connected()); +} + +TEST(TcpIoTest, EmptyWriteIsSuccessfulNoOp) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto connection_result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(connection_result); + + NativeSocket peer = accept_one(listener.socket.get()); + ASSERT_TRUE(peer); + + auto write_result = connection_result.value().write_all( + {}, + std::chrono::milliseconds{0}); + + ASSERT_TRUE(write_result); + EXPECT_EQ(write_result.value(), 0U); + EXPECT_TRUE(connection_result.value().connected()); +} + +TEST(TcpIoTest, ZeroCapacityReadIsSuccessfulNoOp) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto connection_result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(connection_result); + + NativeSocket peer = accept_one(listener.socket.get()); + ASSERT_TRUE(peer); + + auto read_result = connection_result.value().read_some( + nullptr, + 0, + std::chrono::milliseconds{0}); + + ASSERT_TRUE(read_result); + EXPECT_EQ(read_result.value(), 0U); + EXPECT_TRUE(connection_result.value().connected()); +} + +} // namespace diff --git a/vcpkg.json b/vcpkg.json index ddd812e..81e1b28 100644 --- a/vcpkg.json +++ b/vcpkg.json @@ -4,7 +4,6 @@ "description": "A lightweight, dependency-free HTTP client library written from scratch in C/C++ to understand low-level network communication and TCP socket programming.", "builtin-baseline": "a1cae005c39be7b18ba319fced856b68d7276271", "dependencies": [ - "benchmark", "gtest" ] }