[Core] Simplify Raylet Client (#9420)

This commit is contained in:
Siyuan (Ryans) Zhuang
2020-07-12 12:42:54 -07:00
committed by GitHub
parent 8c985dc797
commit 381c242f6b
5 changed files with 90 additions and 161 deletions
+31 -2
View File
@@ -14,21 +14,50 @@
#include "client_connection.h"
#include <stdio.h>
#include <boost/asio/buffer.hpp>
#include <boost/asio/generic/stream_protocol.hpp>
#include <boost/asio/placeholders.hpp>
#include <boost/asio/read.hpp>
#include <boost/asio/write.hpp>
#include <boost/bind.hpp>
#include <chrono>
#include <sstream>
#include <thread>
#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> ServerConnection::Create(local_stream_socket &&socket) {
std::shared_ptr<ServerConnection> self(new ServerConnection(std::move(socket)));
return self;
+7 -4
View File
@@ -14,23 +14,26 @@
#pragma once
#include <deque>
#include <memory>
#include <boost/asio/basic_stream_socket.hpp>
#include <boost/asio/buffer.hpp>
#include <boost/asio/error.hpp>
#include <boost/asio/generic/stream_protocol.hpp>
#include <boost/enable_shared_from_this.hpp>
#include <deque>
#include <memory>
#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_protocol> 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
@@ -113,7 +113,7 @@ Status CoreWorkerPlasmaStoreProvider::Create(const std::shared_ptr<Buffer> &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;
+41 -136
View File
@@ -14,17 +14,7 @@
#include "raylet_client.h"
#include <inttypes.h>
#include <netdb.h>
#include <netinet/in.h>
#include <stdarg.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <sys/ioctl.h>
#include <sys/socket.h>
#include <sys/types.h>
#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<uint8_t[]> &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<uint8_t[]>(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<int64_t>(MessageType::DisconnectClient);
length = 0;
}
if (type_field == static_cast<int64_t>(MessageType::DisconnectClient)) {
return Status::IOError("[RayletClient] Raylet connection closed.");
}
if (type_field != static_cast<int64_t>(type)) {
return Status::TypeError(
std::string("[RayletClient] Raylet connection corrupted. ") +
"Expected message type: " + std::to_string(static_cast<int64_t>(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<std::mutex> 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<int64_t>(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<int64_t>(type), length, bytes);
}
Status raylet::RayletConnection::AtomicRequestReply(
MessageType request_type, MessageType reply_type,
std::unique_ptr<uint8_t[]> &reply_message, flatbuffers::FlatBufferBuilder *fbb) {
Status raylet::RayletConnection::AtomicRequestReply(MessageType request_type,
MessageType reply_type,
std::vector<uint8_t> *reply_message,
flatbuffers::FlatBufferBuilder *fbb) {
std::unique_lock<std::mutex> 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<int64_t>(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<uint8_t[]> reply;
std::vector<uint8_t> 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<protocol::RegisterClientReply>(reply.get());
auto reply_message = flatbuffers::GetRoot<protocol::RegisterClientReply>(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 &current_task_id) {
@@ -282,12 +193,11 @@ Status raylet::RayletClient::Wait(const std::vector<ObjectID> &object_ids,
num_returns, timeout_milliseconds, wait_local, mark_worker_blocked,
to_flatbuf(fbb, current_task_id));
fbb.Finish(message);
std::unique_ptr<uint8_t[]> reply;
auto status = conn_->AtomicRequestReply(MessageType::WaitRequest,
MessageType::WaitReply, reply, &fbb);
if (!status.ok()) return status;
std::vector<uint8_t> reply;
RAY_RETURN_NOT_OK(conn_->AtomicRequestReply(MessageType::WaitRequest,
MessageType::WaitReply, &reply, &fbb));
// Parse the flatbuffer object.
auto reply_message = flatbuffers::GetRoot<protocol::WaitReply>(reply.get());
auto reply_message = flatbuffers::GetRoot<protocol::WaitReply>(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<ObjectID> &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<uint8_t[]> reply;
auto status =
conn_->AtomicRequestReply(MessageType::PrepareActorCheckpointRequest,
MessageType::PrepareActorCheckpointReply, reply, &fbb);
if (!status.ok()) return status;
std::vector<uint8_t> reply;
RAY_RETURN_NOT_OK(conn_->AtomicRequestReply(MessageType::PrepareActorCheckpointRequest,
MessageType::PrepareActorCheckpointReply,
&reply, &fbb));
auto reply_message =
flatbuffers::GetRoot<protocol::PrepareActorCheckpointReply>(reply.get());
flatbuffers::GetRoot<protocol::PrepareActorCheckpointReply>(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);
}
+10 -18
View File
@@ -15,14 +15,13 @@
#pragma once
#include <ray/protobuf/gcs.pb.h>
#include <unistd.h>
#include <boost/asio/detail/socket_holder.hpp>
#include <mutex>
#include <unordered_map>
#include <vector>
#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<ObjectID>, std::vector<ObjectID>>;
namespace ray {
typedef boost::asio::generic::stream_protocol local_stream_protocol;
typedef boost::asio::basic_stream_socket<local_stream_protocol> 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<uint8_t[]> &message);
ray::Status WriteMessage(MessageType type,
flatbuffers::FlatBufferBuilder *fbb = nullptr);
ray::Status AtomicRequestReply(MessageType request_type, MessageType reply_type,
std::unique_ptr<uint8_t[]> &reply_message,
std::vector<uint8_t> *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<ServerConnection> 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<ray::rpc::NodeManagerWorkerClient> 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.
///