From 54d5969ceacabad2741af278ef9ed4db0af66835 Mon Sep 17 00:00:00 2001 From: Zhijun Fu <37800433+zhijunfu@users.noreply.github.com> Date: Thu, 4 Jul 2019 20:16:42 +0800 Subject: [PATCH] [grpc] Add grpc server to worker (#5054) * refactor grpc server * format * change GetTask() to PushTask() * change PushTask to AssignTask * format * update * fix test * format * Update src/ray/rpc/worker_client.h Co-Authored-By: Hao Chen * Update BUILD.bazel * Update src/ray/core_worker/task_execution.cc Co-Authored-By: Stephanie Wang * update * format * address comments * format * Update src/ray/rpc/worker/worker_server.h Co-Authored-By: Stephanie Wang * Update src/ray/protobuf/worker.proto Co-Authored-By: Stephanie Wang * format * fix * format --- BUILD.bazel | 36 ++++++++++ src/ray/core_worker/core_worker.cc | 28 +++++--- src/ray/core_worker/core_worker.h | 10 ++- src/ray/core_worker/object_interface.cc | 6 +- src/ray/core_worker/object_interface.h | 8 ++- .../store_provider/plasma_store_provider.cc | 12 ++-- .../store_provider/plasma_store_provider.h | 4 +- .../store_provider/store_provider.h | 3 +- src/ray/core_worker/task_execution.cc | 19 ++++-- src/ray/core_worker/task_execution.h | 18 ++++- src/ray/core_worker/task_interface.cc | 4 +- src/ray/core_worker/task_interface.h | 3 +- .../core_worker/transport/raylet_transport.cc | 30 ++++++-- .../core_worker/transport/raylet_transport.h | 33 +++++++-- src/ray/core_worker/transport/transport.h | 2 + src/ray/protobuf/worker.proto | 23 +++++++ src/ray/raylet/format/node_manager.fbs | 4 ++ src/ray/raylet/node_manager.cc | 4 +- src/ray/raylet/raylet_client.cc | 11 ++- src/ray/raylet/raylet_client.h | 4 +- src/ray/raylet/worker.cc | 5 +- src/ray/raylet/worker.h | 6 +- src/ray/raylet/worker_pool.cc | 3 +- src/ray/raylet/worker_pool_test.cc | 2 +- src/ray/rpc/client_call.h | 3 + src/ray/rpc/grpc_server.cc | 4 +- src/ray/rpc/grpc_server.h | 2 +- src/ray/rpc/worker/worker_client.h | 58 ++++++++++++++++ src/ray/rpc/worker/worker_server.h | 68 +++++++++++++++++++ 29 files changed, 351 insertions(+), 62 deletions(-) create mode 100644 src/ray/protobuf/worker.proto create mode 100644 src/ray/rpc/worker/worker_client.h create mode 100644 src/ray/rpc/worker/worker_server.h diff --git a/BUILD.bazel b/BUILD.bazel index 054c57bb4..f2d458e31 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -47,6 +47,16 @@ cc_proto_library( deps = ["object_manager_proto"], ) +proto_library( + name = "worker_proto", + srcs = ["src/ray/protobuf/worker.proto"], +) + +cc_proto_library( + name = "worker_cc_proto", + deps = ["worker_proto"], +) + proto_library( name = "core_worker_proto", srcs = ["src/ray/protobuf/core_worker.proto"], @@ -127,6 +137,30 @@ cc_library( ], ) +# worker gRPC lib. +cc_grpc_library( + name = "worker_cc_grpc", + srcs = [":worker_proto"], + grpc_only = True, + deps = [":worker_cc_proto"], +) + +# worker server and client. +cc_library( + name = "worker_rpc", + hdrs = glob([ + "src/ray/rpc/worker/*.h", + ]), + copts = COPTS, + deps = [ + ":grpc_common_lib", + ":ray_common", + ":worker_cc_grpc", + "@boost//:asio", + "@com_github_grpc_grpc//:grpc++", + ], +) + # === End of rpc definitions === cc_binary( @@ -216,6 +250,7 @@ cc_library( ":ray_common", ":ray_util", ":stats_lib", + ":worker_rpc", "@boost//:asio", "@com_github_jupp0r_prometheus_cpp//pull", "@com_google_absl//absl/base:core_headers", @@ -253,6 +288,7 @@ cc_library( ":ray_common", ":ray_util", ":raylet_lib", + ":worker_rpc", ], ) diff --git a/src/ray/core_worker/core_worker.cc b/src/ray/core_worker/core_worker.cc index 57e35d4fd..fe6ed290d 100644 --- a/src/ray/core_worker/core_worker.cc +++ b/src/ray/core_worker/core_worker.cc @@ -10,15 +10,23 @@ CoreWorker::CoreWorker(const enum WorkerType worker_type, const ::Language langu language_(language), raylet_socket_(raylet_socket), worker_context_(worker_type, job_id), - // TODO(zhijunfu): currently RayletClient would crash in its constructor - // if it cannot connect to Raylet after a number of retries, this needs - // to be changed so that the worker (java/python .etc) can retrieve and - // handle the error instead of crashing. - raylet_client_(raylet_socket_, - ClientID::FromBinary(worker_context_.GetWorkerID().Binary()), - (worker_type_ == ray::WorkerType::WORKER), - worker_context_.GetCurrentJobID(), language_), task_interface_(worker_context_, raylet_client_), - object_interface_(worker_context_, raylet_client_, store_socket), - task_execution_interface_(worker_context_, raylet_client_, object_interface_) {} + object_interface_(worker_context_, raylet_client_, store_socket) { + int rpc_server_port = 0; + if (worker_type_ == ray::WorkerType::WORKER) { + task_execution_interface_ = std::unique_ptr( + new CoreWorkerTaskExecutionInterface(worker_context_, raylet_client_, + object_interface_)); + rpc_server_port = task_execution_interface_->worker_server_.GetPort(); + } + // TODO(zhijunfu): currently RayletClient would crash in its constructor if it cannot + // connect to Raylet after a number of retries, this can be changed later + // so that the worker (java/python .etc) can retrieve and handle the error + // instead of crashing. + raylet_client_ = std::unique_ptr(new RayletClient( + raylet_socket_, ClientID::FromBinary(worker_context_.GetWorkerID().Binary()), + (worker_type_ == ray::WorkerType::WORKER), worker_context_.GetCurrentJobID(), + language_, rpc_server_port)); +} + } // namespace ray diff --git a/src/ray/core_worker/core_worker.h b/src/ray/core_worker/core_worker.h index 3f567f0af..8a3347b2a 100644 --- a/src/ray/core_worker/core_worker.h +++ b/src/ray/core_worker/core_worker.h @@ -43,7 +43,10 @@ class CoreWorker { /// Return the `CoreWorkerTaskExecutionInterface` that contains methods related to /// task execution. - CoreWorkerTaskExecutionInterface &Execution() { return task_execution_interface_; } + CoreWorkerTaskExecutionInterface &Execution() { + RAY_CHECK(task_execution_interface_ != nullptr); + return *task_execution_interface_; + } private: /// Type of this worker. @@ -59,7 +62,7 @@ class CoreWorker { WorkerContext worker_context_; /// Raylet client. - RayletClient raylet_client_; + std::unique_ptr raylet_client_; /// The `CoreWorkerTaskInterface` instance. CoreWorkerTaskInterface task_interface_; @@ -68,7 +71,8 @@ class CoreWorker { CoreWorkerObjectInterface object_interface_; /// The `CoreWorkerTaskExecutionInterface` instance. - CoreWorkerTaskExecutionInterface task_execution_interface_; + /// This is only available if it's not a driver. + std::unique_ptr task_execution_interface_; }; } // namespace ray diff --git a/src/ray/core_worker/object_interface.cc b/src/ray/core_worker/object_interface.cc index 53836431f..8a4eb4642 100644 --- a/src/ray/core_worker/object_interface.cc +++ b/src/ray/core_worker/object_interface.cc @@ -4,9 +4,9 @@ namespace ray { -CoreWorkerObjectInterface::CoreWorkerObjectInterface(WorkerContext &worker_context, - RayletClient &raylet_client, - const std::string &store_socket) +CoreWorkerObjectInterface::CoreWorkerObjectInterface( + WorkerContext &worker_context, std::unique_ptr &raylet_client, + const std::string &store_socket) : worker_context_(worker_context), raylet_client_(raylet_client) { store_providers_.emplace( static_cast(StoreProviderType::PLASMA), diff --git a/src/ray/core_worker/object_interface.h b/src/ray/core_worker/object_interface.h index 7d3d42a2d..f1d66e803 100644 --- a/src/ray/core_worker/object_interface.h +++ b/src/ray/core_worker/object_interface.h @@ -17,7 +17,8 @@ class CoreWorkerStoreProvider; /// The interface that contains all `CoreWorker` methods that are related to object store. class CoreWorkerObjectInterface { public: - CoreWorkerObjectInterface(WorkerContext &worker_context, RayletClient &raylet_client, + CoreWorkerObjectInterface(WorkerContext &worker_context, + std::unique_ptr &raylet_client, const std::string &store_socket); /// Put an object into object store. @@ -59,7 +60,8 @@ class CoreWorkerObjectInterface { /// \param[in] local_only Whether only delete the objects in local node, or all nodes in /// the cluster. /// \param[in] delete_creating_tasks Whether also delete the tasks that - /// created these objects. \return Status. + /// created these objects. + /// \return Status. Status Delete(const std::vector &object_ids, bool local_only, bool delete_creating_tasks); @@ -67,7 +69,7 @@ class CoreWorkerObjectInterface { /// Reference to the parent CoreWorker's context. WorkerContext &worker_context_; /// Reference to the parent CoreWorker's raylet client. - RayletClient &raylet_client_; + std::unique_ptr &raylet_client_; /// All the store providers supported. std::unordered_map> store_providers_; 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 652e64c02..53c330dc0 100644 --- a/src/ray/core_worker/store_provider/plasma_store_provider.cc +++ b/src/ray/core_worker/store_provider/plasma_store_provider.cc @@ -7,7 +7,7 @@ namespace ray { CoreWorkerPlasmaStoreProvider::CoreWorkerPlasmaStoreProvider( - const std::string &store_socket, RayletClient &raylet_client) + const std::string &store_socket, std::unique_ptr &raylet_client) : raylet_client_(raylet_client) { auto status = store_client_.Connect(store_socket); if (!status.ok()) { @@ -70,7 +70,7 @@ Status CoreWorkerPlasmaStoreProvider::Get( } // TODO(zhijunfu): can call `fetchOrReconstruct` in batches as an optimization. - RAY_CHECK_OK(raylet_client_.FetchOrReconstruct(unready_ids, fetch_only, task_id)); + RAY_CHECK_OK(raylet_client_->FetchOrReconstruct(unready_ids, fetch_only, task_id)); // Get the objects from the object store, and parse the result. int64_t get_timeout; @@ -109,7 +109,7 @@ Status CoreWorkerPlasmaStoreProvider::Get( } if (was_blocked) { - RAY_CHECK_OK(raylet_client_.NotifyUnblocked(task_id)); + RAY_CHECK_OK(raylet_client_->NotifyUnblocked(task_id)); } return Status::OK(); @@ -120,8 +120,8 @@ Status CoreWorkerPlasmaStoreProvider::Wait(const std::vector &object_i const TaskID &task_id, std::vector *results) { WaitResultPair result_pair; - auto status = raylet_client_.Wait(object_ids, num_objects, timeout_ms, false, task_id, - &result_pair); + auto status = raylet_client_->Wait(object_ids, num_objects, timeout_ms, false, task_id, + &result_pair); std::unordered_set ready_ids; for (const auto &entry : result_pair.first) { ready_ids.insert(entry); @@ -141,7 +141,7 @@ Status CoreWorkerPlasmaStoreProvider::Wait(const std::vector &object_i Status CoreWorkerPlasmaStoreProvider::Delete(const std::vector &object_ids, bool local_only, bool delete_creating_tasks) { - return raylet_client_.FreeObjects(object_ids, local_only, delete_creating_tasks); + return raylet_client_->FreeObjects(object_ids, local_only, delete_creating_tasks); } } // namespace ray diff --git a/src/ray/core_worker/store_provider/plasma_store_provider.h b/src/ray/core_worker/store_provider/plasma_store_provider.h index db9c39639..9aa2f914a 100644 --- a/src/ray/core_worker/store_provider/plasma_store_provider.h +++ b/src/ray/core_worker/store_provider/plasma_store_provider.h @@ -18,7 +18,7 @@ class CoreWorker; class CoreWorkerPlasmaStoreProvider : public CoreWorkerStoreProvider { public: CoreWorkerPlasmaStoreProvider(const std::string &store_socket, - RayletClient &raylet_client); + std::unique_ptr &raylet_client); /// Put an object with specified ID into object store. /// @@ -67,7 +67,7 @@ class CoreWorkerPlasmaStoreProvider : public CoreWorkerStoreProvider { std::mutex store_client_mutex_; /// Raylet client. - RayletClient &raylet_client_; + std::unique_ptr &raylet_client_; }; } // namespace ray diff --git a/src/ray/core_worker/store_provider/store_provider.h b/src/ray/core_worker/store_provider/store_provider.h index 25b932437..fdba0d887 100644 --- a/src/ray/core_worker/store_provider/store_provider.h +++ b/src/ray/core_worker/store_provider/store_provider.h @@ -77,7 +77,8 @@ class CoreWorkerStoreProvider { /// \param[in] local_only Whether only delete the objects in local node, or all nodes in /// the cluster. /// \param[in] delete_creating_tasks Whether also delete the tasks that - /// created these objects. \return Status. + /// created these objects. + /// \return Status. virtual Status Delete(const std::vector &object_ids, bool local_only, bool delete_creating_tasks) = 0; }; diff --git a/src/ray/core_worker/task_execution.cc b/src/ray/core_worker/task_execution.cc index ce7c88986..bbcfb39ef 100644 --- a/src/ray/core_worker/task_execution.cc +++ b/src/ray/core_worker/task_execution.cc @@ -6,19 +6,26 @@ namespace ray { CoreWorkerTaskExecutionInterface::CoreWorkerTaskExecutionInterface( - WorkerContext &worker_context, RayletClient &raylet_client, + WorkerContext &worker_context, std::unique_ptr &raylet_client, CoreWorkerObjectInterface &object_interface) - : worker_context_(worker_context), object_interface_(object_interface) { - task_receivers.emplace(static_cast(TaskTransportType::RAYLET), - std::unique_ptr( - new CoreWorkerRayletTaskReceiver(raylet_client))); + : worker_context_(worker_context), + object_interface_(object_interface), + worker_server_("Worker", 0 /* let grpc choose port */), + main_work_(main_service_) { + task_receivers_.emplace( + static_cast(TaskTransportType::RAYLET), + std::unique_ptr(new CoreWorkerRayletTaskReceiver( + raylet_client, main_service_, worker_server_))); + + // Start RPC server after all the task receivers are properly initialized. + worker_server_.Run(); } Status CoreWorkerTaskExecutionInterface::Run(const TaskExecutor &executor) { while (true) { std::vector tasks; auto status = - task_receivers[static_cast(TaskTransportType::RAYLET)]->GetTasks(&tasks); + task_receivers_[static_cast(TaskTransportType::RAYLET)]->GetTasks(&tasks); if (!status.ok()) { RAY_LOG(ERROR) << "Getting task failed with error: " << ray::Status::IOError(status.message()); diff --git a/src/ray/core_worker/task_execution.h b/src/ray/core_worker/task_execution.h index 22491aa03..27f43e9fb 100644 --- a/src/ray/core_worker/task_execution.h +++ b/src/ray/core_worker/task_execution.h @@ -7,6 +7,9 @@ #include "ray/core_worker/context.h" #include "ray/core_worker/object_interface.h" #include "ray/core_worker/transport/transport.h" +#include "ray/rpc/client_call.h" +#include "ray/rpc/worker/worker_client.h" +#include "ray/rpc/worker/worker_server.h" namespace ray { @@ -21,7 +24,7 @@ class TaskSpecification; class CoreWorkerTaskExecutionInterface { public: CoreWorkerTaskExecutionInterface(WorkerContext &worker_context, - RayletClient &raylet_client, + std::unique_ptr &raylet_client, CoreWorkerObjectInterface &object_interface); /// The callback provided app-language workers that executes tasks. @@ -56,7 +59,18 @@ class CoreWorkerTaskExecutionInterface { CoreWorkerObjectInterface &object_interface_; /// All the task task receivers supported. - std::unordered_map> task_receivers; + std::unordered_map> task_receivers_; + + /// The RPC server. + rpc::GrpcServer worker_server_; + + /// Event loop where tasks are processed. + boost::asio::io_service main_service_; + + /// The asio work to keep main_service_ alive. + boost::asio::io_service::work main_work_; + + friend class CoreWorker; }; } // namespace ray diff --git a/src/ray/core_worker/task_interface.cc b/src/ray/core_worker/task_interface.cc index c1510726d..e85ce28c9 100644 --- a/src/ray/core_worker/task_interface.cc +++ b/src/ray/core_worker/task_interface.cc @@ -92,8 +92,8 @@ std::vector ActorHandle::NewActorHandles() const { void ActorHandle::ClearNewActorHandles() { new_actor_handles_.clear(); } -CoreWorkerTaskInterface::CoreWorkerTaskInterface(WorkerContext &worker_context, - RayletClient &raylet_client) +CoreWorkerTaskInterface::CoreWorkerTaskInterface( + WorkerContext &worker_context, std::unique_ptr &raylet_client) : worker_context_(worker_context) { task_submitters_.emplace(static_cast(TaskTransportType::RAYLET), std::unique_ptr( diff --git a/src/ray/core_worker/task_interface.h b/src/ray/core_worker/task_interface.h index 75bd9cc64..aa5876c00 100644 --- a/src/ray/core_worker/task_interface.h +++ b/src/ray/core_worker/task_interface.h @@ -111,7 +111,8 @@ class ActorHandle { /// submission. class CoreWorkerTaskInterface { public: - CoreWorkerTaskInterface(WorkerContext &worker_context, RayletClient &raylet_client); + CoreWorkerTaskInterface(WorkerContext &worker_context, + std::unique_ptr &raylet_client); /// Submit a normal task. /// diff --git a/src/ray/core_worker/transport/raylet_transport.cc b/src/ray/core_worker/transport/raylet_transport.cc index 14906acfe..b63a5ef89 100644 --- a/src/ray/core_worker/transport/raylet_transport.cc +++ b/src/ray/core_worker/transport/raylet_transport.cc @@ -1,21 +1,20 @@ #include "ray/core_worker/transport/raylet_transport.h" +#include "ray/raylet/task.h" namespace ray { -CoreWorkerRayletTaskSubmitter::CoreWorkerRayletTaskSubmitter(RayletClient &raylet_client) +CoreWorkerRayletTaskSubmitter::CoreWorkerRayletTaskSubmitter( + std::unique_ptr &raylet_client) : raylet_client_(raylet_client) {} Status CoreWorkerRayletTaskSubmitter::SubmitTask(const TaskSpec &task) { - return raylet_client_.SubmitTask(task.GetDependencies(), task.GetTaskSpecification()); + return raylet_client_->SubmitTask(task.GetDependencies(), task.GetTaskSpecification()); } -CoreWorkerRayletTaskReceiver::CoreWorkerRayletTaskReceiver(RayletClient &raylet_client) - : raylet_client_(raylet_client) {} - Status CoreWorkerRayletTaskReceiver::GetTasks(std::vector *tasks) { std::unique_ptr task_spec; - auto status = raylet_client_.GetTask(&task_spec); + auto status = raylet_client_->GetTask(&task_spec); if (!status.ok()) { RAY_LOG(ERROR) << "Get task from raylet failed with error: " << ray::Status::IOError(status.message()); @@ -29,4 +28,23 @@ Status CoreWorkerRayletTaskReceiver::GetTasks(std::vector *tasks) { return Status::OK(); } +CoreWorkerRayletTaskReceiver::CoreWorkerRayletTaskReceiver( + std::unique_ptr &raylet_client, boost::asio::io_service &io_service, + rpc::GrpcServer &server) + : raylet_client_(raylet_client), task_service_(io_service, *this) { + server.RegisterService(task_service_); +} + +void CoreWorkerRayletTaskReceiver::HandleAssignTask( + const rpc::AssignTaskRequest &request, rpc::AssignTaskReply *reply, + rpc::RequestDoneCallback done_callback) { + const std::string &task_message = request.task_spec(); + const raylet::Task task(*flatbuffers::GetRoot( + reinterpret_cast(task_message.data()))); + const auto &spec = task.GetTaskSpecification(); + + auto status = task_handler_(spec); + done_callback(status); +} + } // namespace ray diff --git a/src/ray/core_worker/transport/raylet_transport.h b/src/ray/core_worker/transport/raylet_transport.h index 03bf82f29..cdf8deebb 100644 --- a/src/ray/core_worker/transport/raylet_transport.h +++ b/src/ray/core_worker/transport/raylet_transport.h @@ -5,6 +5,7 @@ #include "ray/core_worker/transport/transport.h" #include "ray/raylet/raylet_client.h" +#include "ray/rpc/worker/worker_server.h" namespace ray { @@ -14,7 +15,7 @@ namespace ray { class CoreWorkerRayletTaskSubmitter : public CoreWorkerTaskSubmitter { public: - CoreWorkerRayletTaskSubmitter(RayletClient &raylet_client); + CoreWorkerRayletTaskSubmitter(std::unique_ptr &raylet_client); /// Submit a task for execution to raylet. /// @@ -24,19 +25,41 @@ class CoreWorkerRayletTaskSubmitter : public CoreWorkerTaskSubmitter { private: /// Raylet client. - RayletClient &raylet_client_; + std::unique_ptr &raylet_client_; }; -class CoreWorkerRayletTaskReceiver : public CoreWorkerTaskReceiver { +class CoreWorkerRayletTaskReceiver : public CoreWorkerTaskReceiver, + public rpc::WorkerTaskHandler { public: - CoreWorkerRayletTaskReceiver(RayletClient &raylet_client); + CoreWorkerRayletTaskReceiver(std::unique_ptr &raylet_client, + boost::asio::io_service &io_service, + rpc::GrpcServer &server); // Get tasks for execution from raylet. virtual Status GetTasks(std::vector *tasks) override; + /// TODO(zhijunfu): This is currently unused. Later when we migrate from worker "get + /// task" to raylet "assign task", this method will be used and the `GetTask` above will + /// be removed. + /// + /// Handle a `AssignTask` request. + /// The implementation can handle this request asynchronously. When hanling is done, the + /// `done_callback` should be called. + /// + /// \param[in] request The request message. + /// \param[out] reply The reply message. + /// \param[in] done_callback The callback to be called when the request is done. + void HandleAssignTask(const rpc::AssignTaskRequest &request, + rpc::AssignTaskReply *reply, + rpc::RequestDoneCallback done_callback) override; + private: /// Raylet client. - RayletClient &raylet_client_; + std::unique_ptr &raylet_client_; + /// The callback function to process a task. + TaskHandler task_handler_; + /// The rpc service for `WorkerTaskService`. + rpc::WorkerTaskGrpcService task_service_; }; } // namespace ray diff --git a/src/ray/core_worker/transport/transport.h b/src/ray/core_worker/transport/transport.h index 44be74b98..8433b3b17 100644 --- a/src/ray/core_worker/transport/transport.h +++ b/src/ray/core_worker/transport/transport.h @@ -32,6 +32,8 @@ class CoreWorkerTaskSubmitter { /// This class receives tasks for execution. class CoreWorkerTaskReceiver { public: + using TaskHandler = std::function; + // Get tasks for execution. virtual Status GetTasks(std::vector *tasks) = 0; }; diff --git a/src/ray/protobuf/worker.proto b/src/ray/protobuf/worker.proto new file mode 100644 index 000000000..b63782925 --- /dev/null +++ b/src/ray/protobuf/worker.proto @@ -0,0 +1,23 @@ +syntax = "proto3"; + +package ray.rpc; + +message AssignTaskRequest { + // The ID of the task to be pushed. + bytes task_id = 1; + // The task to be pushed. This should include task_id. + // TODO(hchen): Currently, `task_spec` are represented as + // flatbutters-serialized bytes. This is because the flatbuffers-defined Task data + // structure is being used in many places. We should move Task and all related data + // structures to protobuf. + bytes task_spec = 2; +} + +message AssignTaskReply { +} + +// Service for worker. +service WorkerTaskService { + // Push a task to a worker. + rpc AssignTask(AssignTaskRequest) returns (AssignTaskReply); +} diff --git a/src/ray/raylet/format/node_manager.fbs b/src/ray/raylet/format/node_manager.fbs index d5f5f2c7c..be2a6ead5 100644 --- a/src/ray/raylet/format/node_manager.fbs +++ b/src/ray/raylet/format/node_manager.fbs @@ -138,6 +138,10 @@ table RegisterClientRequest { job_id: string; // Language of this worker. language: Language; + // Port that this worker is listening on. + // If port > 0, then worker will listen to this port and wait for + // raylet to push tasks, instead of invoking GetTask(). + port: int; } table RegisterClientReply { diff --git a/src/ray/raylet/node_manager.cc b/src/ray/raylet/node_manager.cc index 35da34fe5..fa7de5dad 100644 --- a/src/ray/raylet/node_manager.cc +++ b/src/ray/raylet/node_manager.cc @@ -845,8 +845,8 @@ void NodeManager::ProcessRegisterClientRequestMessage( const std::shared_ptr &client, const uint8_t *message_data) { auto message = flatbuffers::GetRoot(message_data); client->SetClientID(from_flatbuf(*message->worker_id())); - auto worker = - std::make_shared(message->worker_pid(), message->language(), client); + auto worker = std::make_shared(message->worker_pid(), message->language(), + message->port(), client); if (message->is_worker()) { // Register the new worker. worker_pool_.RegisterWorker(std::move(worker)); diff --git a/src/ray/raylet/raylet_client.cc b/src/ray/raylet/raylet_client.cc index 2c4cbf60c..da1b09e27 100644 --- a/src/ray/raylet/raylet_client.cc +++ b/src/ray/raylet/raylet_client.cc @@ -202,15 +202,20 @@ ray::Status RayletConnection::AtomicRequestReply( } RayletClient::RayletClient(const std::string &raylet_socket, const ClientID &client_id, - bool is_worker, const JobID &job_id, const Language &language) - : client_id_(client_id), is_worker_(is_worker), job_id_(job_id), language_(language) { + bool is_worker, const JobID &job_id, const Language &language, + int port) + : client_id_(client_id), + is_worker_(is_worker), + job_id_(job_id), + language_(language), + port_(port) { // For C++14, we could use std::make_unique conn_ = std::unique_ptr(new RayletConnection(raylet_socket, -1, -1)); flatbuffers::FlatBufferBuilder fbb; auto message = ray::protocol::CreateRegisterClientRequest( fbb, is_worker, to_flatbuf(fbb, client_id), getpid(), to_flatbuf(fbb, job_id), - language); + language, port); 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. diff --git a/src/ray/raylet/raylet_client.h b/src/ray/raylet/raylet_client.h index 53e880452..5f887cc93 100644 --- a/src/ray/raylet/raylet_client.h +++ b/src/ray/raylet/raylet_client.h @@ -69,7 +69,8 @@ class RayletClient { /// \param job_id The ID of the driver. This is non-nil if the client is a driver. /// \return The connection information. RayletClient(const std::string &raylet_socket, const ClientID &client_id, - bool is_worker, const JobID &job_id, const Language &language); + bool is_worker, const JobID &job_id, const Language &language, + int port = -1); ray::Status Disconnect() { return conn_->Disconnect(); }; @@ -188,6 +189,7 @@ class RayletClient { const bool is_worker_; const JobID job_id_; const Language language_; + const int port_; /// A map from resource name to the resource IDs that are currently reserved /// for this worker. Each pair consists of the resource ID and the fraction /// of that resource allocated for this worker. diff --git a/src/ray/raylet/worker.cc b/src/ray/raylet/worker.cc index 359754340..d3d9833a3 100644 --- a/src/ray/raylet/worker.cc +++ b/src/ray/raylet/worker.cc @@ -10,10 +10,11 @@ namespace ray { namespace raylet { /// A constructor responsible for initializing the state of a worker. -Worker::Worker(pid_t pid, const Language &language, +Worker::Worker(pid_t pid, const Language &language, int port, std::shared_ptr connection) : pid_(pid), language_(language), + port_(port), connection_(connection), dead_(false), blocked_(false) {} @@ -32,6 +33,8 @@ pid_t Worker::Pid() const { return pid_; } Language Worker::GetLanguage() const { return language_; } +int Worker::Port() const { return port_; } + void Worker::AssignTaskId(const TaskID &task_id) { assigned_task_id_ = task_id; } const TaskID &Worker::GetAssignedTaskId() const { return assigned_task_id_; } diff --git a/src/ray/raylet/worker.h b/src/ray/raylet/worker.h index 7cd8d5e1d..6720c8ce3 100644 --- a/src/ray/raylet/worker.h +++ b/src/ray/raylet/worker.h @@ -17,7 +17,7 @@ namespace raylet { class Worker { public: /// A constructor that initializes a worker object. - Worker(pid_t pid, const Language &language, + Worker(pid_t pid, const Language &language, int port, std::shared_ptr connection); /// A destructor responsible for freeing all worker state. ~Worker() {} @@ -29,6 +29,7 @@ class Worker { /// Return the worker's PID. pid_t Pid() const; Language GetLanguage() const; + int Port() const; void AssignTaskId(const TaskID &task_id); const TaskID &GetAssignedTaskId() const; bool AddBlockedTaskId(const TaskID &task_id); @@ -56,6 +57,9 @@ class Worker { pid_t pid_; /// The language type of this worker. Language language_; + /// Port that this worker listens on. + /// If port <= 0, this indicates that the worker will not listen to a port. + int port_; /// Connection state of a worker. std::shared_ptr connection_; /// The worker's currently assigned task. diff --git a/src/ray/raylet/worker_pool.cc b/src/ray/raylet/worker_pool.cc index f15df88ae..3afe78b18 100644 --- a/src/ray/raylet/worker_pool.cc +++ b/src/ray/raylet/worker_pool.cc @@ -175,7 +175,8 @@ pid_t WorkerPool::StartProcess(const std::vector &worker_command_a void WorkerPool::RegisterWorker(const std::shared_ptr &worker) { const auto pid = worker->Pid(); - RAY_LOG(DEBUG) << "Registering worker with pid " << pid; + const auto port = worker->Port(); + RAY_LOG(DEBUG) << "Registering worker with pid " << pid << ", port: " << port; auto &state = GetStateForLanguage(worker->GetLanguage()); state.registered_workers.insert(std::move(worker)); diff --git a/src/ray/raylet/worker_pool_test.cc b/src/ray/raylet/worker_pool_test.cc index 698387a8f..715e89417 100644 --- a/src/ray/raylet/worker_pool_test.cc +++ b/src/ray/raylet/worker_pool_test.cc @@ -86,7 +86,7 @@ class WorkerPoolTest : public ::testing::Test { auto client = LocalClientConnection::Create(client_handler, message_handler, std::move(socket), "worker", {}, error_message_type_); - return std::shared_ptr(new Worker(pid, language, client)); + return std::shared_ptr(new Worker(pid, language, -1, client)); } void SetWorkerCommands( diff --git a/src/ray/rpc/client_call.h b/src/ray/rpc/client_call.h index a134c05a6..b132c66a4 100644 --- a/src/ray/rpc/client_call.h +++ b/src/ray/rpc/client_call.h @@ -30,6 +30,8 @@ class ClientCall { /// The callback to be called by `ClientCallManager` when the reply of this request is /// received. virtual void OnReplyReceived() = 0; + /// Return status. + virtual ray::Status GetStatus() = 0; virtual ~ClientCall() = default; }; @@ -49,6 +51,7 @@ using ClientCallback = std::function class ClientCallImpl : public ClientCall { public: + Status GetStatus() override { return GrpcStatusToRayStatus(status_); } void OnReplyReceived() override { if (callback_ != nullptr) { callback_(GrpcStatusToRayStatus(status_), reply_); diff --git a/src/ray/rpc/grpc_server.cc b/src/ray/rpc/grpc_server.cc index 08d064304..80ad81d98 100644 --- a/src/ray/rpc/grpc_server.cc +++ b/src/ray/rpc/grpc_server.cc @@ -13,7 +13,7 @@ void GrpcServer::Run() { builder.AddListeningPort(server_address, grpc::InsecureServerCredentials(), &port_); // Register all the services to this server. if (services_.size() == 0) { - RAY_LOG(WARNING) << "No service is found when start grpc server."; + RAY_LOG(WARNING) << "No service is found when start grpc server " << name_; } for (auto &entry : services_) { builder.RegisterService(&entry.get()); @@ -34,6 +34,8 @@ void GrpcServer::Run() { } // Start a thread that polls incoming requests. polling_thread_ = std::thread(&GrpcServer::PollEventsFromCompletionQueue, this); + // Set the server as running. + is_closed_ = false; } void GrpcServer::RegisterService(GrpcService &service) { diff --git a/src/ray/rpc/grpc_server.h b/src/ray/rpc/grpc_server.h index 01db49421..4af07ed27 100644 --- a/src/ray/rpc/grpc_server.h +++ b/src/ray/rpc/grpc_server.h @@ -33,7 +33,7 @@ class GrpcServer { /// \param[in] main_service The main event loop, to which service handler functions /// will be posted. GrpcServer(const std::string &name, const uint32_t port) - : name_(name), port_(port), is_closed_(false) {} + : name_(name), port_(port), is_closed_(true) {} /// Destruct this gRPC server. ~GrpcServer() { Shutdown(); } diff --git a/src/ray/rpc/worker/worker_client.h b/src/ray/rpc/worker/worker_client.h new file mode 100644 index 000000000..91ded57ac --- /dev/null +++ b/src/ray/rpc/worker/worker_client.h @@ -0,0 +1,58 @@ +#ifndef RAY_RPC_WORKER_CLIENT_H +#define RAY_RPC_WORKER_CLIENT_H + +#include + +#include + +#include "ray/common/status.h" +#include "ray/rpc/client_call.h" +#include "ray/util/logging.h" +#include "src/ray/protobuf/worker.grpc.pb.h" +#include "src/ray/protobuf/worker.pb.h" + +namespace ray { +namespace rpc { + +/// Client used for communicating with a remote worker server. +class WorkerTaskClient { + public: + /// Constructor. + /// + /// \param[in] address Address of the worker server. + /// \param[in] port Port of the worker server. + /// \param[in] client_call_manager The `ClientCallManager` used for managing requests. + WorkerTaskClient(const std::string &address, const int port, + ClientCallManager &client_call_manager) + : client_call_manager_(client_call_manager) { + std::shared_ptr channel = grpc::CreateChannel( + address + ":" + std::to_string(port), grpc::InsecureChannelCredentials()); + stub_ = WorkerTaskService::NewStub(channel); + }; + + /// Assign a task to the work. + /// + /// \param[in] request The request message. + /// \param[in] callback The callback function that handles reply. + /// \return if the rpc call succeeds + ray::Status AssignTask(const AssignTaskRequest &request, + const ClientCallback &callback) { + auto call = client_call_manager_ + .CreateCall( + *stub_, &WorkerTaskService::Stub::PrepareAsyncAssignTask, request, + callback); + return call->GetStatus(); + } + + private: + /// The gRPC-generated stub. + std::unique_ptr stub_; + + /// The `ClientCallManager` used for managing requests. + ClientCallManager &client_call_manager_; +}; + +} // namespace rpc +} // namespace ray + +#endif // RAY_RPC_WORKER_CLIENT_H diff --git a/src/ray/rpc/worker/worker_server.h b/src/ray/rpc/worker/worker_server.h new file mode 100644 index 000000000..adf9bb982 --- /dev/null +++ b/src/ray/rpc/worker/worker_server.h @@ -0,0 +1,68 @@ +#ifndef RAY_RPC_WORKER_SERVER_H +#define RAY_RPC_WORKER_SERVER_H + +#include "ray/rpc/grpc_server.h" +#include "ray/rpc/server_call.h" + +#include "src/ray/protobuf/worker.grpc.pb.h" +#include "src/ray/protobuf/worker.pb.h" + +namespace ray { +namespace rpc { + +/// Interface of the `WorkerService`, see `src/ray/protobuf/worker.proto`. +class WorkerTaskHandler { + public: + /// Handle a `AssignTask` request. + /// The implementation can handle this request asynchronously. When handling is done, + /// the `done_callback` should be called. + /// + /// \param[in] request The request message. + /// \param[out] reply The reply message. + /// \param[in] done_callback The callback to be called when the request is done. + virtual void HandleAssignTask(const AssignTaskRequest &request, AssignTaskReply *reply, + RequestDoneCallback done_callback) = 0; +}; + +/// The `GrpcServer` for `WorkerService`. +class WorkerTaskGrpcService : public GrpcService { + public: + /// Constructor. + /// + /// \param[in] main_service See super class. + /// \param[in] handler The service handler that actually handle the requests. + WorkerTaskGrpcService(boost::asio::io_service &main_service, + WorkerTaskHandler &service_handler) + : GrpcService(main_service), service_handler_(service_handler){}; + + protected: + grpc::Service &GetGrpcService() override { return service_; } + + void InitServerCallFactories( + const std::unique_ptr &cq, + std::vector, int>> + *server_call_factories_and_concurrencies) override { + // Initialize the Factory for `AssignTask` requests. + std::unique_ptr push_task_call_Factory( + new ServerCallFactoryImpl( + service_, &WorkerTaskService::AsyncService::RequestAssignTask, + service_handler_, &WorkerTaskHandler::HandleAssignTask, cq, main_service_)); + + // Set `AssignTask`'s accept concurrency to 5. + server_call_factories_and_concurrencies->emplace_back( + std::move(push_task_call_Factory), 5); + } + + private: + /// The grpc async service object. + WorkerTaskService::AsyncService service_; + + /// The service handler that actually handle the requests. + WorkerTaskHandler &service_handler_; +}; + +} // namespace rpc +} // namespace ray + +#endif