From 1000e3322db3dbeffa4c10572950ff390c7a2574 Mon Sep 17 00:00:00 2001 From: fangfengbin <869218239a@zju.edu.cn> Date: Mon, 6 Jan 2020 11:09:32 +0800 Subject: [PATCH] Add gcs server task info handler (#6695) --- src/ray/gcs/gcs_server/gcs_server.cc | 11 +++ src/ray/gcs/gcs_server/gcs_server.h | 8 +- .../gcs/gcs_server/task_info_handler_impl.cc | 73 +++++++++++++++++++ .../gcs/gcs_server/task_info_handler_impl.h | 32 ++++++++ .../gcs_server/test/gcs_server_rpc_test.cc | 72 +++++++++++++++++- src/ray/gcs/redis_accessor.cc | 3 + src/ray/protobuf/gcs_service.proto | 32 ++++++++ src/ray/rpc/gcs_server/gcs_rpc_client.h | 26 ++++--- src/ray/rpc/gcs_server/gcs_rpc_server.h | 53 ++++++++++++++ 9 files changed, 296 insertions(+), 14 deletions(-) create mode 100644 src/ray/gcs/gcs_server/task_info_handler_impl.cc create mode 100644 src/ray/gcs/gcs_server/task_info_handler_impl.h diff --git a/src/ray/gcs/gcs_server/gcs_server.cc b/src/ray/gcs/gcs_server/gcs_server.cc index 2c8d652fc..52aff19b8 100644 --- a/src/ray/gcs/gcs_server/gcs_server.cc +++ b/src/ray/gcs/gcs_server/gcs_server.cc @@ -3,6 +3,7 @@ #include "job_info_handler_impl.h" #include "node_info_handler_impl.h" #include "object_info_handler_impl.h" +#include "task_info_handler_impl.h" namespace ray { namespace gcs { @@ -38,6 +39,11 @@ void GcsServer::Start() { new rpc::ObjectInfoGrpcService(main_service_, *object_info_handler_)); rpc_server_.RegisterService(*object_info_service_); + task_info_handler_ = InitTaskInfoHandler(); + task_info_service_.reset( + new rpc::TaskInfoGrpcService(main_service_, *task_info_handler_)); + rpc_server_.RegisterService(*task_info_service_); + // Run rpc server. rpc_server_.Run(); @@ -84,5 +90,10 @@ std::unique_ptr GcsServer::InitObjectInfoHandler() { new rpc::DefaultObjectInfoHandler(*redis_gcs_client_)); } +std::unique_ptr GcsServer::InitTaskInfoHandler() { + return std::unique_ptr( + new rpc::DefaultTaskInfoHandler(*redis_gcs_client_)); +} + } // namespace gcs } // namespace ray diff --git a/src/ray/gcs/gcs_server/gcs_server.h b/src/ray/gcs/gcs_server/gcs_server.h index 48103d464..4b33feea0 100644 --- a/src/ray/gcs/gcs_server/gcs_server.h +++ b/src/ray/gcs/gcs_server/gcs_server.h @@ -3,8 +3,6 @@ #include #include -#include -#include namespace ray { namespace gcs { @@ -58,6 +56,9 @@ class GcsServer { /// The object info handler virtual std::unique_ptr InitObjectInfoHandler(); + /// The task info handler + virtual std::unique_ptr InitTaskInfoHandler(); + private: /// Gcs server configuration GcsServerConfig config_; @@ -77,6 +78,9 @@ class GcsServer { /// Object info handler and service std::unique_ptr object_info_handler_; std::unique_ptr object_info_service_; + /// Task info handler and service + std::unique_ptr task_info_handler_; + std::unique_ptr task_info_service_; /// Backend client std::shared_ptr redis_gcs_client_; }; diff --git a/src/ray/gcs/gcs_server/task_info_handler_impl.cc b/src/ray/gcs/gcs_server/task_info_handler_impl.cc new file mode 100644 index 000000000..2c62466c3 --- /dev/null +++ b/src/ray/gcs/gcs_server/task_info_handler_impl.cc @@ -0,0 +1,73 @@ +#include "task_info_handler_impl.h" + +namespace ray { +namespace rpc { + +void DefaultTaskInfoHandler::HandleAddTask(const AddTaskRequest &request, + AddTaskReply *reply, + SendReplyCallback send_reply_callback) { + JobID job_id = JobID::FromBinary(request.task_data().task().task_spec().job_id()); + TaskID task_id = TaskID::FromBinary(request.task_data().task().task_spec().task_id()); + RAY_LOG(DEBUG) << "Adding task, task id = " << task_id << ", job id = " << job_id; + auto task_table_data = std::make_shared(); + task_table_data->CopyFrom(request.task_data()); + auto on_done = [job_id, task_id, request, send_reply_callback](Status status) { + if (!status.ok()) { + RAY_LOG(ERROR) << "Failed to add task, task id = " << task_id + << ", job id = " << job_id; + } + send_reply_callback(status, nullptr, nullptr); + }; + + Status status = gcs_client_.Tasks().AsyncAdd(task_table_data, on_done); + if (!status.ok()) { + on_done(status); + } + RAY_LOG(DEBUG) << "Finished adding task, task id = " << task_id + << ", job id = " << job_id; +} + +void DefaultTaskInfoHandler::HandleGetTask(const GetTaskRequest &request, + GetTaskReply *reply, + SendReplyCallback send_reply_callback) { + TaskID task_id = TaskID::FromBinary(request.task_id()); + RAY_LOG(DEBUG) << "Getting task, task id = " << task_id; + auto on_done = [task_id, request, reply, send_reply_callback]( + Status status, const boost::optional &result) { + if (status.ok()) { + RAY_DCHECK(result); + reply->mutable_task_data()->CopyFrom(*result); + } else { + RAY_LOG(ERROR) << "Failed to get task, task id = " << task_id; + } + send_reply_callback(status, nullptr, nullptr); + }; + + Status status = gcs_client_.Tasks().AsyncGet(task_id, on_done); + if (!status.ok()) { + on_done(status, boost::none); + } + RAY_LOG(DEBUG) << "Finished getting task, task id = " << task_id; +} + +void DefaultTaskInfoHandler::HandleDeleteTasks(const DeleteTasksRequest &request, + DeleteTasksReply *reply, + SendReplyCallback send_reply_callback) { + std::vector task_ids = IdVectorFromProtobuf(request.task_id_list()); + RAY_LOG(DEBUG) << "Deleting tasks, task id list size = " << task_ids.size(); + auto on_done = [task_ids, request, send_reply_callback](Status status) { + if (!status.ok()) { + RAY_LOG(ERROR) << "Failed to delete tasks, task id list size = " << task_ids.size(); + } + send_reply_callback(status, nullptr, nullptr); + }; + + Status status = gcs_client_.Tasks().AsyncDelete(task_ids, on_done); + if (!status.ok()) { + on_done(status); + } + RAY_LOG(DEBUG) << "Finished deleting tasks, task id list size = " << task_ids.size(); +} + +} // namespace rpc +} // namespace ray diff --git a/src/ray/gcs/gcs_server/task_info_handler_impl.h b/src/ray/gcs/gcs_server/task_info_handler_impl.h new file mode 100644 index 000000000..0d210ac2f --- /dev/null +++ b/src/ray/gcs/gcs_server/task_info_handler_impl.h @@ -0,0 +1,32 @@ +#ifndef RAY_GCS_TASK_INFO_HANDLER_IMPL_H +#define RAY_GCS_TASK_INFO_HANDLER_IMPL_H + +#include "ray/gcs/redis_gcs_client.h" +#include "ray/rpc/gcs_server/gcs_rpc_server.h" + +namespace ray { +namespace rpc { + +/// This implementation class of `TaskInfoHandler`. +class DefaultTaskInfoHandler : public rpc::TaskInfoHandler { + public: + explicit DefaultTaskInfoHandler(gcs::RedisGcsClient &gcs_client) + : gcs_client_(gcs_client) {} + + void HandleAddTask(const AddTaskRequest &request, AddTaskReply *reply, + SendReplyCallback send_reply_callback) override; + + void HandleGetTask(const GetTaskRequest &request, GetTaskReply *reply, + SendReplyCallback send_reply_callback) override; + + void HandleDeleteTasks(const DeleteTasksRequest &request, DeleteTasksReply *reply, + SendReplyCallback send_reply_callback) override; + + private: + gcs::RedisGcsClient &gcs_client_; +}; + +} // namespace rpc +} // namespace ray + +#endif // RAY_GCS_TASK_INFO_HANDLER_IMPL_H diff --git a/src/ray/gcs/gcs_server/test/gcs_server_rpc_test.cc b/src/ray/gcs/gcs_server/test/gcs_server_rpc_test.cc index 3bea2a317..7b3c5848b 100644 --- a/src/ray/gcs/gcs_server/test/gcs_server_rpc_test.cc +++ b/src/ray/gcs/gcs_server/test/gcs_server_rpc_test.cc @@ -1,7 +1,5 @@ #include "gtest/gtest.h" -#include "ray/gcs/gcs_server/actor_info_handler_impl.h" #include "ray/gcs/gcs_server/gcs_server.h" -#include "ray/gcs/gcs_server/job_info_handler_impl.h" #include "ray/rpc/gcs_server/gcs_rpc_client.h" #include "ray/util/test_util.h" @@ -289,6 +287,43 @@ class GcsServerTest : public RedisServiceManagerForTest { return object_locations; } + bool AddTask(const rpc::AddTaskRequest &request) { + std::promise promise; + client_->AddTask(request, + [&promise](const Status &status, const rpc::AddTaskReply &reply) { + RAY_CHECK_OK(status); + promise.set_value(true); + }); + return WaitReady(promise.get_future(), timeout_ms_); + } + + rpc::TaskTableData GetTask(const std::string &task_id) { + rpc::TaskTableData task_data; + rpc::GetTaskRequest request; + request.set_task_id(task_id); + std::promise promise; + client_->GetTask(request, [&task_data, &promise](const Status &status, + const rpc::GetTaskReply &reply) { + if (status.ok()) { + task_data.CopyFrom(reply.task_data()); + } + promise.set_value(true); + }); + + EXPECT_TRUE(WaitReady(promise.get_future(), timeout_ms_)); + return task_data; + } + + bool DeleteTasks(const rpc::DeleteTasksRequest &request) { + std::promise promise; + client_->DeleteTasks( + request, [&promise](const Status &status, const rpc::DeleteTasksReply &reply) { + RAY_CHECK_OK(status); + promise.set_value(true); + }); + return WaitReady(promise.get_future(), timeout_ms_); + } + bool WaitReady(const std::future &future, uint64_t timeout_ms) { auto status = future.wait_for(std::chrono::milliseconds(timeout_ms)); return status == std::future_status::ready; @@ -323,6 +358,18 @@ class GcsServerTest : public RedisServiceManagerForTest { return gcs_node_info; } + rpc::TaskTableData GenTaskTableData(const std::string &job_id, + const std::string &task_id) { + rpc::TaskTableData task_table_data; + rpc::Task task; + rpc::TaskSpec task_spec; + task_spec.set_job_id(job_id); + task_spec.set_task_id(task_id); + task.mutable_task_spec()->CopyFrom(task_spec); + task_table_data.mutable_task()->CopyFrom(task); + return task_table_data; + } + protected: // Gcs server std::unique_ptr gcs_server_; @@ -480,6 +527,27 @@ TEST_F(GcsServerTest, TestObjectInfo) { ASSERT_TRUE(object_locations[0].manager() == node2_id.Binary()); } +TEST_F(GcsServerTest, TestTaskInfo) { + // Create task_table_data + JobID job_id = JobID::FromInt(1); + TaskID task_id = TaskID::ForDriverTask(job_id); + rpc::TaskTableData job_table_data = GenTaskTableData(job_id.Binary(), task_id.Binary()); + + // Add task + rpc::AddTaskRequest add_task_request; + add_task_request.mutable_task_data()->CopyFrom(job_table_data); + ASSERT_TRUE(AddTask(add_task_request)); + rpc::TaskTableData result = GetTask(task_id.Binary()); + ASSERT_TRUE(result.task().task_spec().job_id() == job_id.Binary()); + + // Delete task + rpc::DeleteTasksRequest delete_tasks_request; + delete_tasks_request.add_task_id_list(task_id.Binary()); + ASSERT_TRUE(DeleteTasks(delete_tasks_request)); + result = GetTask(task_id.Binary()); + ASSERT_TRUE(!result.has_task()); +} + } // namespace ray int main(int argc, char **argv) { diff --git a/src/ray/gcs/redis_accessor.cc b/src/ray/gcs/redis_accessor.cc index a1fb72aab..bffd71392 100644 --- a/src/ray/gcs/redis_accessor.cc +++ b/src/ray/gcs/redis_accessor.cc @@ -286,6 +286,9 @@ Status RedisTaskInfoAccessor::AsyncDelete(const std::vector &task_ids, const StatusCallback &callback) { raylet::TaskTable &task_table = client_impl_->raylet_task_table(); task_table.Delete(JobID::Nil(), task_ids); + if (callback) { + callback(Status::OK()); + } // TODO(micafan) Always return OK here. // Confirm if we need to handle the deletion failure and how to handle it. return Status::OK(); diff --git a/src/ray/protobuf/gcs_service.proto b/src/ray/protobuf/gcs_service.proto index 3e373b2d1..d98dcef0a 100644 --- a/src/ray/protobuf/gcs_service.proto +++ b/src/ray/protobuf/gcs_service.proto @@ -218,3 +218,35 @@ service ObjectInfoGcsService { rpc RemoveObjectLocation(RemoveObjectLocationRequest) returns (RemoveObjectLocationReply); } + +message AddTaskRequest { + TaskTableData task_data = 1; +} + +message AddTaskReply { +} + +message GetTaskRequest { + bytes task_id = 1; +} + +message GetTaskReply { + TaskTableData task_data = 1; +} + +message DeleteTasksRequest { + repeated bytes task_id_list = 1; +} + +message DeleteTasksReply { +} + +// Service for task info access. +service TaskInfoGcsService { + // Add a task to GCS Service. + rpc AddTask(AddTaskRequest) returns (AddTaskReply); + // Get task information from GCS Service. + rpc GetTask(GetTaskRequest) returns (GetTaskReply); + // Delete tasks from GCS Service. + rpc DeleteTasks(DeleteTasksRequest) returns (DeleteTasksReply); +} diff --git a/src/ray/rpc/gcs_server/gcs_rpc_client.h b/src/ray/rpc/gcs_server/gcs_rpc_client.h index 531b693fc..3a2a7af88 100644 --- a/src/ray/rpc/gcs_server/gcs_rpc_client.h +++ b/src/ray/rpc/gcs_server/gcs_rpc_client.h @@ -1,11 +1,6 @@ #ifndef RAY_RPC_GCS_RPC_CLIENT_H #define RAY_RPC_GCS_RPC_CLIENT_H -#include - -#include - -#include "src/ray/protobuf/gcs_service.pb.h" #include "src/ray/rpc/grpc_client.h" namespace ray { @@ -20,8 +15,7 @@ class GcsRpcClient { /// \param[in] port Port of the gcs server. /// \param[in] client_call_manager The `ClientCallManager` used for managing requests. GcsRpcClient(const std::string &address, const int port, - ClientCallManager &client_call_manager) - : client_call_manager_(client_call_manager) { + ClientCallManager &client_call_manager) { job_info_grpc_client_ = std::unique_ptr>( new GrpcClient(address, port, client_call_manager)); actor_info_grpc_client_ = std::unique_ptr>( @@ -30,6 +24,8 @@ class GcsRpcClient { new GrpcClient(address, port, client_call_manager)); object_info_grpc_client_ = std::unique_ptr>( new GrpcClient(address, port, client_call_manager)); + task_info_grpc_client_ = std::unique_ptr>( + new GrpcClient(address, port, client_call_manager)); }; /// Add job info to gcs server. @@ -108,15 +104,25 @@ class GcsRpcClient { VOID_RPC_CLIENT_METHOD(ObjectInfoGcsService, RemoveObjectLocation, request, callback, object_info_grpc_client_, ) + /// Add a task to GCS Service. + VOID_RPC_CLIENT_METHOD(TaskInfoGcsService, AddTask, request, callback, + task_info_grpc_client_) + + /// Get task information from GCS Service. + VOID_RPC_CLIENT_METHOD(TaskInfoGcsService, GetTask, request, callback, + task_info_grpc_client_) + + /// Delete tasks from GCS Service. + VOID_RPC_CLIENT_METHOD(TaskInfoGcsService, DeleteTasks, request, callback, + task_info_grpc_client_) + private: /// The gRPC-generated stub. std::unique_ptr> job_info_grpc_client_; std::unique_ptr> actor_info_grpc_client_; std::unique_ptr> node_info_grpc_client_; std::unique_ptr> object_info_grpc_client_; - - /// The `ClientCallManager` used for managing requests. - ClientCallManager &client_call_manager_; + std::unique_ptr> task_info_grpc_client_; }; } // namespace rpc diff --git a/src/ray/rpc/gcs_server/gcs_rpc_server.h b/src/ray/rpc/gcs_server/gcs_rpc_server.h index 7998a7630..a3f28055d 100644 --- a/src/ray/rpc/gcs_server/gcs_rpc_server.h +++ b/src/ray/rpc/gcs_server/gcs_rpc_server.h @@ -45,6 +45,15 @@ namespace rpc { server_call_factories_and_concurrencies->emplace_back( \ std::move(HANDLER##_call_factory), CONCURRENCY); +#define TASK_INFO_SERVICE_RPC_HANDLER(HANDLER, CONCURRENCY) \ + std::unique_ptr HANDLER##_call_factory( \ + new ServerCallFactoryImpl( \ + service_, &TaskInfoGcsService::AsyncService::Request##HANDLER, \ + service_handler_, &TaskInfoHandler::Handle##HANDLER, cq, main_service_)); \ + server_call_factories_and_concurrencies->emplace_back( \ + std::move(HANDLER##_call_factory), CONCURRENCY); + class JobInfoHandler { public: virtual ~JobInfoHandler() = default; @@ -263,6 +272,50 @@ class ObjectInfoGrpcService : public GrpcService { ObjectInfoHandler &service_handler_; }; +class TaskInfoHandler { + public: + virtual ~TaskInfoHandler() = default; + + virtual void HandleAddTask(const AddTaskRequest &request, AddTaskReply *reply, + SendReplyCallback send_reply_callback) = 0; + + virtual void HandleGetTask(const GetTaskRequest &request, GetTaskReply *reply, + SendReplyCallback send_reply_callback) = 0; + + virtual void HandleDeleteTasks(const DeleteTasksRequest &request, + DeleteTasksReply *reply, + SendReplyCallback send_reply_callback) = 0; +}; + +/// The `GrpcService` for `TaskInfoGcsService`. +class TaskInfoGrpcService : public GrpcService { + public: + /// Constructor. + /// + /// \param[in] handler The service handler that actually handle the requests. + explicit TaskInfoGrpcService(boost::asio::io_service &io_service, + TaskInfoHandler &handler) + : GrpcService(io_service), service_handler_(handler){}; + + protected: + grpc::Service &GetGrpcService() override { return service_; } + + void InitServerCallFactories( + const std::unique_ptr &cq, + std::vector, int>> + *server_call_factories_and_concurrencies) override { + TASK_INFO_SERVICE_RPC_HANDLER(AddTask, 1); + TASK_INFO_SERVICE_RPC_HANDLER(GetTask, 1); + TASK_INFO_SERVICE_RPC_HANDLER(DeleteTasks, 1); + } + + private: + /// The grpc async service object. + TaskInfoGcsService::AsyncService service_; + /// The service handler that actually handle the requests. + TaskInfoHandler &service_handler_; +}; + } // namespace rpc } // namespace ray