Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
207 changes: 205 additions & 2 deletions src/net/tcp_connection.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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 <chrono>

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<std::chrono::milliseconds>(
deadline - now);
if (remaining.count() <= 0) {
remaining = std::chrono::milliseconds{1};
}
return remaining;
}

} // namespace

Result<TcpConnection> TcpConnection::connect(
const std::vector<Endpoint>& endpoints,
Expand Down Expand Up @@ -52,8 +72,7 @@ Result<TcpConnection> TcpConnection::connect(
}

created_socket = true;
const auto remaining = std::chrono::duration_cast<std::chrono::milliseconds>(
deadline - std::chrono::steady_clock::now());
const auto remaining = remaining_timeout(deadline);

const auto attempt = platform::connect_with_timeout(
socket.get(),
Expand Down Expand Up @@ -82,4 +101,188 @@ Result<TcpConnection> TcpConnection::connect(
return Error{ErrorCode::ConnectFailed, last_native_error};
}

Result<std::size_t> 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<std::size_t>(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<std::size_t> 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<std::size_t>(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
11 changes: 11 additions & 0 deletions src/net/tcp_connection.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@
#include <cpp_request/result.hpp>

#include <chrono>
#include <cstddef>
#include <string_view>
#include <utility>
#include <vector>

Expand All @@ -24,6 +26,15 @@ class TcpConnection final {
const std::vector<Endpoint>& endpoints,
std::chrono::milliseconds timeout);

[[nodiscard]] Result<std::size_t> write_all(
std::string_view data,
std::chrono::milliseconds timeout);

[[nodiscard]] Result<std::size_t> read_some(
char* buffer,
std::size_t capacity,
std::chrono::milliseconds timeout);

[[nodiscard]] bool connected() const noexcept {
return socket_.valid();
}
Expand Down
Loading
Loading