diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 4a818fc..b758946 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -3,7 +3,9 @@ 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_connect.cpp ${CMAKE_CURRENT_SOURCE_DIR}/net/resolver.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/net/tcp_connection.cpp ) add_library(${PROJECT_NAME}::internal ALIAS internal) diff --git a/src/net/tcp_connection.cpp b/src/net/tcp_connection.cpp new file mode 100644 index 0000000..faeebc2 --- /dev/null +++ b/src/net/tcp_connection.cpp @@ -0,0 +1,85 @@ +#include "net/tcp_connection.hpp" + +#include "platform/network_runtime.hpp" +#include "platform/socket_connect.hpp" + +#include + +namespace cpp_request::detail::net { + +Result TcpConnection::connect( + const std::vector& endpoints, + std::chrono::milliseconds timeout) { + if (timeout.count() <= 0) { + return Error{ErrorCode::ConnectTimeout}; + } + + if (endpoints.empty()) { + return Error{ErrorCode::ConnectFailed}; + } + + const auto runtime = platform::ensure_network_runtime(); + if (!runtime.ok) { + return Error{ErrorCode::SocketCreateFailed, runtime.native_code}; + } + + const auto deadline = std::chrono::steady_clock::now() + timeout; + bool created_socket = false; + bool saw_valid_endpoint = false; + int last_native_error = 0; + + for (const Endpoint& endpoint : endpoints) { + if (!endpoint.valid()) { + continue; + } + + saw_valid_endpoint = true; + const auto now = std::chrono::steady_clock::now(); + if (now >= deadline) { + return Error{ErrorCode::ConnectTimeout, last_native_error}; + } + + int create_error = 0; + auto socket = platform::create_tcp_socket( + endpoint.family(), + endpoint.socket_type(), + endpoint.protocol(), + create_error); + + if (!socket) { + last_native_error = create_error; + continue; + } + + created_socket = true; + const auto remaining = std::chrono::duration_cast( + deadline - std::chrono::steady_clock::now()); + + const auto attempt = platform::connect_with_timeout( + socket.get(), + endpoint.native_address(), + endpoint.native_address_size(), + remaining); + + if (attempt.status == platform::ConnectStatus::Connected) { + return TcpConnection{std::move(socket)}; + } + + last_native_error = attempt.native_code; + if (attempt.status == platform::ConnectStatus::TimedOut) { + return Error{ErrorCode::ConnectTimeout, last_native_error}; + } + } + + if (!saw_valid_endpoint) { + return Error{ErrorCode::ConnectFailed}; + } + + if (!created_socket) { + return Error{ErrorCode::SocketCreateFailed, last_native_error}; + } + + return Error{ErrorCode::ConnectFailed, last_native_error}; +} + +} // namespace cpp_request::detail::net diff --git a/src/net/tcp_connection.hpp b/src/net/tcp_connection.hpp new file mode 100644 index 0000000..23141ee --- /dev/null +++ b/src/net/tcp_connection.hpp @@ -0,0 +1,46 @@ +#pragma once + +#include "net/endpoint.hpp" +#include "platform/native_socket.hpp" + +#include + +#include +#include +#include + +namespace cpp_request::detail::net { + +class TcpConnection final { +public: + TcpConnection() noexcept = default; + + TcpConnection(const TcpConnection&) = delete; + TcpConnection& operator=(const TcpConnection&) = delete; + TcpConnection(TcpConnection&&) noexcept = default; + TcpConnection& operator=(TcpConnection&&) noexcept = default; + + [[nodiscard]] static Result connect( + const std::vector& endpoints, + std::chrono::milliseconds timeout); + + [[nodiscard]] bool connected() const noexcept { + return socket_.valid(); + } + + [[nodiscard]] platform::NativeSocketHandle native_handle() const noexcept { + return socket_.get(); + } + + void close() noexcept { + socket_.close(); + } + +private: + explicit TcpConnection(platform::NativeSocket socket) noexcept + : socket_(std::move(socket)) {} + + platform::NativeSocket socket_; +}; + +} // namespace cpp_request::detail::net diff --git a/src/platform/socket_connect.cpp b/src/platform/socket_connect.cpp new file mode 100644 index 0000000..82bd63e --- /dev/null +++ b/src/platform/socket_connect.cpp @@ -0,0 +1,170 @@ +#include "platform/socket_connect.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; +} +#endif + +ConnectAttemptResult socket_error_result(NativeSocketHandle socket) noexcept { + int error = 0; +#ifdef _WIN32 + int length = sizeof(error); + if (::getsockopt(socket, SOL_SOCKET, SO_ERROR, reinterpret_cast(&error), &length) != 0) { + return {ConnectStatus::Failed, ::WSAGetLastError()}; + } +#else + socklen_t length = sizeof(error); + if (::getsockopt(socket, SOL_SOCKET, SO_ERROR, &error, &length) != 0) { + return {ConnectStatus::Failed, errno}; + } +#endif + return error == 0 + ? ConnectAttemptResult{ConnectStatus::Connected, 0} + : ConnectAttemptResult{ConnectStatus::Failed, error}; +} + +} // namespace + +NativeSocket create_tcp_socket( + int family, + int socket_type, + int protocol, + int& native_code) noexcept { + native_code = 0; + const NativeSocketHandle handle = ::socket(family, socket_type, protocol); + if (handle == kInvalidSocket) { + native_code = last_socket_error(); + return {}; + } + return NativeSocket{handle}; +} + +ConnectAttemptResult connect_with_timeout( + NativeSocketHandle socket, + const sockaddr* address, + std::size_t address_size, + std::chrono::milliseconds timeout) noexcept { + if (timeout.count() <= 0) { + return {ConnectStatus::TimedOut, 0}; + } + + int native_code = 0; + if (!set_nonblocking(socket, true, native_code)) { + return {ConnectStatus::Failed, native_code}; + } + +#ifdef _WIN32 + const int connect_result = ::connect( + socket, + address, + static_cast(address_size)); +#else + const int connect_result = ::connect( + socket, + address, + static_cast(address_size)); +#endif + + if (connect_result == 0) { + if (!set_nonblocking(socket, false, native_code)) { + return {ConnectStatus::Failed, native_code}; + } + return {ConnectStatus::Connected, 0}; + } + + native_code = last_socket_error(); + if (!connect_in_progress(native_code)) { + 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}; + } + if (wait_result == SOCKET_ERROR) { + return {ConnectStatus::Failed, ::WSAGetLastError()}; + } +#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)) { + return {ConnectStatus::Failed, native_code}; + } + } + return result; +} + +} // namespace cpp_request::detail::platform diff --git a/src/platform/socket_connect.hpp b/src/platform/socket_connect.hpp new file mode 100644 index 0000000..7a6882b --- /dev/null +++ b/src/platform/socket_connect.hpp @@ -0,0 +1,35 @@ +#pragma once + +#include "platform/native_socket.hpp" + +#include +#include + +struct sockaddr; + +namespace cpp_request::detail::platform { + +enum class ConnectStatus { + Connected, + Failed, + TimedOut +}; + +struct ConnectAttemptResult { + ConnectStatus status{ConnectStatus::Failed}; + int native_code{0}; +}; + +[[nodiscard]] NativeSocket create_tcp_socket( + int family, + int socket_type, + int protocol, + int& native_code) noexcept; + +[[nodiscard]] ConnectAttemptResult connect_with_timeout( + NativeSocketHandle socket, + const sockaddr* address, + std::size_t address_size, + std::chrono::milliseconds timeout) noexcept; + +} // namespace cpp_request::detail::platform diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 8a48995..e0be62b 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -4,6 +4,7 @@ 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 ) target_include_directories(cpp_request_tests diff --git a/tests/tcp_connection_test.cpp b/tests/tcp_connection_test.cpp new file mode 100644 index 0000000..b0a346a --- /dev/null +++ b/tests/tcp_connection_test.cpp @@ -0,0 +1,150 @@ +#include + +#include "net/tcp_connection.hpp" +#include "platform/native_socket.hpp" +#include "platform/network_runtime.hpp" + +#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::kInvalidSocket; + +static_assert(!std::is_copy_constructible_v); +static_assert(!std::is_copy_assignable_v); +static_assert(std::is_nothrow_move_constructible_v); +static_assert(std::is_nothrow_move_assignable_v); + +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 {}; + } + + Endpoint endpoint{ + reinterpret_cast(&address), + sizeof(address), + SOCK_STREAM, + IPPROTO_TCP}; + + return {std::move(listener), endpoint}; +} + +TEST(TcpConnectionTest, ConnectsToIpv4LoopbackListener) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + ASSERT_TRUE(listener.endpoint.valid()); + + auto result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + + ASSERT_TRUE(result); + EXPECT_TRUE(result.value().connected()); + EXPECT_NE(result.value().native_handle(), kInvalidSocket); +} + +TEST(TcpConnectionTest, FallsBackToLaterCandidate) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + const Endpoint invalid_socket_candidate{ + listener.endpoint.native_address(), + listener.endpoint.native_address_size(), + -1, + listener.endpoint.protocol()}; + + std::vector candidates; + candidates.push_back(invalid_socket_candidate); + candidates.push_back(listener.endpoint); + + auto result = TcpConnection::connect(candidates, std::chrono::seconds{2}); + + ASSERT_TRUE(result); + EXPECT_TRUE(result.value().connected()); +} + +TEST(TcpConnectionTest, ZeroTimeoutFailsImmediately) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::milliseconds{0}); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::ConnectTimeout); +} + +TEST(TcpConnectionTest, EmptyCandidateListIsConnectFailure) { + auto result = TcpConnection::connect({}, std::chrono::seconds{1}); + + ASSERT_FALSE(result); + EXPECT_EQ(result.error().code, ErrorCode::ConnectFailed); +} + +TEST(TcpConnectionTest, CloseInvalidatesConnection) { + auto listener = make_ipv4_listener(); + ASSERT_TRUE(listener.socket); + + auto result = TcpConnection::connect( + std::vector{listener.endpoint}, + std::chrono::seconds{2}); + ASSERT_TRUE(result); + + result.value().close(); + EXPECT_FALSE(result.value().connected()); +} + +} // namespace