diff --git a/src/ray/gcs/subscription_executor.cc b/src/ray/gcs/subscription_executor.cc index c55660c3c..9b2ff20c9 100644 --- a/src/ray/gcs/subscription_executor.cc +++ b/src/ray/gcs/subscription_executor.cc @@ -10,19 +10,39 @@ Status SubscriptionExecutor::AsyncSubscribe( const StatusCallback &done) { // TODO(micafan) Optimize the lock when necessary. // Consider avoiding locking in single-threaded processes. - std::lock_guard lock(mutex_); + std::unique_lock lock(mutex_); if (subscribe_all_callback_ != nullptr) { RAY_LOG(DEBUG) << "Duplicate subscription! Already subscribed to all elements."; return Status::Invalid("Duplicate subscription!"); } - if (registered_) { + if (registration_status_ != RegistrationStatus::kNotRegistered) { if (subscribe != nullptr) { RAY_LOG(DEBUG) << "Duplicate subscription! Already subscribed to specific elements" ", can't subscribe to all elements."; return Status::Invalid("Duplicate subscription!"); } + } + + if (registration_status_ == RegistrationStatus::kRegistered) { + // Already registered to GCS, just invoke the `done` callback. + lock.unlock(); + if (done != nullptr) { + done(Status::OK()); + } + return Status::OK(); + } + + // Registration to GCS is not finished yet, add the `done` callback to the pending list + // to be invoked when registration is done. + if (done != nullptr) { + pending_subscriptions_.emplace_back(done); + } + + // If there's another registration request that's already on-going, then wait for it + // to finish. + if (registration_status_ == RegistrationStatus::kRegistering) { return Status::OK(); } @@ -37,7 +57,7 @@ Status SubscriptionExecutor::AsyncSubscribe( SubscribeCallback sub_one_callback = nullptr; SubscribeCallback sub_all_callback = nullptr; { - std::lock_guard lock(mutex_); + std::unique_lock lock(mutex_); const auto it = id_to_callback_map_.find(id); if (it != id_to_callback_map_.end()) { sub_one_callback = it->second; @@ -53,15 +73,23 @@ Status SubscriptionExecutor::AsyncSubscribe( } }; - auto on_done = [done](RedisGcsClient *client) { - if (done != nullptr) { - done(Status::OK()); + auto on_done = [this](RedisGcsClient *client) { + std::list pending_callbacks; + { + std::unique_lock lock(mutex_); + registration_status_ = RegistrationStatus::kRegistered; + pending_callbacks.swap(pending_subscriptions_); + RAY_CHECK(pending_subscriptions_.empty()); + } + + for (const auto &callback : pending_callbacks) { + callback(Status::OK()); } }; Status status = table_.Subscribe(JobID::Nil(), client_id, on_subscribe, on_done); if (status.ok()) { - registered_ = true; + registration_status_ = RegistrationStatus::kRegistering; subscribe_all_callback_ = subscribe; } @@ -72,35 +100,49 @@ template Status SubscriptionExecutor::AsyncSubscribe( const ClientID &client_id, const ID &id, const SubscribeCallback &subscribe, const StatusCallback &done) { - Status status = AsyncSubscribe(client_id, nullptr, nullptr); - if (!status.ok()) { - return status; - } + RAY_CHECK(client_id != ClientID::Nil()); - auto on_done = [this, done, id](Status status) { - if (!status.ok()) { - std::lock_guard lock(mutex_); - id_to_callback_map_.erase(id); - } - if (done != nullptr) { - done(status); + // NOTE(zhijunfu): `Subscribe` and other operations use different redis contexts, + // thus we need to call `RequestNotifications` in the Subscribe callback to ensure + // it's processed after the `Subscribe` request. Otherwise if `RequestNotifications` + // is processed first we will miss the initial notification. + auto on_subscribe_done = [this, client_id, id, subscribe, done](Status status) { + auto on_request_notification_done = [this, done, id](Status status) { + if (!status.ok()) { + std::unique_lock lock(mutex_); + id_to_callback_map_.erase(id); + } + if (done != nullptr) { + done(status); + } + }; + + { + std::unique_lock lock(mutex_); + status = table_.RequestNotifications(JobID::Nil(), id, client_id, + on_request_notification_done); + if (!status.ok()) { + id_to_callback_map_.erase(id); + } } }; { - std::lock_guard lock(mutex_); + std::unique_lock lock(mutex_); const auto it = id_to_callback_map_.find(id); if (it != id_to_callback_map_.end()) { RAY_LOG(DEBUG) << "Duplicate subscription to id " << id << " client_id " << client_id; return Status::Invalid("Duplicate subscription to element!"); } - status = table_.RequestNotifications(JobID::Nil(), id, client_id, on_done); - if (status.ok()) { - id_to_callback_map_[id] = subscribe; - } + id_to_callback_map_[id] = subscribe; } + auto status = AsyncSubscribe(client_id, nullptr, on_subscribe_done); + if (!status.ok()) { + std::unique_lock lock(mutex_); + id_to_callback_map_.erase(id); + } return status; } @@ -108,7 +150,7 @@ template Status SubscriptionExecutor::AsyncUnsubscribe( const ClientID &client_id, const ID &id, const StatusCallback &done) { { - std::lock_guard lock(mutex_); + std::unique_lock lock(mutex_); const auto it = id_to_callback_map_.find(id); if (it == id_to_callback_map_.end()) { RAY_LOG(DEBUG) << "Invalid Unsubscribe! id " << id << " client_id " << client_id; @@ -118,7 +160,7 @@ Status SubscriptionExecutor::AsyncUnsubscribe( auto on_done = [this, id, done](Status status) { if (status.ok()) { - std::lock_guard lock(mutex_); + std::unique_lock lock(mutex_); const auto it = id_to_callback_map_.find(id); if (it != id_to_callback_map_.end()) { id_to_callback_map_.erase(it); diff --git a/src/ray/gcs/subscription_executor.h b/src/ray/gcs/subscription_executor.h index 167e1f274..1221b70e0 100644 --- a/src/ray/gcs/subscription_executor.h +++ b/src/ray/gcs/subscription_executor.h @@ -2,6 +2,7 @@ #define RAY_GCS_SUBSCRIPTION_EXECUTOR_H #include +#include #include #include "ray/gcs/callback.h" #include "ray/gcs/tables.h" @@ -67,8 +68,18 @@ class SubscriptionExecutor { std::mutex mutex_; + enum class RegistrationStatus : uint8_t { + kNotRegistered, + kRegistering, + kRegistered, + }; + /// Whether successfully registered subscription to GCS. - bool registered_{false}; + RegistrationStatus registration_status_{RegistrationStatus::kNotRegistered}; + + /// List of subscriptions before registration to GCS is done, these callbacks + /// will be called when the registration to GCS finishes. + std::list pending_subscriptions_; /// Subscribe Callback of all elements. SubscribeCallback subscribe_all_callback_{nullptr}; diff --git a/src/ray/gcs/subscription_executor_test.cc b/src/ray/gcs/subscription_executor_test.cc index 6f477173e..5e7367b6e 100644 --- a/src/ray/gcs/subscription_executor_test.cc +++ b/src/ray/gcs/subscription_executor_test.cc @@ -95,25 +95,6 @@ TEST_F(SubscriptionExecutorTest, SubscribeAllTest) { WaitPendingDone(sub_pending_count_, wait_pending_timeout_); } -TEST_F(SubscriptionExecutorTest, SubscribeOneTest) { - Status status; - for (const auto &item : id_to_data_) { - ++do_sub_pending_count_; - status = actor_sub_executor_->AsyncSubscribe(ClientID::Nil(), item.first, subscribe_, - sub_done_); - ASSERT_TRUE(status.ok()); - } - WaitPendingDone(do_sub_pending_count_, wait_pending_timeout_); - sub_pending_count_ = id_to_data_.size(); - AsyncRegisterActorToGcs(); - for (const auto &item : id_to_data_) { - status = actor_sub_executor_->AsyncSubscribe(ClientID::Nil(), item.first, subscribe_, - sub_done_); - ASSERT_TRUE(status.IsInvalid()); - } - WaitPendingDone(sub_pending_count_, wait_pending_timeout_); -} - TEST_F(SubscriptionExecutorTest, SubscribeOneWithClientIDTest) { const auto &item = id_to_data_.begin(); ++do_sub_pending_count_; @@ -124,6 +105,24 @@ TEST_F(SubscriptionExecutorTest, SubscribeOneWithClientIDTest) { ASSERT_TRUE(status.ok()); AsyncRegisterActorToGcs(); WaitPendingDone(sub_pending_count_, wait_pending_timeout_); + status = actor_sub_executor_->AsyncSubscribe(ClientID::FromRandom(), item->first, + subscribe_, sub_done_); + ASSERT_TRUE(status.IsInvalid()); +} + +TEST_F(SubscriptionExecutorTest, SubscribeOneAfterActorRegistrationWithClientIDTest) { + const auto &item = id_to_data_.begin(); + ++do_sub_pending_count_; + ++sub_pending_count_; + AsyncRegisterActorToGcs(); + Status status = actor_sub_executor_->AsyncSubscribe(ClientID::FromRandom(), item->first, + subscribe_, sub_done_); + WaitPendingDone(do_sub_pending_count_, wait_pending_timeout_); + ASSERT_TRUE(status.ok()); + WaitPendingDone(sub_pending_count_, wait_pending_timeout_); + status = actor_sub_executor_->AsyncSubscribe(ClientID::FromRandom(), item->first, + subscribe_, sub_done_); + ASSERT_TRUE(status.IsInvalid()); } TEST_F(SubscriptionExecutorTest, SubscribeAllAndSubscribeOneTest) { @@ -133,8 +132,8 @@ TEST_F(SubscriptionExecutorTest, SubscribeAllAndSubscribeOneTest) { ASSERT_TRUE(status.ok()); WaitPendingDone(do_sub_pending_count_, wait_pending_timeout_); for (const auto &item : id_to_data_) { - status = actor_sub_executor_->AsyncSubscribe(ClientID::Nil(), item.first, subscribe_, - sub_done_); + status = actor_sub_executor_->AsyncSubscribe(ClientID::FromRandom(), item.first, + subscribe_, sub_done_); ASSERT_FALSE(status.ok()); } sub_pending_count_ = id_to_data_.size(); @@ -143,51 +142,48 @@ TEST_F(SubscriptionExecutorTest, SubscribeAllAndSubscribeOneTest) { } TEST_F(SubscriptionExecutorTest, UnsubscribeTest) { + ClientID client_id = ClientID::FromRandom(); Status status; for (const auto &item : id_to_data_) { - status = - actor_sub_executor_->AsyncUnsubscribe(ClientID::Nil(), item.first, unsub_done_); + status = actor_sub_executor_->AsyncUnsubscribe(client_id, item.first, unsub_done_); ASSERT_TRUE(status.IsInvalid()); } for (const auto &item : id_to_data_) { ++do_sub_pending_count_; - status = actor_sub_executor_->AsyncSubscribe(ClientID::Nil(), item.first, subscribe_, - sub_done_); + status = + actor_sub_executor_->AsyncSubscribe(client_id, item.first, subscribe_, sub_done_); ASSERT_TRUE(status.ok()); } WaitPendingDone(do_sub_pending_count_, wait_pending_timeout_); for (const auto &item : id_to_data_) { ++do_unsub_pending_count_; - status = - actor_sub_executor_->AsyncUnsubscribe(ClientID::Nil(), item.first, unsub_done_); + status = actor_sub_executor_->AsyncUnsubscribe(client_id, item.first, unsub_done_); ASSERT_TRUE(status.ok()); } WaitPendingDone(do_unsub_pending_count_, wait_pending_timeout_); for (const auto &item : id_to_data_) { - status = - actor_sub_executor_->AsyncUnsubscribe(ClientID::Nil(), item.first, unsub_done_); + status = actor_sub_executor_->AsyncUnsubscribe(client_id, item.first, unsub_done_); ASSERT_TRUE(!status.ok()); } for (const auto &item : id_to_data_) { ++do_sub_pending_count_; - status = actor_sub_executor_->AsyncSubscribe(ClientID::Nil(), item.first, subscribe_, - sub_done_); + status = + actor_sub_executor_->AsyncSubscribe(client_id, item.first, subscribe_, sub_done_); ASSERT_TRUE(status.ok()); } WaitPendingDone(do_sub_pending_count_, wait_pending_timeout_); for (const auto &item : id_to_data_) { ++do_unsub_pending_count_; - status = - actor_sub_executor_->AsyncUnsubscribe(ClientID::Nil(), item.first, unsub_done_); + status = actor_sub_executor_->AsyncUnsubscribe(client_id, item.first, unsub_done_); ASSERT_TRUE(status.ok()); } WaitPendingDone(do_unsub_pending_count_, wait_pending_timeout_); for (const auto &item : id_to_data_) { ++do_sub_pending_count_; - status = actor_sub_executor_->AsyncSubscribe(ClientID::Nil(), item.first, subscribe_, - sub_done_); + status = + actor_sub_executor_->AsyncSubscribe(client_id, item.first, subscribe_, sub_done_); ASSERT_TRUE(status.ok()); } WaitPendingDone(do_sub_pending_count_, wait_pending_timeout_);