diff --git a/BUILD.bazel b/BUILD.bazel index d25da120e..cfe4f7cfc 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -678,6 +678,16 @@ cc_test( ], ) +cc_test( + name = "sequencer_test", + srcs = ["src/ray/util/sequencer_test.cc"], + copts = COPTS, + deps = [ + ":ray_util", + "@com_google_googletest//:gtest_main", + ], +) + cc_test( name = "stats_test", srcs = ["src/ray/stats/stats_test.cc"], @@ -830,6 +840,7 @@ cc_library( ":sha256", "@boost//:asio", "@com_github_google_glog//:glog", + "@com_google_absl//absl/synchronization", "@com_google_absl//absl/time", "@plasma//:plasma_client", ], diff --git a/src/ray/gcs/gcs_client/service_based_accessor.cc b/src/ray/gcs/gcs_client/service_based_accessor.cc index 9c3627267..f580a3044 100644 --- a/src/ray/gcs/gcs_client/service_based_accessor.cc +++ b/src/ray/gcs/gcs_client/service_based_accessor.cc @@ -94,15 +94,22 @@ Status ServiceBasedActorInfoAccessor::AsyncRegister( RAY_LOG(DEBUG) << "Registering actor info, actor id = " << actor_id; rpc::RegisterActorInfoRequest request; request.mutable_actor_table_data()->CopyFrom(*data_ptr); - client_impl_->GetGcsRpcClient().RegisterActorInfo( - request, [actor_id, callback](const Status &status, - const rpc::RegisterActorInfoReply &reply) { - if (callback) { - callback(status); - } - RAY_LOG(DEBUG) << "Finished registering actor info, status = " << status - << ", actor id = " << actor_id; - }); + + auto operation = [this, request, actor_id, + callback](SequencerDoneCallback done_callback) { + client_impl_->GetGcsRpcClient().RegisterActorInfo( + request, [actor_id, callback, done_callback]( + const Status &status, const rpc::RegisterActorInfoReply &reply) { + if (callback) { + callback(status); + } + RAY_LOG(DEBUG) << "Finished registering actor info, status = " << status + << ", actor id = " << actor_id; + done_callback(); + }); + }; + + sequencer_.Post(actor_id, operation); return Status::OK(); } @@ -113,15 +120,22 @@ Status ServiceBasedActorInfoAccessor::AsyncUpdate( rpc::UpdateActorInfoRequest request; request.set_actor_id(actor_id.Binary()); request.mutable_actor_table_data()->CopyFrom(*data_ptr); - client_impl_->GetGcsRpcClient().UpdateActorInfo( - request, - [actor_id, callback](const Status &status, const rpc::UpdateActorInfoReply &reply) { - if (callback) { - callback(status); - } - RAY_LOG(DEBUG) << "Finished updating actor info, status = " << status - << ", actor id = " << actor_id; - }); + + auto operation = [this, request, actor_id, + callback](SequencerDoneCallback done_callback) { + client_impl_->GetGcsRpcClient().UpdateActorInfo( + request, [actor_id, callback, done_callback]( + const Status &status, const rpc::UpdateActorInfoReply &reply) { + if (callback) { + callback(status); + } + RAY_LOG(DEBUG) << "Finished updating actor info, status = " << status + << ", actor id = " << actor_id; + done_callback(); + }); + }; + + sequencer_.Post(actor_id, operation); return Status::OK(); } @@ -167,16 +181,23 @@ Status ServiceBasedActorInfoAccessor::AsyncAddCheckpoint( << ", checkpoint id = " << checkpoint_id; rpc::AddActorCheckpointRequest request; request.mutable_checkpoint_data()->CopyFrom(*data_ptr); - client_impl_->GetGcsRpcClient().AddActorCheckpoint( - request, [actor_id, checkpoint_id, callback]( - const Status &status, const rpc::AddActorCheckpointReply &reply) { - if (callback) { - callback(status); - } - RAY_LOG(DEBUG) << "Finished adding actor checkpoint, status = " << status - << ", actor id = " << actor_id - << ", checkpoint id = " << checkpoint_id; - }); + + auto operation = [this, request, actor_id, checkpoint_id, + callback](SequencerDoneCallback done_callback) { + client_impl_->GetGcsRpcClient().AddActorCheckpoint( + request, [actor_id, checkpoint_id, callback, done_callback]( + const Status &status, const rpc::AddActorCheckpointReply &reply) { + if (callback) { + callback(status); + } + RAY_LOG(DEBUG) << "Finished adding actor checkpoint, status = " << status + << ", actor id = " << actor_id + << ", checkpoint id = " << checkpoint_id; + done_callback(); + }); + }; + + sequencer_.Post(actor_id, operation); return Status::OK(); } @@ -392,15 +413,22 @@ Status ServiceBasedNodeInfoAccessor::AsyncUpdateResources( for (auto &resource : resources) { (*request.mutable_resources())[resource.first] = *resource.second; } - client_impl_->GetGcsRpcClient().UpdateResources( - request, - [node_id, callback](const Status &status, const rpc::UpdateResourcesReply &reply) { - if (callback) { - callback(status); - } - RAY_LOG(DEBUG) << "Finished updating node resources, status = " << status - << ", node id = " << node_id; - }); + + auto operation = [this, request, node_id, + callback](SequencerDoneCallback done_callback) { + client_impl_->GetGcsRpcClient().UpdateResources( + request, [node_id, callback, done_callback]( + const Status &status, const rpc::UpdateResourcesReply &reply) { + if (callback) { + callback(status); + } + RAY_LOG(DEBUG) << "Finished updating node resources, status = " << status + << ", node id = " << node_id; + done_callback(); + }); + }; + + sequencer_.Post(node_id, operation); return Status::OK(); } @@ -413,15 +441,22 @@ Status ServiceBasedNodeInfoAccessor::AsyncDeleteResources( for (auto &resource_name : resource_names) { request.add_resource_name_list(resource_name); } - client_impl_->GetGcsRpcClient().DeleteResources( - request, - [node_id, callback](const Status &status, const rpc::DeleteResourcesReply &reply) { - if (callback) { - callback(status); - } - RAY_LOG(DEBUG) << "Finished deleting node resources, status = " << status - << ", node id = " << node_id; - }); + + auto operation = [this, request, node_id, + callback](SequencerDoneCallback done_callback) { + client_impl_->GetGcsRpcClient().DeleteResources( + request, [node_id, callback, done_callback]( + const Status &status, const rpc::DeleteResourcesReply &reply) { + if (callback) { + callback(status); + } + RAY_LOG(DEBUG) << "Finished deleting node resources, status = " << status + << ", node id = " << node_id; + done_callback(); + }); + }; + + sequencer_.Post(node_id, operation); return Status::OK(); } @@ -679,15 +714,23 @@ Status ServiceBasedObjectInfoAccessor::AsyncAddLocation(const ObjectID &object_i rpc::AddObjectLocationRequest request; request.set_object_id(object_id.Binary()); request.set_node_id(node_id.Binary()); - client_impl_->GetGcsRpcClient().AddObjectLocation( - request, [object_id, node_id, callback](const Status &status, - const rpc::AddObjectLocationReply &reply) { - if (callback) { - callback(status); - } - RAY_LOG(DEBUG) << "Finished adding object location, status = " << status - << ", object id = " << object_id << ", node id = " << node_id; - }); + + auto operation = [this, request, object_id, node_id, + callback](SequencerDoneCallback done_callback) { + client_impl_->GetGcsRpcClient().AddObjectLocation( + request, [object_id, node_id, callback, done_callback]( + const Status &status, const rpc::AddObjectLocationReply &reply) { + if (callback) { + callback(status); + } + + RAY_LOG(DEBUG) << "Finished adding object location, status = " << status + << ", object id = " << object_id << ", node id = " << node_id; + done_callback(); + }); + }; + + sequencer_.Post(object_id, operation); return Status::OK(); } @@ -698,15 +741,22 @@ Status ServiceBasedObjectInfoAccessor::AsyncRemoveLocation( rpc::RemoveObjectLocationRequest request; request.set_object_id(object_id.Binary()); request.set_node_id(node_id.Binary()); - client_impl_->GetGcsRpcClient().RemoveObjectLocation( - request, [object_id, node_id, callback]( - const Status &status, const rpc::RemoveObjectLocationReply &reply) { - if (callback) { - callback(status); - } - RAY_LOG(DEBUG) << "Finished removing object location, status = " << status - << ", object id = " << object_id << ", node id = " << node_id; - }); + + auto operation = [this, request, object_id, node_id, + callback](SequencerDoneCallback done_callback) { + client_impl_->GetGcsRpcClient().RemoveObjectLocation( + request, [object_id, node_id, callback, done_callback]( + const Status &status, const rpc::RemoveObjectLocationReply &reply) { + if (callback) { + callback(status); + } + RAY_LOG(DEBUG) << "Finished removing object location, status = " << status + << ", object id = " << object_id << ", node id = " << node_id; + done_callback(); + }); + }; + + sequencer_.Post(object_id, operation); return Status::OK(); } diff --git a/src/ray/gcs/gcs_client/service_based_accessor.h b/src/ray/gcs/gcs_client/service_based_accessor.h index 2f781c6f9..9f1444d49 100644 --- a/src/ray/gcs/gcs_client/service_based_accessor.h +++ b/src/ray/gcs/gcs_client/service_based_accessor.h @@ -1,8 +1,9 @@ #ifndef RAY_GCS_SERVICE_BASED_ACCESSOR_H #define RAY_GCS_SERVICE_BASED_ACCESSOR_H -#include "src/ray/gcs/accessor.h" -#include "src/ray/gcs/subscription_executor.h" +#include "ray/gcs/accessor.h" +#include "ray/gcs/subscription_executor.h" +#include "ray/util/sequencer.h" namespace ray { namespace gcs { @@ -82,6 +83,8 @@ class ServiceBasedActorInfoAccessor : public ActorInfoAccessor { typedef SubscriptionExecutor ActorSubscriptionExecutor; ActorSubscriptionExecutor actor_sub_executor_; + + Sequencer sequencer_; }; /// \class ServiceBasedNodeInfoAccessor @@ -165,6 +168,8 @@ class ServiceBasedNodeInfoAccessor : public NodeInfoAccessor { GcsNodeInfo local_node_info_; ClientID local_node_id_; + + Sequencer sequencer_; }; /// \class ServiceBasedTaskInfoAccessor @@ -255,6 +260,8 @@ class ServiceBasedObjectInfoAccessor : public ObjectInfoAccessor { typedef SubscriptionExecutor ObjectSubscriptionExecutor; ObjectSubscriptionExecutor object_sub_executor_; + + Sequencer sequencer_; }; /// \class ServiceBasedStatsInfoAccessor diff --git a/src/ray/gcs/gcs_client/test/service_based_gcs_client_test.cc b/src/ray/gcs/gcs_client/test/service_based_gcs_client_test.cc index 2b323cc28..168257610 100644 --- a/src/ray/gcs/gcs_client/test/service_based_gcs_client_test.cc +++ b/src/ray/gcs/gcs_client/test/service_based_gcs_client_test.cc @@ -32,7 +32,7 @@ class ServiceBasedGcsGcsClientTest : public RedisServiceManagerForTest { thread_gcs_server_.reset(new std::thread([this] { gcs_server_->Start(); })); // Wait until server starts listening. - while (gcs_server_->GetPort() == 0) { + while (!gcs_server_->IsStarted()) { std::this_thread::sleep_for(std::chrono::milliseconds(10)); } diff --git a/src/ray/gcs/gcs_server/gcs_server.cc b/src/ray/gcs/gcs_server/gcs_server.cc index 29247cc7f..8b849df9a 100644 --- a/src/ray/gcs/gcs_server/gcs_server.cc +++ b/src/ray/gcs/gcs_server/gcs_server.cc @@ -66,6 +66,7 @@ void GcsServer::Start() { // Store gcs rpc server address in redis StoreGcsServerAddressInRedis(); + is_started_ = true; // Run the event loop. // Using boost::asio::io_context::work to avoid ending the event loop when diff --git a/src/ray/gcs/gcs_server/gcs_server.h b/src/ray/gcs/gcs_server/gcs_server.h index 9321440c0..4d3e9e9f3 100644 --- a/src/ray/gcs/gcs_server/gcs_server.h +++ b/src/ray/gcs/gcs_server/gcs_server.h @@ -38,6 +38,9 @@ class GcsServer { /// Get the port of this gcs server. int GetPort() const { return rpc_server_.GetPort(); } + /// Check if gcs server is started + bool IsStarted() const { return is_started_; } + protected: /// Initialize the backend storage client /// The gcs server is just the proxy between the gcs client and reliable storage @@ -108,6 +111,8 @@ class GcsServer { std::unique_ptr worker_info_service_; /// Backend client std::shared_ptr redis_gcs_client_; + /// Gcs service init flag + bool is_started_ = false; }; } // namespace gcs diff --git a/src/ray/gcs/gcs_server/object_info_handler_impl.cc b/src/ray/gcs/gcs_server/object_info_handler_impl.cc index 69c4f7d23..29175a84b 100644 --- a/src/ray/gcs/gcs_server/object_info_handler_impl.cc +++ b/src/ray/gcs/gcs_server/object_info_handler_impl.cc @@ -66,7 +66,7 @@ void DefaultObjectInfoHandler::HandleRemoveObjectLocation( auto on_done = [object_id, node_id, send_reply_callback](Status status) { if (!status.ok()) { - RAY_LOG(ERROR) << "Failed to add object location: " << status.ToString() + RAY_LOG(ERROR) << "Failed to remove object location: " << status.ToString() << ", object id = " << object_id << ", node id = " << node_id; } send_reply_callback(status, nullptr, nullptr); diff --git a/src/ray/rpc/gcs_server/gcs_rpc_server.h b/src/ray/rpc/gcs_server/gcs_rpc_server.h index c92260c0d..f2079ef56 100644 --- a/src/ray/rpc/gcs_server/gcs_rpc_server.h +++ b/src/ray/rpc/gcs_server/gcs_rpc_server.h @@ -33,6 +33,8 @@ namespace rpc { #define WORKER_INFO_SERVICE_RPC_HANDLER(HANDLER, CONCURRENCY) \ RPC_SERVICE_HANDLER(WorkerInfoGcsService, HANDLER, CONCURRENCY) +#define SERVER_CALL_CONCURRENCY 9999 + class JobInfoGcsServiceHandler { public: virtual ~JobInfoGcsServiceHandler() = default; @@ -62,8 +64,8 @@ class JobInfoGrpcService : public GrpcService { const std::unique_ptr &cq, std::vector, int>> *server_call_factories_and_concurrencies) override { - JOB_INFO_SERVICE_RPC_HANDLER(AddJob, 1); - JOB_INFO_SERVICE_RPC_HANDLER(MarkJobFinished, 1); + JOB_INFO_SERVICE_RPC_HANDLER(AddJob, SERVER_CALL_CONCURRENCY); + JOB_INFO_SERVICE_RPC_HANDLER(MarkJobFinished, SERVER_CALL_CONCURRENCY); } private: @@ -119,12 +121,12 @@ class ActorInfoGrpcService : public GrpcService { const std::unique_ptr &cq, std::vector, int>> *server_call_factories_and_concurrencies) override { - ACTOR_INFO_SERVICE_RPC_HANDLER(GetActorInfo, 1); - ACTOR_INFO_SERVICE_RPC_HANDLER(RegisterActorInfo, 1); - ACTOR_INFO_SERVICE_RPC_HANDLER(UpdateActorInfo, 1); - ACTOR_INFO_SERVICE_RPC_HANDLER(AddActorCheckpoint, 1); - ACTOR_INFO_SERVICE_RPC_HANDLER(GetActorCheckpoint, 1); - ACTOR_INFO_SERVICE_RPC_HANDLER(GetActorCheckpointID, 1); + ACTOR_INFO_SERVICE_RPC_HANDLER(GetActorInfo, SERVER_CALL_CONCURRENCY); + ACTOR_INFO_SERVICE_RPC_HANDLER(RegisterActorInfo, SERVER_CALL_CONCURRENCY); + ACTOR_INFO_SERVICE_RPC_HANDLER(UpdateActorInfo, SERVER_CALL_CONCURRENCY); + ACTOR_INFO_SERVICE_RPC_HANDLER(AddActorCheckpoint, SERVER_CALL_CONCURRENCY); + ACTOR_INFO_SERVICE_RPC_HANDLER(GetActorCheckpoint, SERVER_CALL_CONCURRENCY); + ACTOR_INFO_SERVICE_RPC_HANDLER(GetActorCheckpointID, SERVER_CALL_CONCURRENCY); } private: @@ -188,14 +190,14 @@ class NodeInfoGrpcService : public GrpcService { const std::unique_ptr &cq, std::vector, int>> *server_call_factories_and_concurrencies) override { - NODE_INFO_SERVICE_RPC_HANDLER(RegisterNode, 1); - NODE_INFO_SERVICE_RPC_HANDLER(UnregisterNode, 1); - NODE_INFO_SERVICE_RPC_HANDLER(GetAllNodeInfo, 1); - NODE_INFO_SERVICE_RPC_HANDLER(ReportHeartbeat, 1); - NODE_INFO_SERVICE_RPC_HANDLER(ReportBatchHeartbeat, 1); - NODE_INFO_SERVICE_RPC_HANDLER(GetResources, 1); - NODE_INFO_SERVICE_RPC_HANDLER(UpdateResources, 1); - NODE_INFO_SERVICE_RPC_HANDLER(DeleteResources, 1); + NODE_INFO_SERVICE_RPC_HANDLER(RegisterNode, SERVER_CALL_CONCURRENCY); + NODE_INFO_SERVICE_RPC_HANDLER(UnregisterNode, SERVER_CALL_CONCURRENCY); + NODE_INFO_SERVICE_RPC_HANDLER(GetAllNodeInfo, SERVER_CALL_CONCURRENCY); + NODE_INFO_SERVICE_RPC_HANDLER(ReportHeartbeat, SERVER_CALL_CONCURRENCY); + NODE_INFO_SERVICE_RPC_HANDLER(ReportBatchHeartbeat, SERVER_CALL_CONCURRENCY); + NODE_INFO_SERVICE_RPC_HANDLER(GetResources, SERVER_CALL_CONCURRENCY); + NODE_INFO_SERVICE_RPC_HANDLER(UpdateResources, SERVER_CALL_CONCURRENCY); + NODE_INFO_SERVICE_RPC_HANDLER(DeleteResources, SERVER_CALL_CONCURRENCY); } private: @@ -239,9 +241,9 @@ class ObjectInfoGrpcService : public GrpcService { const std::unique_ptr &cq, std::vector, int>> *server_call_factories_and_concurrencies) override { - OBJECT_INFO_SERVICE_RPC_HANDLER(GetObjectLocations, 1); - OBJECT_INFO_SERVICE_RPC_HANDLER(AddObjectLocation, 1); - OBJECT_INFO_SERVICE_RPC_HANDLER(RemoveObjectLocation, 1); + OBJECT_INFO_SERVICE_RPC_HANDLER(GetObjectLocations, SERVER_CALL_CONCURRENCY); + OBJECT_INFO_SERVICE_RPC_HANDLER(AddObjectLocation, SERVER_CALL_CONCURRENCY); + OBJECT_INFO_SERVICE_RPC_HANDLER(RemoveObjectLocation, SERVER_CALL_CONCURRENCY); } private: @@ -291,11 +293,11 @@ class TaskInfoGrpcService : public GrpcService { 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); - TASK_INFO_SERVICE_RPC_HANDLER(AddTaskLease, 1); - TASK_INFO_SERVICE_RPC_HANDLER(AttemptTaskReconstruction, 1); + TASK_INFO_SERVICE_RPC_HANDLER(AddTask, SERVER_CALL_CONCURRENCY); + TASK_INFO_SERVICE_RPC_HANDLER(GetTask, SERVER_CALL_CONCURRENCY); + TASK_INFO_SERVICE_RPC_HANDLER(DeleteTasks, SERVER_CALL_CONCURRENCY); + TASK_INFO_SERVICE_RPC_HANDLER(AddTaskLease, SERVER_CALL_CONCURRENCY); + TASK_INFO_SERVICE_RPC_HANDLER(AttemptTaskReconstruction, SERVER_CALL_CONCURRENCY); } private: @@ -331,7 +333,7 @@ class StatsGrpcService : public GrpcService { const std::unique_ptr &cq, std::vector, int>> *server_call_factories_and_concurrencies) override { - STATS_SERVICE_RPC_HANDLER(AddProfileData, 1); + STATS_SERVICE_RPC_HANDLER(AddProfileData, SERVER_CALL_CONCURRENCY); } private: @@ -367,7 +369,7 @@ class ErrorInfoGrpcService : public GrpcService { const std::unique_ptr &cq, std::vector, int>> *server_call_factories_and_concurrencies) override { - ERROR_INFO_SERVICE_RPC_HANDLER(ReportJobError, 1); + ERROR_INFO_SERVICE_RPC_HANDLER(ReportJobError, SERVER_CALL_CONCURRENCY); } private: @@ -403,7 +405,7 @@ class WorkerInfoGrpcService : public GrpcService { const std::unique_ptr &cq, std::vector, int>> *server_call_factories_and_concurrencies) override { - WORKER_INFO_SERVICE_RPC_HANDLER(ReportWorkerFailure, 1); + WORKER_INFO_SERVICE_RPC_HANDLER(ReportWorkerFailure, SERVER_CALL_CONCURRENCY); } private: diff --git a/src/ray/util/sequencer.h b/src/ray/util/sequencer.h new file mode 100644 index 000000000..952cd7094 --- /dev/null +++ b/src/ray/util/sequencer.h @@ -0,0 +1,67 @@ +#ifndef RAY_UTIL_SEQUENCER_H_ +#define RAY_UTIL_SEQUENCER_H_ + +#include +#include +#include +#include "absl/synchronization/mutex.h" + +namespace ray { + +/// This callback is used to notify when a operation completes. +using SequencerDoneCallback = std::function; + +/// \class Sequencer +/// Sequencer guarantees that all operations with the same key are sequenced. +/// This class is thread safe. +template +class Sequencer { + public: + /// This function is used to ask the sequencer to execute the given operation. + /// The sequencer guarantees that all operations with the same key are sequenced. + /// + /// \param key The key of operation. + /// \param operation The operation to be called. + void Post(KEY key, std::function operation) { + mutex_.Lock(); + pending_operations_[key].push_back(operation); + int queue_size = pending_operations_[key].size(); + mutex_.Unlock(); + + if (1 == queue_size) { + auto done_callback = [this, key]() { PostExecute(key); }; + operation(done_callback); + } + } + + private: + /// This function is used when a operation completes. + /// If the sequencer has operations with the same key, we will execute next operation. + /// + /// \param key The key of operation. + void PostExecute(const KEY key) { + mutex_.Lock(); + pending_operations_[key].pop_front(); + if (pending_operations_[key].empty()) { + pending_operations_.erase(key); + mutex_.Unlock(); + } else { + auto operation = pending_operations_[key].front(); + mutex_.Unlock(); + + auto done_callback = [this, key]() { PostExecute(key); }; + operation(done_callback); + } + } + + // Mutex to protect the pending_operations_ field. + absl::Mutex mutex_; + + std::unordered_map>> + pending_operations_ GUARDED_BY(mutex_); +}; + +} // namespace ray + +#endif // RAY_UTIL_SEQUENCER_H_ diff --git a/src/ray/util/sequencer_test.cc b/src/ray/util/sequencer_test.cc new file mode 100644 index 000000000..b8e42121d --- /dev/null +++ b/src/ray/util/sequencer_test.cc @@ -0,0 +1,37 @@ +#include "ray/util/sequencer.h" +#include +#include "gtest/gtest.h" +#include "ray/util/logging.h" + +namespace ray { + +TEST(SequencerTest, ExecuteOrderedTest) { + Sequencer sequencer; + std::deque queue; + int key = 1; + int size = 100; + for (int index = 0; index < size; ++index) { + auto operation = [index, &queue](SequencerDoneCallback done_callback) { + usleep(1000); + queue.push_back(index); + done_callback(); + }; + sequencer.Post(key, operation); + } + + while (queue.size() < (size_t)size) { + usleep(1000); + } + + for (int index = 0; index < size; ++index) { + ASSERT_EQ(queue.front(), index); + queue.pop_front(); + } +} + +} // namespace ray + +int main(int argc, char **argv) { + ::testing::InitGoogleTest(&argc, argv); + return RUN_ALL_TESTS(); +} \ No newline at end of file