From 10d49a3f6fb5f113933d8cf128de28dac9509160 Mon Sep 17 00:00:00 2001 From: mehrdadn Date: Sun, 1 Dec 2019 13:27:52 -0800 Subject: [PATCH] Use Boost's socket_holder instead of manually managing the socket (#6314) * Use Boost's socket_holder instead of manually managing sockets. Socket types are not ints on Windows, and we need to use wrapper for proper lifetime management regardless. --- src/ray/raylet/raylet_client.cc | 37 ++++++++++++++++----------------- src/ray/raylet/raylet_client.h | 8 ++++--- src/ray/rpc/grpc_server.cc | 13 +++++------- 3 files changed, 28 insertions(+), 30 deletions(-) diff --git a/src/ray/raylet/raylet_client.cc b/src/ray/raylet/raylet_client.cc index 72fc4e6bb..9b917c0cc 100644 --- a/src/ray/raylet/raylet_client.cc +++ b/src/ray/raylet/raylet_client.cc @@ -21,40 +21,37 @@ using MessageType = ray::protocol::MessageType; // TODO(rkn): The io methods below should be removed. -int connect_ipc_sock(const std::string &socket_pathname) { +bool connect_ipc_sock(Socket &sock, const std::string &socket_pathname) { struct sockaddr_un socket_address; - int socket_fd; - socket_fd = socket(AF_UNIX, SOCK_STREAM, 0); - if (socket_fd < 0) { + sock.reset(socket(AF_UNIX, SOCK_STREAM, 0)); + if (sock.get() < 0) { RAY_LOG(ERROR) << "socket() failed for pathname " << socket_pathname; - return -1; + return false; } memset(&socket_address, 0, sizeof(socket_address)); socket_address.sun_family = AF_UNIX; if (socket_pathname.length() + 1 > sizeof(socket_address.sun_path)) { RAY_LOG(ERROR) << "Socket pathname is too long."; - close(socket_fd); - return -1; + return false; } strncpy(socket_address.sun_path, socket_pathname.c_str(), socket_pathname.length() + 1); - if (connect(socket_fd, (struct sockaddr *)&socket_address, sizeof(socket_address)) != + if (connect(sock.get(), (struct sockaddr *)&socket_address, sizeof(socket_address)) != 0) { - close(socket_fd); - return -1; + return false; } - return socket_fd; + return true; } -int read_bytes(int socket_fd, uint8_t *cursor, size_t length) { +int read_bytes(Socket &conn, uint8_t *cursor, size_t length) { ssize_t nbytes = 0; // Termination condition: EOF or read 'length' bytes total. size_t bytesleft = length; size_t offset = 0; while (bytesleft > 0) { - nbytes = read(socket_fd, cursor + offset, bytesleft); + nbytes = read(conn.get(), cursor + offset, bytesleft); if (nbytes < 0) { if (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR) { continue; @@ -71,14 +68,14 @@ int read_bytes(int socket_fd, uint8_t *cursor, size_t length) { return 0; } -int write_bytes(int socket_fd, uint8_t *cursor, size_t length) { +int write_bytes(Socket &conn, uint8_t *cursor, size_t length) { ssize_t nbytes = 0; size_t bytesleft = length; size_t offset = 0; while (bytesleft > 0) { // While we haven't written the whole message, write to the file // descriptor, advance the cursor, and decrease the amount left to write. - nbytes = write(socket_fd, cursor + offset, bytesleft); + nbytes = write(conn.get(), cursor + offset, bytesleft); if (nbytes < 0) { if (errno == EAGAIN || errno == EWOULDBLOCK || errno == EINTR) { continue; @@ -105,10 +102,12 @@ RayletConnection::RayletConnection(const std::string &raylet_socket, int num_ret timeout = RayConfig::instance().connect_timeout_milliseconds(); } RAY_CHECK(!raylet_socket.empty()); - conn_ = -1; + bool connected = false; for (int num_attempts = 0; num_attempts < num_retries; ++num_attempts) { - conn_ = connect_ipc_sock(raylet_socket); - if (conn_ >= 0) break; + connected = connect_ipc_sock(conn_, raylet_socket); + if (connected) { + break; + } if (num_attempts > 0) { RAY_LOG(ERROR) << "Retrying to connect to socket for pathname " << raylet_socket << " (num_attempts = " << num_attempts @@ -118,7 +117,7 @@ RayletConnection::RayletConnection(const std::string &raylet_socket, int num_ret usleep(timeout * 1000); } // If we could not connect to the socket, exit. - if (conn_ == -1) { + if (!connected) { RAY_LOG(FATAL) << "Could not connect to socket " << raylet_socket; } } diff --git a/src/ray/raylet/raylet_client.h b/src/ray/raylet/raylet_client.h index 4671bf24a..f1e48c4de 100644 --- a/src/ray/raylet/raylet_client.h +++ b/src/ray/raylet/raylet_client.h @@ -7,6 +7,8 @@ #include #include +#include + #include "ray/common/status.h" #include "ray/common/task/task_spec.h" #include "ray/rpc/node_manager/node_manager_client.h" @@ -25,6 +27,7 @@ using ray::rpc::ProfileTableData; using MessageType = ray::protocol::MessageType; using ResourceMappingType = std::unordered_map>>; +using Socket = boost::asio::detail::socket_holder; using WaitResultPair = std::pair, std::vector>; class RayletConnection { @@ -40,7 +43,6 @@ class RayletConnection { /// \return The connection information. RayletConnection(const std::string &raylet_socket, int num_retries, int64_t timeout); - ~RayletConnection() { close(conn_); } /// Notify the raylet that this client is disconnecting gracefully. This /// is used by actors to exit gracefully so that the raylet doesn't /// propagate an error message to the driver. @@ -55,8 +57,8 @@ class RayletConnection { flatbuffers::FlatBufferBuilder *fbb = nullptr); private: - /// File descriptor of the Unix domain socket that connects to raylet. - int conn_; + /// The Unix domain socket that connects to raylet. + Socket conn_; /// A mutex to protect stateful operations of the raylet client. std::mutex mutex_; /// A mutex to protect write operations of the raylet client. diff --git a/src/ray/rpc/grpc_server.cc b/src/ray/rpc/grpc_server.cc index 4061e3598..c90cf6adb 100644 --- a/src/ray/rpc/grpc_server.cc +++ b/src/ray/rpc/grpc_server.cc @@ -1,21 +1,18 @@ - #include "src/ray/rpc/grpc_server.h" + #include +#include namespace { bool PortNotInUse(int port) { - int fd = socket(AF_INET, SOCK_STREAM, 0); - if (fd == -1) { - return false; - } + boost::asio::detail::socket_holder fd(socket(AF_INET, SOCK_STREAM, 0)); struct sockaddr_in server_addr = {0}; server_addr.sin_family = AF_INET; server_addr.sin_addr.s_addr = htonl(INADDR_ANY); server_addr.sin_port = htons(port); - int err = bind(fd, (struct sockaddr *)&server_addr, sizeof(server_addr)); - close(fd); - return err == 0; + return fd.get() >= 0 && + bind(fd.get(), (struct sockaddr *)&server_addr, sizeof(server_addr)) == 0; } } // namespace