diff --git a/src/ray/common/client_connection.cc b/src/ray/common/client_connection.cc index 9a5210fb3..361c66ce7 100644 --- a/src/ray/common/client_connection.cc +++ b/src/ray/common/client_connection.cc @@ -14,21 +14,50 @@ #include "client_connection.h" -#include - #include #include #include #include #include #include +#include #include +#include #include "ray/common/ray_config.h" #include "ray/util/util.h" namespace ray { +Status ConnectSocketRetry(local_stream_socket &socket, const std::string &endpoint, + int num_retries, int64_t timeout_in_ms) { + RAY_CHECK(num_retries != 0); + // Pick the default values if the user did not specify. + if (num_retries < 0) { + num_retries = RayConfig::instance().raylet_client_num_connect_attempts(); + } + if (timeout_in_ms < 0) { + timeout_in_ms = RayConfig::instance().raylet_client_connect_timeout_milliseconds(); + } + boost::system::error_code ec; + for (int num_attempts = 0; num_attempts < num_retries; ++num_attempts) { + // The latest boost::asio always returns void for connect(). Do not + // treat its return value as error code anymore. + socket.connect(ParseUrlEndpoint(endpoint), ec); + if (!ec) { + break; + } + if (num_attempts > 0) { + RAY_LOG(ERROR) << "Retrying to connect to socket for endpoint " << endpoint + << " (num_attempts = " << num_attempts + << ", num_retries = " << num_retries << ")"; + } + // Sleep for timeout milliseconds. + std::this_thread::sleep_for(std::chrono::milliseconds(timeout_in_ms)); + } + return boost_to_ray_status(ec); +} + std::shared_ptr ServerConnection::Create(local_stream_socket &&socket) { std::shared_ptr self(new ServerConnection(std::move(socket))); return self; diff --git a/src/ray/common/client_connection.h b/src/ray/common/client_connection.h index fab43bfae..6916c72a6 100644 --- a/src/ray/common/client_connection.h +++ b/src/ray/common/client_connection.h @@ -14,23 +14,26 @@ #pragma once -#include -#include - #include #include #include #include -#include +#include +#include #include "ray/common/id.h" #include "ray/common/status.h" +#include "ray/util/util.h" namespace ray { typedef boost::asio::generic::stream_protocol local_stream_protocol; typedef boost::asio::basic_stream_socket local_stream_socket; +/// Connect to a socket with retry times. +Status ConnectSocketRetry(local_stream_socket &socket, const std::string &endpoint, + int num_retries = -1, int64_t timeout_in_ms = -1); + /// \typename ServerConnection /// /// A generic type representing a client connection to a server. This typename diff --git a/src/ray/core_worker/store_provider/plasma_store_provider.cc b/src/ray/core_worker/store_provider/plasma_store_provider.cc index 4adb6d3f3..654d3259d 100644 --- a/src/ray/core_worker/store_provider/plasma_store_provider.cc +++ b/src/ray/core_worker/store_provider/plasma_store_provider.cc @@ -113,7 +113,7 @@ Status CoreWorkerPlasmaStoreProvider::Create(const std::shared_ptr &meta if (on_store_full_) { on_store_full_(); } - usleep(1000 * delay); + std::this_thread::sleep_for(std::chrono::milliseconds(delay)); delay *= 2; retries += 1; should_retry = true; diff --git a/src/ray/raylet/raylet_client.cc b/src/ray/raylet/raylet_client.cc index 9b7f023a3..d114a71f0 100644 --- a/src/ray/raylet/raylet_client.cc +++ b/src/ray/raylet/raylet_client.cc @@ -14,17 +14,7 @@ #include "raylet_client.h" -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - +#include "ray/common/client_connection.h" #include "ray/common/common_protocol.h" #include "ray/common/ray_config.h" #include "ray/common/task/task_spec.h" @@ -54,125 +44,33 @@ AddressesToFlatbuffer(flatbuffers::FlatBufferBuilder &fbb, namespace ray { -static int read_bytes(local_stream_socket &conn, void *cursor, size_t length) { - boost::system::error_code ec; - size_t nread = boost::asio::read(conn, boost::asio::buffer(cursor, length), ec); - return nread == length ? 0 : -1; -} - -static int write_bytes(local_stream_socket &conn, void *cursor, size_t length) { - boost::system::error_code ec; - size_t nread = boost::asio::write(conn, boost::asio::buffer(cursor, length), ec); - return nread == length ? 0 : -1; -} - raylet::RayletConnection::RayletConnection(boost::asio::io_service &io_service, const std::string &raylet_socket, - int num_retries, int64_t timeout) - : conn_(io_service) { - // Pick the default values if the user did not specify. - if (num_retries < 0) { - num_retries = RayConfig::instance().raylet_client_num_connect_attempts(); - } - if (timeout < 0) { - timeout = RayConfig::instance().raylet_client_connect_timeout_milliseconds(); - } - RAY_CHECK(!raylet_socket.empty()); - boost::system::error_code ec; - for (int num_attempts = 0; num_attempts < num_retries; ++num_attempts) { - if (!conn_.connect(ParseUrlEndpoint(raylet_socket), ec)) { - break; - } - if (num_attempts > 0) { - RAY_LOG(ERROR) << "Retrying to connect to socket for pathname " << raylet_socket - << " (num_attempts = " << num_attempts - << ", num_retries = " << num_retries << ")"; - } - // Sleep for timeout milliseconds. - usleep(timeout * 1000); - } + int num_retries, int64_t timeout) { + local_stream_socket socket(io_service); + Status s = ConnectSocketRetry(socket, raylet_socket, num_retries, timeout); // If we could not connect to the socket, exit. - if (ec) { + if (!s.ok()) { RAY_LOG(FATAL) << "Could not connect to socket " << raylet_socket; } -} - -Status raylet::RayletConnection::Disconnect() { - flatbuffers::FlatBufferBuilder fbb; - auto message = protocol::CreateDisconnectClient(fbb); - fbb.Finish(message); - auto status = WriteMessage(MessageType::IntentionalDisconnectClient, &fbb); - // Don't be too strict for disconnection errors. - // Just create logs and prevent it from crash. - if (!status.ok()) { - RAY_LOG(ERROR) << status.ToString() - << " [RayletClient] Failed to disconnect from raylet."; - } - return Status::OK(); -} - -Status raylet::RayletConnection::ReadMessage(MessageType type, - std::unique_ptr &message) { - int64_t cookie; - int64_t type_field; - int64_t length; - int closed = read_bytes(conn_, &cookie, sizeof(cookie)); - if (closed) goto disconnected; - RAY_CHECK(cookie == RayConfig::instance().ray_cookie()); - closed = read_bytes(conn_, &type_field, sizeof(type_field)); - if (closed) goto disconnected; - closed = read_bytes(conn_, &length, sizeof(length)); - if (closed) goto disconnected; - message = std::unique_ptr(new uint8_t[length]); - closed = read_bytes(conn_, message.get(), length); - if (closed) { - // Handle the case in which the socket is closed. - message.reset(nullptr); - disconnected: - message = nullptr; - type_field = static_cast(MessageType::DisconnectClient); - length = 0; - } - if (type_field == static_cast(MessageType::DisconnectClient)) { - return Status::IOError("[RayletClient] Raylet connection closed."); - } - if (type_field != static_cast(type)) { - return Status::TypeError( - std::string("[RayletClient] Raylet connection corrupted. ") + - "Expected message type: " + std::to_string(static_cast(type)) + - "; got message type: " + std::to_string(type_field) + - ". Check logs or dmesg for previous errors."); - } - return Status::OK(); + conn_ = ServerConnection::Create(std::move(socket)); } Status raylet::RayletConnection::WriteMessage(MessageType type, flatbuffers::FlatBufferBuilder *fbb) { std::unique_lock guard(write_mutex_); - int64_t cookie = RayConfig::instance().ray_cookie(); int64_t length = fbb ? fbb->GetSize() : 0; uint8_t *bytes = fbb ? fbb->GetBufferPointer() : nullptr; - int64_t type_field = static_cast(type); - auto io_error = Status::IOError("[RayletClient] Connection closed unexpectedly."); - int closed; - closed = write_bytes(conn_, &cookie, sizeof(cookie)); - if (closed) return io_error; - closed = write_bytes(conn_, &type_field, sizeof(type_field)); - if (closed) return io_error; - closed = write_bytes(conn_, &length, sizeof(length)); - if (closed) return io_error; - closed = write_bytes(conn_, bytes, length * sizeof(char)); - if (closed) return io_error; - return Status::OK(); + return conn_->WriteMessage(static_cast(type), length, bytes); } -Status raylet::RayletConnection::AtomicRequestReply( - MessageType request_type, MessageType reply_type, - std::unique_ptr &reply_message, flatbuffers::FlatBufferBuilder *fbb) { +Status raylet::RayletConnection::AtomicRequestReply(MessageType request_type, + MessageType reply_type, + std::vector *reply_message, + flatbuffers::FlatBufferBuilder *fbb) { std::unique_lock guard(mutex_); - auto status = WriteMessage(request_type, fbb); - if (!status.ok()) return status; - return ReadMessage(reply_type, reply_message); + RAY_RETURN_NOT_OK(WriteMessage(request_type, fbb)); + return conn_->ReadMessage(static_cast(reply_type), reply_message); } raylet::RayletClient::RayletClient( @@ -198,11 +96,11 @@ raylet::RayletClient::RayletClient( fbb.Finish(message); // Register the process ID with the raylet. // NOTE(swang): If raylet exits and we are registered as a worker, we will get killed. - std::unique_ptr reply; + std::vector reply; auto status = conn_->AtomicRequestReply(MessageType::RegisterClientRequest, - MessageType::RegisterClientReply, reply, &fbb); + MessageType::RegisterClientReply, &reply, &fbb); RAY_CHECK_OK_PREPEND(status, "[RayletClient] Unable to register worker with raylet."); - auto reply_message = flatbuffers::GetRoot(reply.get()); + auto reply_message = flatbuffers::GetRoot(reply.data()); *raylet_id = ClientID::FromBinary(reply_message->raylet_id()->str()); *port = reply_message->port(); @@ -215,6 +113,20 @@ raylet::RayletClient::RayletClient( } } +Status raylet::RayletClient::Disconnect() { + flatbuffers::FlatBufferBuilder fbb; + auto message = protocol::CreateDisconnectClient(fbb); + fbb.Finish(message); + auto status = conn_->WriteMessage(MessageType::IntentionalDisconnectClient, &fbb); + // Don't be too strict for disconnection errors. + // Just create logs and prevent it from crash. + if (!status.ok()) { + RAY_LOG(ERROR) << status.ToString() + << " [RayletClient] Failed to disconnect from raylet."; + } + return Status::OK(); +} + Status raylet::RayletClient::AnnounceWorkerPort(int port) { flatbuffers::FlatBufferBuilder fbb; auto message = protocol::CreateAnnounceWorkerPort(fbb, port); @@ -245,8 +157,7 @@ Status raylet::RayletClient::FetchOrReconstruct( fbb, object_ids_message, AddressesToFlatbuffer(fbb, owner_addresses), fetch_only, mark_worker_blocked, to_flatbuf(fbb, current_task_id)); fbb.Finish(message); - auto status = conn_->WriteMessage(MessageType::FetchOrReconstruct, &fbb); - return status; + return conn_->WriteMessage(MessageType::FetchOrReconstruct, &fbb); } Status raylet::RayletClient::NotifyUnblocked(const TaskID ¤t_task_id) { @@ -282,12 +193,11 @@ Status raylet::RayletClient::Wait(const std::vector &object_ids, num_returns, timeout_milliseconds, wait_local, mark_worker_blocked, to_flatbuf(fbb, current_task_id)); fbb.Finish(message); - std::unique_ptr reply; - auto status = conn_->AtomicRequestReply(MessageType::WaitRequest, - MessageType::WaitReply, reply, &fbb); - if (!status.ok()) return status; + std::vector reply; + RAY_RETURN_NOT_OK(conn_->AtomicRequestReply(MessageType::WaitRequest, + MessageType::WaitReply, &reply, &fbb)); // Parse the flatbuffer object. - auto reply_message = flatbuffers::GetRoot(reply.get()); + auto reply_message = flatbuffers::GetRoot(reply.data()); auto found = reply_message->found(); for (size_t i = 0; i < found->size(); i++) { ObjectID object_id = ObjectID::FromBinary(found->Get(i)->str()); @@ -324,7 +234,6 @@ Status raylet::RayletClient::PushError(const JobID &job_id, const std::string &t fbb, to_flatbuf(fbb, job_id), fbb.CreateString(type), fbb.CreateString(error_message), timestamp); fbb.Finish(message); - return conn_->WriteMessage(MessageType::PushErrorRequest, &fbb); } @@ -348,9 +257,7 @@ Status raylet::RayletClient::FreeObjects(const std::vector &object_ids auto message = protocol::CreateFreeObjectsRequest( fbb, local_only, delete_creating_tasks, to_flatbuf(fbb, object_ids)); fbb.Finish(message); - - auto status = conn_->WriteMessage(MessageType::FreeObjectsInObjectStoreRequest, &fbb); - return status; + return conn_->WriteMessage(MessageType::FreeObjectsInObjectStoreRequest, &fbb); } Status raylet::RayletClient::PrepareActorCheckpoint(const ActorID &actor_id, @@ -360,13 +267,12 @@ Status raylet::RayletClient::PrepareActorCheckpoint(const ActorID &actor_id, protocol::CreatePrepareActorCheckpointRequest(fbb, to_flatbuf(fbb, actor_id)); fbb.Finish(message); - std::unique_ptr reply; - auto status = - conn_->AtomicRequestReply(MessageType::PrepareActorCheckpointRequest, - MessageType::PrepareActorCheckpointReply, reply, &fbb); - if (!status.ok()) return status; + std::vector reply; + RAY_RETURN_NOT_OK(conn_->AtomicRequestReply(MessageType::PrepareActorCheckpointRequest, + MessageType::PrepareActorCheckpointReply, + &reply, &fbb)); auto reply_message = - flatbuffers::GetRoot(reply.get()); + flatbuffers::GetRoot(reply.data()); *checkpoint_id = ActorCheckpointID::FromBinary(reply_message->checkpoint_id()->str()); return Status::OK(); } @@ -377,7 +283,6 @@ Status raylet::RayletClient::NotifyActorResumedFromCheckpoint( auto message = protocol::CreateNotifyActorResumedFromCheckpoint( fbb, to_flatbuf(fbb, actor_id), to_flatbuf(fbb, checkpoint_id)); fbb.Finish(message); - return conn_->WriteMessage(MessageType::NotifyActorResumedFromCheckpoint, &fbb); } diff --git a/src/ray/raylet/raylet_client.h b/src/ray/raylet/raylet_client.h index 278f8b017..a5abf2ece 100644 --- a/src/ray/raylet/raylet_client.h +++ b/src/ray/raylet/raylet_client.h @@ -15,14 +15,13 @@ #pragma once #include -#include -#include #include #include #include #include "ray/common/bundle_spec.h" +#include "ray/common/client_connection.h" #include "ray/common/status.h" #include "ray/common/task/task_spec.h" #include "ray/rpc/node_manager/node_manager_client.h" @@ -45,9 +44,6 @@ using WaitResultPair = std::pair, std::vector>; namespace ray { -typedef boost::asio::generic::stream_protocol local_stream_protocol; -typedef boost::asio::basic_stream_socket local_stream_socket; - /// Interface for pinning objects. Abstract for testing. class PinObjectsInterface { public: @@ -133,25 +129,16 @@ class RayletConnection { RayletConnection(boost::asio::io_service &io_service, const std::string &raylet_socket, int num_retries, int64_t timeout); - /// 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. - /// - /// \return ray::Status. - ray::Status Disconnect(); - - ray::Status ReadMessage(MessageType type, std::unique_ptr &message); - ray::Status WriteMessage(MessageType type, flatbuffers::FlatBufferBuilder *fbb = nullptr); ray::Status AtomicRequestReply(MessageType request_type, MessageType reply_type, - std::unique_ptr &reply_message, + std::vector *reply_message, flatbuffers::FlatBufferBuilder *fbb = nullptr); private: - /// The Unix domain socket that connects to raylet. - local_stream_socket conn_; + /// The connection to raylet. + std::shared_ptr conn_; /// A mutex to protect stateful operations of the raylet client. std::mutex mutex_; /// A mutex to protect write operations of the raylet client. @@ -190,7 +177,12 @@ class RayletClient : public PinObjectsInterface, /// \param grpc_client gRPC client to the raylet. RayletClient(std::shared_ptr grpc_client); - ray::Status Disconnect() { return conn_->Disconnect(); }; + /// 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. + /// + /// \return ray::Status. + ray::Status Disconnect(); /// Tell the raylet which port this worker's gRPC server is listening on. ///