diff --git a/BUILD.bazel b/BUILD.bazel index 8782dbdf8..669d11a67 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -956,8 +956,8 @@ cc_test( ) cc_test( - name = "task_dependency_manager_test", - srcs = ["src/ray/raylet/task_dependency_manager_test.cc"], + name = "dependency_manager_test", + srcs = ["src/ray/raylet/dependency_manager_test.cc"], copts = COPTS, deps = [ ":raylet_lib", diff --git a/src/ray/common/common_protocol.h b/src/ray/common/common_protocol.h index 8ff224922..4ac5b5403 100644 --- a/src/ray/common/common_protocol.h +++ b/src/ray/common/common_protocol.h @@ -20,6 +20,7 @@ #include "ray/common/id.h" #include "ray/util/logging.h" +#include "src/ray/protobuf/common.pb.h" /// Convert an unique ID to a flatbuffer string. /// @@ -201,3 +202,24 @@ to_flatbuf(flatbuffers::FlatBufferBuilder &fbb, const std::unordered_set &id } return fbb.CreateVector(results); } + +static inline ray::rpc::ObjectReference ObjectIdToRef( + const ray::ObjectID &object_id, const ray::rpc::Address owner_address) { + ray::rpc::ObjectReference ref; + ref.set_object_id(object_id.Binary()); + ref.mutable_owner_address()->CopyFrom(owner_address); + return ref; +} + +static inline ray::ObjectID ObjectRefToId(const ray::rpc::ObjectReference &object_ref) { + return ray::ObjectID::FromBinary(object_ref.object_id()); +} + +static inline std::vector ObjectRefsToIds( + const std::vector &object_refs) { + std::vector object_ids; + for (const auto &ref : object_refs) { + object_ids.push_back(ObjectRefToId(ref)); + } + return object_ids; +} diff --git a/src/ray/object_manager/object_manager.cc b/src/ray/object_manager/object_manager.cc index 760909dc0..8ae4d723e 100644 --- a/src/ray/object_manager/object_manager.cc +++ b/src/ray/object_manager/object_manager.cc @@ -103,8 +103,11 @@ ObjectManager::ObjectManager(asio::io_service &main_service, const NodeID &self_ [this](const object_manager::protocol::ObjectInfoT &object_info) { HandleObjectAdded(object_info); }); - store_notification_->SubscribeObjDeleted( - [this](const ObjectID &oid) { NotifyDirectoryObjectDeleted(oid); }); + store_notification_->SubscribeObjDeleted([this](const ObjectID &oid) { + // TODO(swang): We may want to force the pull manager to fetch this object + // again, in case it was needed by an active pull request. + NotifyDirectoryObjectDeleted(oid); + }); // Start object manager rpc server and send & receive request threads StartRpcService(); @@ -169,10 +172,6 @@ void ObjectManager::HandleObjectAdded( } unfulfilled_push_requests_.erase(iter); } - - // The object is local, so we no longer need to Pull it from a remote - // manager. Cancel any outstanding Pull requests for this object. - CancelPull(object_id); } void ObjectManager::NotifyDirectoryObjectDeleted(const ObjectID &object_id) { @@ -198,13 +197,9 @@ ray::Status ObjectManager::SubscribeObjDeleted( return ray::Status::OK(); } -ray::Status ObjectManager::Pull(const ObjectID &object_id, - const rpc::Address &owner_address) { - if (!pull_manager_->Pull(object_id, owner_address)) { - // If we don't need to pull, the object is either already local or this is a duplicate - // request. - return Status::OK(); - } +uint64_t ObjectManager::Pull(const std::vector &object_refs) { + std::vector objects_to_locate; + auto request_id = pull_manager_->Pull(object_refs, &objects_to_locate); const auto &callback = [this](const ObjectID &object_id, const std::unordered_set &client_ids, @@ -212,12 +207,25 @@ ray::Status ObjectManager::Pull(const ObjectID &object_id, pull_manager_->OnLocationChange(object_id, client_ids, spilled_url); }; - // Subscribe to object notifications. A notification will be received every - // time the set of node IDs for the object changes. Notifications will also - // be received if the list of locations is empty. The set of node IDs has - // no ordering guarantee between notifications. - return object_directory_->SubscribeObjectLocations(object_directory_pull_callback_id_, - object_id, owner_address, callback); + for (const auto &ref : objects_to_locate) { + // Subscribe to object notifications. A notification will be received every + // time the set of node IDs for the object changes. Notifications will also + // be received if the list of locations is empty. The set of node IDs has + // no ordering guarantee between notifications. + auto object_id = ObjectRefToId(ref); + RAY_CHECK_OK(object_directory_->SubscribeObjectLocations( + object_directory_pull_callback_id_, object_id, ref.owner_address(), callback)); + } + + return request_id; +} + +void ObjectManager::CancelPull(uint64_t request_id) { + const auto objects_to_cancel = pull_manager_->CancelPull(request_id); + for (const auto &object_id : objects_to_cancel) { + RAY_CHECK_OK(object_directory_->UnsubscribeObjectLocations( + object_directory_pull_callback_id_, object_id)); + } } void ObjectManager::SendPullRequest(const ObjectID &object_id, const NodeID &client_id) { @@ -426,16 +434,6 @@ void ObjectManager::SendObjectChunk(const UniqueID &push_id, const ObjectID &obj buffer_pool_.ReleaseGetChunk(object_id, chunk_info.chunk_index); } -void ObjectManager::CancelPull(const ObjectID &object_id) { - if (!pull_manager_->CancelPull(object_id)) { - // We weren't tracking a pull request for this object, so there is nothing to cancel. - return; - } - - RAY_CHECK_OK(object_directory_->UnsubscribeObjectLocations( - object_directory_pull_callback_id_, object_id)); -} - ray::Status ObjectManager::Wait( const std::vector &object_ids, const std::unordered_map &owner_addresses, int64_t timeout_ms, diff --git a/src/ray/object_manager/object_manager.h b/src/ray/object_manager/object_manager.h index fdca6a190..3e793e21c 100644 --- a/src/ray/object_manager/object_manager.h +++ b/src/ray/object_manager/object_manager.h @@ -43,6 +43,7 @@ #include "ray/object_manager/push_manager.h" #include "ray/rpc/object_manager/object_manager_client.h" #include "ray/rpc/object_manager/object_manager_server.h" +#include "src/ray/protobuf/common.pb.h" namespace ray { @@ -95,9 +96,8 @@ class ObjectStoreRunner { class ObjectManagerInterface { public: - virtual ray::Status Pull(const ObjectID &object_id, - const rpc::Address &owner_address) = 0; - virtual void CancelPull(const ObjectID &object_id) = 0; + virtual uint64_t Pull(const std::vector &object_refs) = 0; + virtual void CancelPull(uint64_t request_id) = 0; virtual ~ObjectManagerInterface(){}; }; @@ -238,18 +238,19 @@ class ObjectManager : public ObjectManagerInterface, /// \return Void. void Push(const ObjectID &object_id, const NodeID &node_id); - /// Pull an object from NodeID. + /// Pull a bundle of objects. This will attempt to make all objects in the + /// bundle local until the request is canceled with the returned ID. /// - /// \param object_id The object's object id. - /// \return Status of whether the pull request successfully initiated. - ray::Status Pull(const ObjectID &object_id, const rpc::Address &owner_address) override; + /// \param object_refs The bundle of objects that must be made local. + /// \return A request ID that can be used to cancel the request. + uint64_t Pull(const std::vector &object_refs) override; - /// Cancels all requests (Push/Pull) associated with the given ObjectID. This - /// method is idempotent. + /// Cancels the pull request with the given ID. This cancels any fetches for + /// objects that were passed to the original pull request, if no other pull + /// request requires them. /// - /// \param object_id The ObjectID. - /// \return Void. - void CancelPull(const ObjectID &object_id) override; + /// \param pull_request_id The request to cancel. + void CancelPull(uint64_t pull_request_id) override; /// Callback definition for wait. using WaitCallback = std::function &found, diff --git a/src/ray/object_manager/pull_manager.cc b/src/ray/object_manager/pull_manager.cc index 8ced0f51b..c9fa13177 100644 --- a/src/ray/object_manager/pull_manager.cc +++ b/src/ray/object_manager/pull_manager.cc @@ -1,5 +1,7 @@ #include "ray/object_manager/pull_manager.h" +#include "ray/common/common_protocol.h" + namespace ray { PullManager::PullManager( @@ -15,29 +17,56 @@ PullManager::PullManager( pull_timeout_ms_(pull_timeout_ms), gen_(std::chrono::high_resolution_clock::now().time_since_epoch().count()) {} -bool PullManager::Pull(const ObjectID &object_id, const rpc::Address &owner_address) { - RAY_LOG(DEBUG) << "Pull " - << " of object " << object_id; - // Check if object is already local. - if (object_is_local_(object_id)) { - RAY_LOG(DEBUG) << object_id << " attempted to pull an object that's already local."; - return false; - } - if (pull_requests_.find(object_id) != pull_requests_.end()) { - RAY_LOG(DEBUG) << object_id << " has inflight pull_requests, skipping."; - return false; +uint64_t PullManager::Pull(const std::vector &object_ref_bundle, + std::vector *objects_to_locate) { + auto bundle_it = pull_request_bundles_.emplace(next_req_id_++, object_ref_bundle).first; + RAY_LOG(DEBUG) << "Start pull request " << bundle_it->first; + + for (const auto &ref : object_ref_bundle) { + auto obj_id = ObjectRefToId(ref); + auto it = object_pull_requests_.find(obj_id); + if (it == object_pull_requests_.end()) { + RAY_LOG(DEBUG) << "Pull of object " << obj_id; + // We don't have an active pull for this object yet. Ask the caller to + // send us notifications about the object's location. + objects_to_locate->push_back(ref); + it = object_pull_requests_ + .emplace(obj_id, ObjectPullRequest(get_time_() + pull_timeout_ms_ / 1000)) + .first; + } + it->second.bundle_request_ids.insert(bundle_it->first); } - pull_requests_.emplace(object_id, PullRequest(get_time_() + pull_timeout_ms_ / 1000)); - return true; + return bundle_it->first; +} + +std::vector PullManager::CancelPull(uint64_t request_id) { + std::vector objects_to_cancel; + RAY_LOG(DEBUG) << "Cancel pull request " << request_id; + auto bundle_it = pull_request_bundles_.find(request_id); + RAY_CHECK(bundle_it != pull_request_bundles_.end()); + + for (const auto &ref : bundle_it->second) { + auto obj_id = ObjectRefToId(ref); + auto it = object_pull_requests_.find(obj_id); + RAY_CHECK(it != object_pull_requests_.end()); + RAY_CHECK(it->second.bundle_request_ids.erase(request_id)); + if (it->second.bundle_request_ids.empty()) { + object_pull_requests_.erase(it); + objects_to_cancel.push_back(obj_id); + } + } + + pull_request_bundles_.erase(bundle_it); + return objects_to_cancel; } void PullManager::OnLocationChange(const ObjectID &object_id, const std::unordered_set &client_ids, const std::string &spilled_url) { // Exit if the Pull request has already been fulfilled or canceled. - auto it = pull_requests_.find(object_id); - if (it == pull_requests_.end()) { + auto it = object_pull_requests_.find(object_id); + if (it == object_pull_requests_.end()) { return; } // Reset the list of clients that are now expected to have the object. @@ -45,28 +74,50 @@ void PullManager::OnLocationChange(const ObjectID &object_id, // we may end up sending a duplicate request to the same client as // before. it->second.client_locations = std::vector(client_ids.begin(), client_ids.end()); - if (!spilled_url.empty()) { - RAY_LOG(DEBUG) << "OnLocationChange " << spilled_url << " num clients " - << client_ids.size(); + it->second.spilled_url = spilled_url; + RAY_LOG(DEBUG) << "OnLocationChange " << spilled_url << " num clients " + << client_ids.size(); + + TryToMakeObjectLocal(object_id); +} + +void PullManager::TryToMakeObjectLocal(const ObjectID &object_id) { + if (object_is_local_(object_id)) { + return; + } + auto it = object_pull_requests_.find(object_id); + if (it == object_pull_requests_.end()) { + return; + } + + auto &request = it->second; + if (!request.spilled_url.empty()) { // Try to restore the spilled object. - restore_spilled_object_(object_id, spilled_url, + restore_spilled_object_(object_id, request.spilled_url, [this, object_id](const ray::Status &status) { // Fall back to fetching from another object manager. if (!status.ok()) { - TryPull(object_id); + PullFromRandomLocation(object_id); } }); } else { // New object locations were found, so begin trying to pull from a // client. This will be called every time a new client location // appears. - TryPull(object_id); + PullFromRandomLocation(object_id); } + + const auto time = get_time_(); + auto retry_timeout_len = (pull_timeout_ms_ / 1000.) * (1UL << request.num_retries); + request.next_pull_time = time + retry_timeout_len; + + // Bound the retry time at 10 * 1024 seconds. + request.num_retries = std::min(request.num_retries + 1, 10); } -void PullManager::TryPull(const ObjectID &object_id) { - auto it = pull_requests_.find(object_id); - if (it == pull_requests_.end()) { +void PullManager::PullFromRandomLocation(const ObjectID &object_id) { + auto it = object_pull_requests_.find(object_id); + if (it == object_pull_requests_.end()) { return; } @@ -111,36 +162,20 @@ void PullManager::TryPull(const ObjectID &object_id) { RAY_LOG(DEBUG) << "Sending pull request from " << self_node_id_ << " to " << node_id << " of object " << object_id; - const auto time = get_time_(); - auto &request = it->second; - auto retry_timeout_len = (pull_timeout_ms_ / 1000.) * (1UL << request.num_retries); - request.next_pull_time = time + retry_timeout_len; send_pull_request_(object_id, node_id); } -bool PullManager::CancelPull(const ObjectID &object_id) { - auto it = pull_requests_.find(object_id); - if (it == pull_requests_.end()) { - return false; - } - - pull_requests_.erase(it); - return true; -} - void PullManager::Tick() { - for (auto &pair : pull_requests_) { + for (auto &pair : object_pull_requests_) { const auto &object_id = pair.first; auto &request = pair.second; const auto time = get_time_(); if (time >= request.next_pull_time) { - TryPull(object_id); - // Bound the retry time at 10 * 1024 seconds. - request.num_retries = std::min(request.num_retries + 1, 10); + TryToMakeObjectLocal(object_id); } } } -int PullManager::NumActiveRequests() const { return pull_requests_.size(); } +int PullManager::NumActiveRequests() const { return object_pull_requests_.size(); } } // namespace ray diff --git a/src/ray/object_manager/pull_manager.h b/src/ray/object_manager/pull_manager.h index f312af17a..33710a18e 100644 --- a/src/ray/object_manager/pull_manager.h +++ b/src/ray/object_manager/pull_manager.h @@ -42,33 +42,32 @@ class PullManager { const RestoreSpilledObjectCallback restore_spilled_object, const std::function get_time, int pull_timeout_ms); - /// Begin a new pull request if necessary. + /// Begin a new pull request for a bundle of objects. /// - /// \param object_id The object id to pull. - /// \param owner_address The owner of the object. - /// - /// \return True if a new pull request was necessary. If true, the caller should - /// subscribe to new locations of the object, and call OnLocationChange when necessary. - bool Pull(const ObjectID &object_id, const rpc::Address &owner_address); + /// \param object_refs The bundle of objects that must be made local. + /// \param objects_to_locate The objects whose new locations the caller + /// should subscribe to, and call OnLocationChange for. + /// \return A request ID that can be used to cancel the request. + uint64_t Pull(const std::vector &object_ref_bundle, + std::vector *objects_to_locate); /// Called when the available locations for a given object change. /// /// \param object_id The ID of the object which is now available in a new location. /// \param client_ids The new set of nodes that the object is available on. Not - /// necessarily a super or subset of the previously available nodes. \param spilled_url - /// The location of the object if it was spilled. If non-empty, the object may no longer - /// be on any node. + /// necessarily a super or subset of the previously available nodes. + /// \param spilled_url The location of the object if it was spilled. If + /// non-empty, the object may no longer be on any node. void OnLocationChange(const ObjectID &object_id, const std::unordered_set &client_ids, const std::string &spilled_url); - /// Cancel an existing pull request if necessary. + /// Cancel an existing pull request. /// - /// \param object_id The object id that no longer needs to be pulled. - /// - /// \return True if a pull was cancelled. If there was no pending pull request for the - /// object this method may return false. - bool CancelPull(const ObjectID &object_id); + /// \param request_id The request ID returned by Pull that should be canceled. + /// \return The objects for which the caller should stop subscribing to + /// locations. + std::vector CancelPull(uint64_t request_id); /// Called when the retry timer fires. If this fires, the pull manager may try to pull /// existing objects from other nodes if necessary. @@ -79,14 +78,32 @@ class PullManager { private: /// A helper structure for tracking information about each ongoing object pull. - struct PullRequest { - PullRequest(double first_retry_time) - : client_locations(), next_pull_time(first_retry_time), num_retries(0) {} + struct ObjectPullRequest { + ObjectPullRequest(double first_retry_time) + : client_locations(), + spilled_url(), + next_pull_time(first_retry_time), + num_retries(0), + bundle_request_ids() {} std::vector client_locations; + std::string spilled_url; double next_pull_time; uint8_t num_retries; + absl::flat_hash_set bundle_request_ids; }; + /// Try to make an object local, by restoring the object from external + /// storage or by fetching the object from one of its expected client + /// locations. This does nothing if the object is not needed by any pull + /// request or if it is already local. This also sets a timeout for when to + /// make the next attempt to make the object local. + void TryToMakeObjectLocal(const ObjectID &object_id); + + /// Try to Pull an object from one of its expected client locations. If there + /// are more client locations to try after this attempt, then this method + /// will try each of the other clients in succession. + void PullFromRandomLocation(const ObjectID &object_id); + /// See the constructor's arguments. NodeID self_node_id_; const std::function object_is_local_; @@ -95,22 +112,17 @@ class PullManager { const std::function get_time_; uint64_t pull_timeout_ms_; + /// The next ID to assign to a bundle pull request, so that the caller can + /// cancel. Start at 1 because 0 means null. + uint64_t next_req_id_ = 1; + + std::unordered_map> pull_request_bundles_; + /// The objects that this object manager is currently trying to fetch from /// remote object managers. - std::unordered_map pull_requests_; + std::unordered_map object_pull_requests_; /// Internally maintained random number generator. std::mt19937_64 gen_; - - /// Try to Pull an object from one of its expected client locations. If there - /// are more client locations to try after this attempt, then this method - /// will try each of the other clients in succession, with a timeout between - /// each attempt. If the object is received or if the Pull is Canceled before - /// the timeout, then no more Pull requests for this object will be sent - /// to other node managers until TryPull is called again. - /// - /// \param object_id The object's object id. - /// \return Void. - void TryPull(const ObjectID &object_id); }; } // namespace ray diff --git a/src/ray/object_manager/test/object_manager_stress_test.cc b/src/ray/object_manager/test/object_manager_stress_test.cc index 018bc357b..6a55d6180 100644 --- a/src/ray/object_manager/test/object_manager_stress_test.cc +++ b/src/ray/object_manager/test/object_manager_stress_test.cc @@ -18,6 +18,7 @@ #include #include "gtest/gtest.h" +#include "ray/common/common_protocol.h" #include "ray/common/status.h" #include "ray/common/test_util.h" #include "ray/gcs/gcs_client/service_based_gcs_client.h" @@ -338,8 +339,6 @@ class StressTestObjectManager : public TestObjectManagerBase { NodeID node_id_1 = gcs_client_1->Nodes().GetSelfId(); NodeID node_id_2 = gcs_client_2->Nodes().GetSelfId(); - ray::Status status = ray::Status::OK(); - if (transfer_pattern == TransferPattern::BIDIRECTIONAL_PULL || transfer_pattern == TransferPattern::BIDIRECTIONAL_PUSH || transfer_pattern == TransferPattern::BIDIRECTIONAL_PULL_VARIABLE_DATA_SIZE) { @@ -374,21 +373,25 @@ class StressTestObjectManager : public TestObjectManagerBase { case TransferPattern::PULL_A_B: { for (int i = -1; ++i < num_trials;) { ObjectID oid1 = WriteDataToClient(client1, data_size); - status = server2->object_manager_.Pull(oid1, rpc::Address()); + static_cast( + server2->object_manager_.Pull({ObjectIdToRef(oid1, rpc::Address())})); } } break; case TransferPattern::PULL_B_A: { for (int i = -1; ++i < num_trials;) { ObjectID oid2 = WriteDataToClient(client2, data_size); - status = server1->object_manager_.Pull(oid2, rpc::Address()); + static_cast( + server1->object_manager_.Pull({ObjectIdToRef(oid2, rpc::Address())})); } } break; case TransferPattern::BIDIRECTIONAL_PULL: { for (int i = -1; ++i < num_trials;) { ObjectID oid1 = WriteDataToClient(client1, data_size); - status = server2->object_manager_.Pull(oid1, rpc::Address()); + static_cast( + server2->object_manager_.Pull({ObjectIdToRef(oid1, rpc::Address())})); ObjectID oid2 = WriteDataToClient(client2, data_size); - status = server1->object_manager_.Pull(oid2, rpc::Address()); + static_cast( + server1->object_manager_.Pull({ObjectIdToRef(oid2, rpc::Address())})); } } break; case TransferPattern::BIDIRECTIONAL_PULL_VARIABLE_DATA_SIZE: { @@ -397,9 +400,11 @@ class StressTestObjectManager : public TestObjectManagerBase { std::uniform_int_distribution<> dis(1, 50); for (int i = -1; ++i < num_trials;) { ObjectID oid1 = WriteDataToClient(client1, data_size + dis(gen)); - status = server2->object_manager_.Pull(oid1, rpc::Address()); + static_cast( + server2->object_manager_.Pull({ObjectIdToRef(oid1, rpc::Address())})); ObjectID oid2 = WriteDataToClient(client2, data_size + dis(gen)); - status = server1->object_manager_.Pull(oid2, rpc::Address()); + static_cast( + server1->object_manager_.Pull({ObjectIdToRef(oid2, rpc::Address())})); } } break; default: { diff --git a/src/ray/object_manager/test/pull_manager_test.cc b/src/ray/object_manager/test/pull_manager_test.cc index fb7b1c1c2..21e41f874 100644 --- a/src/ray/object_manager/test/pull_manager_test.cc +++ b/src/ray/object_manager/test/pull_manager_test.cc @@ -1,11 +1,15 @@ #include "ray/object_manager/pull_manager.h" +#include "gmock/gmock.h" #include "gtest/gtest.h" +#include "ray/common/common_protocol.h" #include "ray/common/test_util.h" namespace ray { +using ::testing::ElementsAre; + class PullManagerTest : public ::testing::Test { public: PullManagerTest() @@ -33,28 +37,42 @@ class PullManagerTest : public ::testing::Test { PullManager pull_manager_; }; +std::vector CreateObjectRefs(int num_objs) { + std::vector refs; + for (int i = 0; i < num_objs; i++) { + ObjectID obj = ObjectID::FromRandom(); + rpc::ObjectReference ref; + ref.set_object_id(obj.Binary()); + refs.push_back(ref); + } + return refs; +} + TEST_F(PullManagerTest, TestStaleSubscription) { - ObjectID obj1 = ObjectID::FromRandom(); - rpc::Address addr1; + auto refs = CreateObjectRefs(1); + auto oid = ObjectRefsToIds(refs)[0]; ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); - pull_manager_.Pull(obj1, addr1); + std::vector objects_to_locate; + auto req_id = pull_manager_.Pull(refs, &objects_to_locate); + ASSERT_EQ(ObjectRefsToIds(objects_to_locate), ObjectRefsToIds(refs)); ASSERT_EQ(pull_manager_.NumActiveRequests(), 1); std::unordered_set client_ids; - pull_manager_.OnLocationChange(obj1, client_ids, ""); + pull_manager_.OnLocationChange(oid, client_ids, ""); // There are no client ids to pull from. ASSERT_EQ(num_send_pull_request_calls_, 0); ASSERT_EQ(num_restore_spilled_object_calls_, 0); - pull_manager_.CancelPull(obj1); + auto objects_to_cancel = pull_manager_.CancelPull(req_id); + ASSERT_EQ(objects_to_cancel, ObjectRefsToIds(refs)); ASSERT_EQ(num_send_pull_request_calls_, 0); ASSERT_EQ(num_restore_spilled_object_calls_, 0); ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); client_ids.insert(NodeID::FromRandom()); - pull_manager_.OnLocationChange(obj1, client_ids, ""); + pull_manager_.OnLocationChange(oid, client_ids, ""); // Now we're getting a notification about an object that was already cancelled. ASSERT_EQ(num_send_pull_request_calls_, 0); @@ -63,10 +81,13 @@ TEST_F(PullManagerTest, TestStaleSubscription) { } TEST_F(PullManagerTest, TestRestoreSpilledObject) { - ObjectID obj1 = ObjectID::FromRandom(); + auto refs = CreateObjectRefs(1); + auto obj1 = ObjectRefsToIds(refs)[0]; rpc::Address addr1; ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); - pull_manager_.Pull(obj1, addr1); + std::vector objects_to_locate; + auto req_id = pull_manager_.Pull(refs, &objects_to_locate); + ASSERT_EQ(ObjectRefsToIds(objects_to_locate), ObjectRefsToIds(refs)); ASSERT_EQ(pull_manager_.NumActiveRequests(), 1); std::unordered_set client_ids; @@ -84,15 +105,25 @@ TEST_F(PullManagerTest, TestRestoreSpilledObject) { ASSERT_EQ(num_send_pull_request_calls_, 0); ASSERT_EQ(num_restore_spilled_object_calls_, 2); - pull_manager_.CancelPull(obj1); + // Don't restore an object if it's local. + object_is_local_ = true; + num_restore_spilled_object_calls_ = 0; + pull_manager_.OnLocationChange(obj1, client_ids, "remote_url/foo/bar"); + ASSERT_EQ(num_restore_spilled_object_calls_, 0); + + auto objects_to_cancel = pull_manager_.CancelPull(req_id); + ASSERT_EQ(objects_to_cancel, ObjectRefsToIds(refs)); ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); } TEST_F(PullManagerTest, TestManyUpdates) { - ObjectID obj1 = ObjectID::FromRandom(); + auto refs = CreateObjectRefs(1); + auto obj1 = ObjectRefsToIds(refs)[0]; rpc::Address addr1; ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); - pull_manager_.Pull(obj1, addr1); + std::vector objects_to_locate; + auto req_id = pull_manager_.Pull(refs, &objects_to_locate); + ASSERT_EQ(ObjectRefsToIds(objects_to_locate), ObjectRefsToIds(refs)); ASSERT_EQ(pull_manager_.NumActiveRequests(), 1); std::unordered_set client_ids; @@ -105,15 +136,19 @@ TEST_F(PullManagerTest, TestManyUpdates) { ASSERT_EQ(num_send_pull_request_calls_, 100); ASSERT_EQ(num_restore_spilled_object_calls_, 0); - pull_manager_.CancelPull(obj1); + auto objects_to_cancel = pull_manager_.CancelPull(req_id); + ASSERT_EQ(objects_to_cancel, ObjectRefsToIds(refs)); ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); } TEST_F(PullManagerTest, TestRetryTimer) { - ObjectID obj1 = ObjectID::FromRandom(); + auto refs = CreateObjectRefs(1); + auto obj1 = ObjectRefsToIds(refs)[0]; rpc::Address addr1; ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); - pull_manager_.Pull(obj1, addr1); + std::vector objects_to_locate; + auto req_id = pull_manager_.Pull(refs, &objects_to_locate); + ASSERT_EQ(ObjectRefsToIds(objects_to_locate), ObjectRefsToIds(refs)); ASSERT_EQ(pull_manager_.NumActiveRequests(), 1); std::unordered_set client_ids; @@ -143,26 +178,102 @@ TEST_F(PullManagerTest, TestRetryTimer) { ASSERT_EQ(num_send_pull_request_calls_, 1 + 7 + 127); ASSERT_EQ(num_restore_spilled_object_calls_, 0); - pull_manager_.CancelPull(obj1); + // Don't retry an object if it's local. + object_is_local_ = true; + num_send_pull_request_calls_ = 0; + for (; fake_time_ <= 127 * 10; fake_time_ += 1.) { + pull_manager_.Tick(); + } + ASSERT_EQ(num_send_pull_request_calls_, 0); + + auto objects_to_cancel = pull_manager_.CancelPull(req_id); + ASSERT_EQ(objects_to_cancel, ObjectRefsToIds(refs)); ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); } TEST_F(PullManagerTest, TestBasic) { - ObjectID obj1 = ObjectID::FromRandom(); - rpc::Address addr1; + auto refs = CreateObjectRefs(3); + auto oids = ObjectRefsToIds(refs); ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); - pull_manager_.Pull(obj1, addr1); - ASSERT_EQ(pull_manager_.NumActiveRequests(), 1); + std::vector objects_to_locate; + auto req_id = pull_manager_.Pull(refs, &objects_to_locate); + ASSERT_EQ(ObjectRefsToIds(objects_to_locate), oids); + ASSERT_EQ(pull_manager_.NumActiveRequests(), oids.size()); std::unordered_set client_ids; client_ids.insert(NodeID::FromRandom()); - pull_manager_.OnLocationChange(obj1, client_ids, ""); + for (size_t i = 0; i < oids.size(); i++) { + pull_manager_.OnLocationChange(oids[i], client_ids, ""); + ASSERT_EQ(num_send_pull_request_calls_, i + 1); + ASSERT_EQ(num_restore_spilled_object_calls_, 0); + } - ASSERT_EQ(num_send_pull_request_calls_, 1); - ASSERT_EQ(num_restore_spilled_object_calls_, 0); + // Don't pull an object if it's local. + object_is_local_ = true; + num_send_pull_request_calls_ = 0; + for (size_t i = 0; i < oids.size(); i++) { + pull_manager_.OnLocationChange(oids[i], client_ids, ""); + } + ASSERT_EQ(num_send_pull_request_calls_, 0); - pull_manager_.CancelPull(obj1); + auto objects_to_cancel = pull_manager_.CancelPull(req_id); + ASSERT_EQ(objects_to_cancel, oids); ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); + + // Don't pull a remote object if we've canceled. + object_is_local_ = false; + num_send_pull_request_calls_ = 0; + for (size_t i = 0; i < oids.size(); i++) { + pull_manager_.OnLocationChange(oids[i], client_ids, ""); + } + ASSERT_EQ(num_send_pull_request_calls_, 0); +} + +TEST_F(PullManagerTest, TestDeduplicateBundles) { + auto refs = CreateObjectRefs(3); + auto oids = ObjectRefsToIds(refs); + ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); + std::vector objects_to_locate; + auto req_id1 = pull_manager_.Pull(refs, &objects_to_locate); + ASSERT_EQ(ObjectRefsToIds(objects_to_locate), oids); + ASSERT_EQ(pull_manager_.NumActiveRequests(), oids.size()); + + objects_to_locate.clear(); + auto req_id2 = pull_manager_.Pull(refs, &objects_to_locate); + ASSERT_TRUE(objects_to_locate.empty()); + + std::unordered_set client_ids; + client_ids.insert(NodeID::FromRandom()); + for (size_t i = 0; i < oids.size(); i++) { + pull_manager_.OnLocationChange(oids[i], client_ids, ""); + ASSERT_EQ(num_send_pull_request_calls_, i + 1); + ASSERT_EQ(num_restore_spilled_object_calls_, 0); + } + + // Cancel one request. + auto objects_to_cancel = pull_manager_.CancelPull(req_id1); + ASSERT_TRUE(objects_to_cancel.empty()); + // Objects should still be pulled because the other request is still open. + ASSERT_EQ(pull_manager_.NumActiveRequests(), oids.size()); + num_send_pull_request_calls_ = 0; + for (size_t i = 0; i < oids.size(); i++) { + pull_manager_.OnLocationChange(oids[i], client_ids, ""); + ASSERT_EQ(num_send_pull_request_calls_, i + 1); + ASSERT_EQ(num_restore_spilled_object_calls_, 0); + } + + // Cancel the other request. + objects_to_cancel = pull_manager_.CancelPull(req_id2); + ASSERT_EQ(objects_to_cancel, oids); + ASSERT_EQ(pull_manager_.NumActiveRequests(), 0); + + // Don't pull a remote object if we've canceled. + object_is_local_ = false; + num_send_pull_request_calls_ = 0; + for (size_t i = 0; i < oids.size(); i++) { + pull_manager_.OnLocationChange(oids[i], client_ids, ""); + } + ASSERT_EQ(num_send_pull_request_calls_, 0); } } // namespace ray diff --git a/src/ray/raylet/dependency_manager.cc b/src/ray/raylet/dependency_manager.cc new file mode 100644 index 000000000..988893bea --- /dev/null +++ b/src/ray/raylet/dependency_manager.cc @@ -0,0 +1,311 @@ +#include "ray/raylet/dependency_manager.h" + +namespace ray { + +namespace raylet { + +bool DependencyManager::CheckObjectLocal(const ObjectID &object_id) const { + return local_objects_.count(object_id) == 1; +} + +bool DependencyManager::GetOwnerAddress(const ObjectID &object_id, + rpc::Address *owner_address) const { + auto obj = required_objects_.find(object_id); + if (obj == required_objects_.end()) { + return false; + } + + *owner_address = obj->second.owner_address; + return !owner_address->worker_id().empty(); +} + +void DependencyManager::RemoveObjectIfNotNeeded( + absl::flat_hash_map::iterator + required_object_it) { + const auto &object_id = required_object_it->first; + if (required_object_it->second.Empty()) { + RAY_LOG(DEBUG) << "Object " << object_id << " no longer needed"; + if (required_object_it->second.wait_request_id > 0) { + RAY_LOG(DEBUG) << "Canceling pull for wait request of object " << object_id + << " request: " << required_object_it->second.wait_request_id; + object_manager_.CancelPull(required_object_it->second.wait_request_id); + } + if (!local_objects_.count(object_id)) { + reconstruction_policy_.Cancel(object_id); + } + required_objects_.erase(required_object_it); + } +} + +absl::flat_hash_map::iterator +DependencyManager::GetOrInsertRequiredObject(const ObjectID &object_id, + const rpc::ObjectReference &ref) { + auto it = required_objects_.find(object_id); + if (it == required_objects_.end()) { + it = required_objects_.emplace(object_id, ref).first; + if (local_objects_.count(object_id) == 0) { + reconstruction_policy_.ListenAndMaybeReconstruct(object_id, ref.owner_address()); + } + } + return it; +} + +void DependencyManager::StartOrUpdateWaitRequest( + const WorkerID &worker_id, + const std::vector &required_objects) { + RAY_LOG(DEBUG) << "Starting wait request for worker " << worker_id; + auto &wait_request = wait_requests_[worker_id]; + for (const auto &ref : required_objects) { + const auto obj_id = ObjectRefToId(ref); + if (local_objects_.count(obj_id)) { + // Object is already local. No need to fetch it. + continue; + } + + if (wait_request.insert(obj_id).second) { + RAY_LOG(DEBUG) << "Worker " << worker_id << " called ray.wait on non-local object " + << obj_id; + auto it = GetOrInsertRequiredObject(obj_id, ref); + it->second.dependent_wait_requests.insert(worker_id); + if (it->second.wait_request_id == 0) { + it->second.wait_request_id = object_manager_.Pull({ref}); + RAY_LOG(DEBUG) << "Started pull for wait request for object " << obj_id + << " request: " << it->second.wait_request_id; + } + } + } + + // No new objects to wait on. Delete the empty entry that was created. + if (wait_request.empty()) { + wait_requests_.erase(worker_id); + } +} + +void DependencyManager::CancelWaitRequest(const WorkerID &worker_id) { + RAY_LOG(DEBUG) << "Canceling wait request for worker " << worker_id; + auto it = wait_requests_.find(worker_id); + if (it == wait_requests_.end()) { + return; + } + + for (const auto &obj_id : it->second) { + auto it = required_objects_.find(obj_id); + RAY_CHECK(it != required_objects_.end()); + it->second.dependent_wait_requests.erase(worker_id); + RemoveObjectIfNotNeeded(it); + } + + wait_requests_.erase(it); +} + +void DependencyManager::StartOrUpdateGetRequest( + const WorkerID &worker_id, + const std::vector &required_objects) { + RAY_LOG(DEBUG) << "Starting get request for worker " << worker_id; + auto &get_request = get_requests_[worker_id]; + bool modified = false; + for (const auto &ref : required_objects) { + const auto obj_id = ObjectRefToId(ref); + if (get_request.first.insert(obj_id).second) { + RAY_LOG(DEBUG) << "Worker " << worker_id << " called ray.get on object " << obj_id; + auto it = GetOrInsertRequiredObject(obj_id, ref); + it->second.dependent_get_requests.insert(worker_id); + modified = true; + } + } + + if (modified) { + std::vector refs; + for (auto &obj_id : get_request.first) { + auto it = required_objects_.find(obj_id); + RAY_CHECK(it != required_objects_.end()); + refs.push_back(ObjectIdToRef(obj_id, it->second.owner_address)); + } + // Pull the new dependencies before canceling the old request, in case some + // of the old dependencies are still being fetched. + uint64_t new_request_id = object_manager_.Pull(refs); + if (get_request.second != 0) { + RAY_LOG(DEBUG) << "Canceling pull for get request from worker " << worker_id + << " request: " << get_request.second; + object_manager_.CancelPull(get_request.second); + } + get_request.second = new_request_id; + RAY_LOG(DEBUG) << "Started pull for get request from worker " << worker_id + << " request: " << get_request.second; + } +} + +void DependencyManager::CancelGetRequest(const WorkerID &worker_id) { + RAY_LOG(DEBUG) << "Canceling get request for worker " << worker_id; + auto it = get_requests_.find(worker_id); + if (it == get_requests_.end()) { + return; + } + + RAY_LOG(DEBUG) << "Canceling pull for get request from worker " << worker_id + << " request: " << it->second.second; + object_manager_.CancelPull(it->second.second); + + for (const auto &obj_id : it->second.first) { + auto it = required_objects_.find(obj_id); + RAY_CHECK(it != required_objects_.end()); + it->second.dependent_get_requests.erase(worker_id); + RemoveObjectIfNotNeeded(it); + } + + get_requests_.erase(it); +} + +/// Request dependencies for a queued task. +bool DependencyManager::RequestTaskDependencies( + const TaskID &task_id, const std::vector &required_objects) { + RAY_LOG(DEBUG) << "Adding dependencies for task " << task_id; + auto inserted = queued_task_requests_.emplace(task_id, required_objects); + RAY_CHECK(inserted.second) << "Task depedencies can be requested only once per task."; + auto &task_entry = inserted.first->second; + + for (const auto &ref : required_objects) { + const auto obj_id = ObjectRefToId(ref); + RAY_LOG(DEBUG) << "Task " << task_id << " blocked on object " << obj_id; + + auto it = GetOrInsertRequiredObject(obj_id, ref); + it->second.dependent_tasks.insert(task_id); + + if (local_objects_.count(obj_id)) { + task_entry.num_missing_dependencies--; + } + } + + if (!required_objects.empty()) { + task_entry.pull_request_id = object_manager_.Pull(required_objects); + RAY_LOG(DEBUG) << "Started pull for dependencies of task " << task_id + << " request: " << task_entry.pull_request_id; + } + + return task_entry.num_missing_dependencies == 0; +} + +bool DependencyManager::IsTaskReady(const TaskID &task_id) const { + auto task_entry = queued_task_requests_.find(task_id); + RAY_CHECK(task_entry != queued_task_requests_.end()); + return task_entry->second.num_missing_dependencies == 0; +} + +void DependencyManager::RemoveTaskDependencies(const TaskID &task_id) { + RAY_LOG(DEBUG) << "Removing dependencies for task " << task_id; + auto task_entry = queued_task_requests_.find(task_id); + RAY_CHECK(task_entry != queued_task_requests_.end()) + << "Can't remove dependencies of tasks that are not queued."; + + if (task_entry->second.pull_request_id > 0) { + RAY_LOG(DEBUG) << "Canceling pull for dependencies of task " << task_id + << " request: " << task_entry->second.pull_request_id; + object_manager_.CancelPull(task_entry->second.pull_request_id); + } + + for (const auto &obj_id : task_entry->second.dependencies) { + auto it = required_objects_.find(obj_id); + RAY_CHECK(it != required_objects_.end()); + it->second.dependent_tasks.erase(task_id); + RemoveObjectIfNotNeeded(it); + } + + queued_task_requests_.erase(task_entry); +} + +std::vector DependencyManager::HandleObjectMissing( + const ray::ObjectID &object_id) { + RAY_CHECK(local_objects_.erase(object_id)) + << "Evicted object was not local " << object_id; + + // Find any tasks that are dependent on the missing object. + std::vector waiting_task_ids; + auto object_entry = required_objects_.find(object_id); + if (object_entry != required_objects_.end()) { + for (auto &dependent_task_id : object_entry->second.dependent_tasks) { + auto it = queued_task_requests_.find(dependent_task_id); + RAY_CHECK(it != queued_task_requests_.end()); + auto &task_entry = it->second; + // If the dependent task had all of its arguments ready, it was ready to + // run but must be switched to waiting since one of its arguments is now + // missing. + if (task_entry.num_missing_dependencies == 0) { + waiting_task_ids.push_back(dependent_task_id); + // During normal execution we should be able to include the check + // RAY_CHECK(pending_tasks_.count(dependent_task_id) == 1); + // However, this invariant will not hold during unit test execution. + } + task_entry.num_missing_dependencies++; + } + + // The object is missing and needed so wait for a possible failure again. + reconstruction_policy_.ListenAndMaybeReconstruct(object_entry->first, + object_entry->second.owner_address); + } + + // Process callbacks for all of the tasks dependent on the object that are + // now ready to run. + return waiting_task_ids; +} + +std::vector DependencyManager::HandleObjectLocal(const ray::ObjectID &object_id) { + // Add the object to the table of locally available objects. + auto inserted = local_objects_.insert(object_id); + RAY_CHECK(inserted.second) << "Local object was already local " << object_id; + + // Find all tasks and workers that depend on the newly available object. + std::vector ready_task_ids; + auto object_entry = required_objects_.find(object_id); + if (object_entry != required_objects_.end()) { + // Loop through all tasks that depend on the newly available object. + for (const auto &dependent_task_id : object_entry->second.dependent_tasks) { + auto it = queued_task_requests_.find(dependent_task_id); + RAY_CHECK(it != queued_task_requests_.end()); + auto &task_entry = it->second; + task_entry.num_missing_dependencies--; + // If the dependent task now has all of its arguments ready, it's ready + // to run. + if (task_entry.num_missing_dependencies == 0) { + ready_task_ids.push_back(dependent_task_id); + } + } + + // Remove the dependency from all workers that called `ray.wait` on the + // newly available object. + for (const auto &worker_id : object_entry->second.dependent_wait_requests) { + auto worker_it = wait_requests_.find(worker_id); + RAY_CHECK(worker_it != wait_requests_.end()); + RAY_CHECK(worker_it->second.erase(object_id) > 0); + if (worker_it->second.empty()) { + wait_requests_.erase(worker_it); + } + } + // Clear all workers that called `ray.wait` on this object, since the + // `ray.wait` calls can now return the object as ready. + object_entry->second.dependent_wait_requests.clear(); + if (object_entry->second.wait_request_id > 0) { + RAY_LOG(DEBUG) << "Canceling pull for wait request of object " << object_id + << " request: " << object_entry->second.wait_request_id; + object_manager_.CancelPull(object_entry->second.wait_request_id); + object_entry->second.wait_request_id = 0; + } + reconstruction_policy_.Cancel(object_entry->first); + RemoveObjectIfNotNeeded(object_entry); + } + + return ready_task_ids; +} + +std::string DependencyManager::DebugString() const { + std::stringstream result; + result << "TaskDependencyManager:"; + result << "\n- task deps map size: " << queued_task_requests_.size(); + result << "\n- get req map size: " << get_requests_.size(); + result << "\n- wait req map size: " << wait_requests_.size(); + result << "\n- local objects map size: " << local_objects_.size(); + return result.str(); +} + +} // namespace raylet + +} // namespace ray diff --git a/src/ray/raylet/dependency_manager.h b/src/ray/raylet/dependency_manager.h new file mode 100644 index 000000000..1e7ddfcb1 --- /dev/null +++ b/src/ray/raylet/dependency_manager.h @@ -0,0 +1,270 @@ +// Copyright 2017 The Ray Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +// clang-format off +#include "ray/common/common_protocol.h" +#include "ray/common/id.h" +#include "ray/common/task/task.h" +#include "ray/object_manager/object_manager.h" +#include "ray/raylet/reconstruction_policy.h" +// clang-format on + +namespace ray { + +namespace raylet { + +using rpc::TaskLeaseData; + +class ReconstructionPolicy; + +/// Used for unit-testing the ClusterTaskManager, which requests dependencies +/// for queued tasks. +class TaskDependencyManagerInterface { + public: + virtual bool RequestTaskDependencies( + const TaskID &task_id, + const std::vector &required_objects) = 0; + virtual bool IsTaskReady(const TaskID &task_id) const = 0; + virtual void RemoveTaskDependencies(const TaskID &task_id) = 0; + virtual ~TaskDependencyManagerInterface(){}; +}; + +/// \class DependencyManager +/// +/// Responsible for managing object dependencies for local workers calling +/// `ray.get` or `ray.wait` and arguments of queued tasks. The caller can +/// request object dependencies for a task or worker. The task manager will +/// determine which object dependencies are remote and will request that these +/// objects be made available locally, either via the object manager or by +/// storing an error if the object is lost. +class DependencyManager : public TaskDependencyManagerInterface { + public: + /// Create a task dependency manager. + DependencyManager(ObjectManagerInterface &object_manager, + ReconstructionPolicyInterface &reconstruction_policy) + : object_manager_(object_manager), reconstruction_policy_(reconstruction_policy) {} + + /// Check whether an object is locally available. + /// + /// \param object_id The object to check for. + /// \return Whether the object is local. + bool CheckObjectLocal(const ObjectID &object_id) const; + + /// Get the address of the owner of this object. An address will only be + /// returned if the caller previously specified that this object is required + /// on this node, through a call to SubscribeGetDependencies or + /// SubscribeWaitDependencies. + /// + /// \param[in] object_id The object whose owner to get. + /// \param[out] owner_address The address of the object's owner, if + /// available. + /// \return True if we have owner information for the object. + bool GetOwnerAddress(const ObjectID &object_id, rpc::Address *owner_address) const; + + /// Start or update a worker's `ray.wait` request. This will attempt to make + /// any remote objects local, including previously requested objects. The + /// `ray.wait` request will stay active until the objects are made local or + /// the request for this worker is canceled, whichever occurs first. + /// + /// This method may be called multiple times per worker on the same objects. + /// + /// \param worker_id The ID of the worker that called `ray.wait`. + /// \param required_objects The objects required by the worker. + /// \return Void. + void StartOrUpdateWaitRequest( + const WorkerID &worker_id, + const std::vector &required_objects); + + /// Cancel a worker's `ray.wait` request. We will no longer attempt to fetch + /// any objects that this worker requested previously, if no other task or + /// worker requires them. + /// + /// \param worker_id The ID of the worker whose `ray.wait` request we should + /// cancel. + /// \return Void. + void CancelWaitRequest(const WorkerID &worker_id); + + /// Start or update a worker's `ray.get` request. This will attempt to make + /// any remote objects local, including previously requested objects. The + /// `ray.get` request will stay active until the request for this worker is + /// canceled. + /// + /// This method may be called multiple times per worker on the same objects. + /// + /// \param worker_id The ID of the worker that called `ray.wait`. + /// \param required_objects The objects required by the worker. + /// \return Void. + void StartOrUpdateGetRequest(const WorkerID &worker_id, + const std::vector &required_objects); + + /// Cancel a worker's `ray.get` request. We will no longer attempt to fetch + /// any objects that this worker requested previously, if no other task or + /// worker requires them. + /// + /// \param worker_id The ID of the worker whose `ray.get` request we should + /// cancel. + /// \return Void. + void CancelGetRequest(const WorkerID &worker_id); + + /// Request dependencies for a queued task. This will attempt to make any + /// remote objects local until the caller cancels the task's dependencies. + /// + /// This method can only be called once per task, until the task has been + /// canceled. + /// + /// \param task_id The task that requires the objects. + /// \param required_objects The objects required by the task. + /// \return Void. + bool RequestTaskDependencies(const TaskID &task_id, + const std::vector &required_objects); + + /// Check whether a task is ready to run. The task ID must have been + /// previously added by the caller. + /// + /// \param task_id The ID of the task to check. + /// \return Whether all of the dependencies for the task are + /// local. + bool IsTaskReady(const TaskID &task_id) const; + + /// Cancel a task's dependencies. We will no longer attempt to fetch any + /// remote dependencies, if no other task or worker requires them. + /// + /// This method can only be called on a task whose dependencies were added. + /// + /// \param task_id The task that requires the objects. + /// \param required_objects The objects required by the task. + /// \return Void. + void RemoveTaskDependencies(const TaskID &task_id); + + /// Handle an object becoming locally available. + /// + /// \param object_id The object ID of the object to mark as locally + /// available. + /// \return A list of task IDs. This contains all added tasks that now have + /// all of their dependencies fulfilled. + std::vector HandleObjectLocal(const ray::ObjectID &object_id); + + /// Handle an object that is no longer locally available. + /// + /// \param object_id The object ID of the object that was previously locally + /// available. + /// \return A list of task IDs. This contains all added tasks that previously + /// had all of their dependencies fulfilled, but are now missing this object + /// dependency. + std::vector HandleObjectMissing(const ray::ObjectID &object_id); + + /// Returns debug string for class. + /// + /// \return string. + std::string DebugString() const; + + private: + /// Metadata for an object that is needed by at least one executing worker + /// and/or one queued task. + struct ObjectDependencies { + ObjectDependencies(const rpc::ObjectReference &ref) + : owner_address(ref.owner_address()) {} + /// The tasks that depend on this object, either because the object is a task argument + /// or because the task called `ray.get` on the object. + std::unordered_set dependent_tasks; + /// The workers that depend on this object because they called `ray.get` on the + /// object. + std::unordered_set dependent_get_requests; + /// The workers that depend on this object because they called `ray.wait` on the + /// object. + std::unordered_set dependent_wait_requests; + /// If this object is required by at least one worker that called `ray.wait`, this is + /// the pull request ID. + uint64_t wait_request_id = 0; + /// The address of the worker that owns this object. + rpc::Address owner_address; + + bool Empty() const { + return dependent_tasks.empty() && dependent_get_requests.empty() && + dependent_wait_requests.empty(); + } + }; + + /// A struct to represent the object dependencies of a task. + struct TaskDependencies { + TaskDependencies(const std::vector &deps) + : num_missing_dependencies(deps.size()) { + const auto dep_ids = ObjectRefsToIds(deps); + dependencies.insert(dep_ids.begin(), dep_ids.end()); + } + /// The objects that the task depends on. These are the arguments to the + /// task. These must all be simultaneously local before the task is ready + /// to execute. Objects are removed from this set once + /// UnsubscribeGetDependencies is called. + absl::flat_hash_set dependencies; + /// The number of object arguments that are not available locally. This + /// must be zero before the task is ready to execute. + size_t num_missing_dependencies; + /// Used to identify the pull request for the dependencies to the object + /// manager. + uint64_t pull_request_id = 0; + }; + + /// Stop tracking this object, if it is no longer needed by any worker or + /// queued task. + void RemoveObjectIfNotNeeded( + absl::flat_hash_map::iterator required_object_it); + + /// Start tracking an object that is needed by a worker and/or queued task. + absl::flat_hash_map::iterator GetOrInsertRequiredObject( + const ObjectID &object_id, const rpc::ObjectReference &ref); + + /// The object manager, used to fetch required objects from remote nodes. + ObjectManagerInterface &object_manager_; + /// The reconstruction policy, used to reconstruct required objects that no + /// longer exist on any live nodes. + /// TODO(swang): This class is no longer needed for reconstruction, since the + /// object's owner handles reconstruction. We use this class as a timer to + /// detect the owner's failure. Remove this class and move the timer logic + /// into this class. + ReconstructionPolicyInterface &reconstruction_policy_; + + /// A map from the ID of a queued task to metadata about whether the task's + /// dependencies are all local or not. + absl::flat_hash_map queued_task_requests_; + + /// A map from worker ID to the set of objects that the worker called + /// `ray.get` on and a pull request ID for these objects. The pull request ID + /// should be used to cancel the pull request in the object manager once the + /// worker cancels the `ray.get` request. + absl::flat_hash_map, uint64_t>> + get_requests_; + + /// A map from worker ID to the set of objects that the worker called + /// `ray.wait` on. Objects are removed from the set once they are made local, + /// or the worker cancels the `ray.wait` request. + absl::flat_hash_map> wait_requests_; + + /// Deduplicated pool of objects required by all queued tasks and workers. + /// Objects are removed from this set once there are no more tasks or workers + /// that require it. + absl::flat_hash_map required_objects_; + + /// The set of locally available objects. This is used to determine which + /// tasks are ready to run and which `ray.wait` requests can be finished. + std::unordered_set local_objects_; + + friend class DependencyManagerTest; +}; + +} // namespace raylet + +} // namespace ray diff --git a/src/ray/raylet/dependency_manager_test.cc b/src/ray/raylet/dependency_manager_test.cc new file mode 100644 index 000000000..c6d0ab2ee --- /dev/null +++ b/src/ray/raylet/dependency_manager_test.cc @@ -0,0 +1,372 @@ +// Copyright 2017 The Ray Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "ray/raylet/dependency_manager.h" + +#include + +#include "gmock/gmock.h" +#include "gtest/gtest.h" +#include "ray/common/task/task_util.h" +#include "ray/common/test_util.h" + +namespace ray { + +namespace raylet { + +using ::testing::_; +using ::testing::InSequence; +using ::testing::Return; + +class MockObjectManager : public ObjectManagerInterface { + public: + uint64_t Pull(const std::vector &object_refs) { + active_requests.insert(req_id); + return req_id++; + } + + void CancelPull(uint64_t request_id) { ASSERT_TRUE(active_requests.erase(request_id)); } + + uint64_t req_id = 1; + std::unordered_set active_requests; +}; + +class MockReconstructionPolicy : public ReconstructionPolicyInterface { + public: + MOCK_METHOD2(ListenAndMaybeReconstruct, + void(const ObjectID &object_id, const rpc::Address &owner_address)); + MOCK_METHOD1(Cancel, void(const ObjectID &object_id)); +}; + +class DependencyManagerTest : public ::testing::Test { + public: + DependencyManagerTest() + : object_manager_mock_(), + reconstruction_policy_mock_(), + dependency_manager_(object_manager_mock_, reconstruction_policy_mock_) {} + + void AssertNoLeaks() { + ASSERT_TRUE(dependency_manager_.required_objects_.empty()); + ASSERT_TRUE(dependency_manager_.queued_task_requests_.empty()); + ASSERT_TRUE(dependency_manager_.get_requests_.empty()); + ASSERT_TRUE(dependency_manager_.wait_requests_.empty()); + // All pull requests are canceled. + ASSERT_TRUE(object_manager_mock_.active_requests.empty()); + } + + MockObjectManager object_manager_mock_; + MockReconstructionPolicy reconstruction_policy_mock_; + DependencyManager dependency_manager_; +}; + +/// Test requesting the dependencies for a task. The dependency manager should +/// return the task ID as ready once all of its arguments are local. +TEST_F(DependencyManagerTest, TestSimpleTask) { + // Create a task with 3 arguments. + int num_arguments = 3; + std::vector arguments; + for (int i = 0; i < num_arguments; i++) { + arguments.push_back(ObjectID::FromRandom()); + } + TaskID task_id = RandomTaskId(); + // No objects have been registered in the task dependency manager, so all + // arguments should be remote. + for (const auto &argument_id : arguments) { + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); + } + bool ready = + dependency_manager_.RequestTaskDependencies(task_id, ObjectIdsToRefs(arguments)); + ASSERT_FALSE(ready); + ASSERT_EQ(object_manager_mock_.active_requests.size(), 1); + ASSERT_FALSE(dependency_manager_.IsTaskReady(task_id)); + + // For each argument, tell the task dependency manager that the argument is + // local. All arguments should be canceled as they become available locally. + for (const auto &argument_id : arguments) { + EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); + } + auto ready_task_ids = dependency_manager_.HandleObjectLocal(arguments[0]); + ASSERT_TRUE(ready_task_ids.empty()); + ASSERT_FALSE(dependency_manager_.IsTaskReady(task_id)); + ready_task_ids = dependency_manager_.HandleObjectLocal(arguments[1]); + ASSERT_TRUE(ready_task_ids.empty()); + ASSERT_FALSE(dependency_manager_.IsTaskReady(task_id)); + // The task is ready to run. + ready_task_ids = dependency_manager_.HandleObjectLocal(arguments[2]); + ASSERT_EQ(ready_task_ids.size(), 1); + ASSERT_EQ(ready_task_ids.front(), task_id); + ASSERT_TRUE(dependency_manager_.IsTaskReady(task_id)); + + // Remove the task. + dependency_manager_.RemoveTaskDependencies(task_id); + AssertNoLeaks(); +} + +/// Test multiple tasks that depend on the same object. The dependency manager +/// should return all task IDs as ready once the object is local. +TEST_F(DependencyManagerTest, TestMultipleTasks) { + // Create 3 tasks that are dependent on the same object. + ObjectID argument_id = ObjectID::FromRandom(); + std::vector dependent_tasks; + int num_dependent_tasks = 3; + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); + for (int i = 0; i < num_dependent_tasks; i++) { + TaskID task_id = RandomTaskId(); + dependent_tasks.push_back(task_id); + bool ready = dependency_manager_.RequestTaskDependencies( + task_id, ObjectIdsToRefs({argument_id})); + ASSERT_FALSE(ready); + ASSERT_FALSE(dependency_manager_.IsTaskReady(task_id)); + // The object should be requested from the object manager once for each task. + ASSERT_EQ(object_manager_mock_.active_requests.size(), i + 1); + } + + // Tell the task dependency manager that the object is local. + EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); + auto ready_task_ids = dependency_manager_.HandleObjectLocal(argument_id); + // Check that all tasks are now ready to run. + std::unordered_set added_tasks(dependent_tasks.begin(), dependent_tasks.end()); + for (auto &id : ready_task_ids) { + ASSERT_TRUE(added_tasks.erase(id)); + ASSERT_TRUE(dependency_manager_.IsTaskReady(id)); + } + ASSERT_TRUE(added_tasks.empty()); + + for (auto &id : dependent_tasks) { + dependency_manager_.RemoveTaskDependencies(id); + } + AssertNoLeaks(); +} + +/// Test task with multiple dependencies. The dependency manager should return +/// the task ID as ready once all dependencies are local. If a dependency is +/// later evicted, the dependency manager should return the task ID as waiting. +TEST_F(DependencyManagerTest, TestTaskArgEviction) { + // Add a task with 3 arguments. + int num_arguments = 3; + std::vector arguments; + for (int i = 0; i < num_arguments; i++) { + arguments.push_back(ObjectID::FromRandom()); + } + TaskID task_id = RandomTaskId(); + for (const auto &argument_id : arguments) { + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); + } + bool ready = + dependency_manager_.RequestTaskDependencies(task_id, ObjectIdsToRefs(arguments)); + ASSERT_FALSE(ready); + ASSERT_FALSE(dependency_manager_.IsTaskReady(task_id)); + + // Tell the task dependency manager that each of the arguments is now + // available. + for (const auto &argument_id : arguments) { + EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); + } + for (size_t i = 0; i < arguments.size(); i++) { + std::vector ready_tasks; + ready_tasks = dependency_manager_.HandleObjectLocal(arguments[i]); + if (i == arguments.size() - 1) { + ASSERT_EQ(ready_tasks.size(), 1); + ASSERT_EQ(ready_tasks.front(), task_id); + } else { + ASSERT_TRUE(ready_tasks.empty()); + } + } + ASSERT_TRUE(dependency_manager_.IsTaskReady(task_id)); + + // Simulate each of the arguments getting evicted. Each object should now be + // considered remote. + for (const auto &argument_id : arguments) { + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); + } + for (size_t i = 0; i < arguments.size(); i++) { + std::vector waiting_tasks; + waiting_tasks = dependency_manager_.HandleObjectMissing(arguments[i]); + if (i == 0) { + // The first eviction should cause the task to go back to the waiting + // state. + ASSERT_EQ(waiting_tasks.size(), 1); + ASSERT_EQ(waiting_tasks.front(), task_id); + } else { + // The subsequent evictions shouldn't cause any more tasks to go back to + // the waiting state. + ASSERT_TRUE(waiting_tasks.empty()); + } + ASSERT_FALSE(dependency_manager_.IsTaskReady(task_id)); + } + + // Tell the task dependency manager that each of the arguments is available + // again. + for (const auto &argument_id : arguments) { + EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); + } + for (size_t i = 0; i < arguments.size(); i++) { + std::vector ready_tasks; + ready_tasks = dependency_manager_.HandleObjectLocal(arguments[i]); + if (i == arguments.size() - 1) { + ASSERT_EQ(ready_tasks.size(), 1); + ASSERT_EQ(ready_tasks.front(), task_id); + } else { + ASSERT_TRUE(ready_tasks.empty()); + } + } + ASSERT_TRUE(dependency_manager_.IsTaskReady(task_id)); + + dependency_manager_.RemoveTaskDependencies(task_id); + AssertNoLeaks(); +} + +/// Test `ray.get`. Worker calls ray.get on {oid1}, then {oid1, oid2}, then +/// {oid1, oid2, oid3}. +TEST_F(DependencyManagerTest, TestGet) { + WorkerID worker_id = WorkerID::FromRandom(); + int num_arguments = 3; + std::vector arguments; + for (int i = 0; i < num_arguments; i++) { + // Add the new argument to the list of dependencies to subscribe to. + ObjectID argument_id = ObjectID::FromRandom(); + arguments.push_back(argument_id); + // Subscribe to the task's dependencies. All arguments except the last are + // duplicates of previous subscription calls. Each argument should only be + // requested from the node manager once. + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); + auto prev_pull_reqs = object_manager_mock_.active_requests; + dependency_manager_.StartOrUpdateGetRequest(worker_id, ObjectIdsToRefs(arguments)); + // Previous pull request for this worker should be canceled upon each new + // bundle. + ASSERT_EQ(object_manager_mock_.active_requests.size(), 1); + ASSERT_NE(object_manager_mock_.active_requests, prev_pull_reqs); + } + + // Nothing happens if the same bundle is requested. + auto prev_pull_reqs = object_manager_mock_.active_requests; + dependency_manager_.StartOrUpdateGetRequest(worker_id, ObjectIdsToRefs(arguments)); + ASSERT_EQ(object_manager_mock_.active_requests, prev_pull_reqs); + + // All arguments should be canceled as they become available locally. + for (const auto &argument_id : arguments) { + EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); + } + + // Cancel the pull request once the worker cancels the `ray.get`. + dependency_manager_.CancelGetRequest(worker_id); + AssertNoLeaks(); +} + +/// Test that when one of the objects becomes local after a `ray.wait` call, +/// all requests to remote nodes associated with the object are canceled. +TEST_F(DependencyManagerTest, TestWait) { + // Generate a random worker and objects to wait on. + WorkerID worker_id = WorkerID::FromRandom(); + int num_objects = 3; + std::vector oids; + for (int i = 0; i < num_objects; i++) { + oids.push_back(ObjectID::FromRandom()); + } + // Simulate a worker calling `ray.wait` on some objects. + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(_, _)) + .Times(num_objects); + dependency_manager_.StartOrUpdateWaitRequest(worker_id, ObjectIdsToRefs(oids)); + ASSERT_EQ(object_manager_mock_.active_requests.size(), num_objects); + + for (int i = 0; i < num_objects; i++) { + // Object is local. + EXPECT_CALL(reconstruction_policy_mock_, Cancel(oids[i])); + auto ready_task_ids = dependency_manager_.HandleObjectLocal(oids[i]); + + // Local object gets evicted. The `ray.wait` call should not be + // reactivated. + auto waiting_task_ids = dependency_manager_.HandleObjectMissing(oids[i]); + ASSERT_TRUE(waiting_task_ids.empty()); + ASSERT_EQ(object_manager_mock_.active_requests.size(), num_objects - i - 1); + } + AssertNoLeaks(); +} + +/// Test that when no objects are locally available, a `ray.wait` call makes +/// the correct requests to remote nodes and correctly cancels the requests +/// when the `ray.wait` call is canceled. +TEST_F(DependencyManagerTest, TestWaitThenCancel) { + // Generate a random worker and objects to wait on. + WorkerID worker_id = WorkerID::FromRandom(); + int num_objects = 3; + std::vector oids; + for (int i = 0; i < num_objects; i++) { + oids.push_back(ObjectID::FromRandom()); + } + // Simulate a worker calling `ray.wait` on some objects. + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(_, _)) + .Times(num_objects); + dependency_manager_.StartOrUpdateWaitRequest(worker_id, ObjectIdsToRefs(oids)); + ASSERT_EQ(object_manager_mock_.active_requests.size(), num_objects); + auto prev_pull_reqs = object_manager_mock_.active_requests; + // Check that it's okay to call `ray.wait` on the same objects again. No new + // calls should be made to try and make the objects local. + dependency_manager_.StartOrUpdateWaitRequest(worker_id, ObjectIdsToRefs(oids)); + ASSERT_EQ(object_manager_mock_.active_requests, prev_pull_reqs); + // Cancel the worker's `ray.wait`. + EXPECT_CALL(reconstruction_policy_mock_, Cancel(_)).Times(num_objects); + dependency_manager_.CancelWaitRequest(worker_id); + AssertNoLeaks(); +} + +/// Test that when one of the objects is already local at the time of the +/// `ray.wait` call, the `ray.wait` call does not trigger any requests to +/// remote nodes for that object. +TEST_F(DependencyManagerTest, TestWaitObjectLocal) { + // Generate a random worker and objects to wait on. + WorkerID worker_id = WorkerID::FromRandom(); + int num_objects = 3; + std::vector oids; + for (int i = 0; i < num_objects; i++) { + oids.push_back(ObjectID::FromRandom()); + } + // Simulate one of the objects becoming local. The later `ray.wait` call + // should have no effect because the object is already local. + const ObjectID local_object_id = std::move(oids.back()); + auto ready_task_ids = dependency_manager_.HandleObjectLocal(local_object_id); + ASSERT_TRUE(ready_task_ids.empty()); + + // Simulate a worker calling `ray.wait` on the objects. It should only make + // requests for the objects that are not local. + for (const auto &object_id : oids) { + if (object_id != local_object_id) { + EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(object_id, _)); + } + } + dependency_manager_.StartOrUpdateWaitRequest(worker_id, ObjectIdsToRefs(oids)); + ASSERT_EQ(object_manager_mock_.active_requests.size(), num_objects - 1); + // Simulate the local object getting evicted. The `ray.wait` call should not + // be reactivated. + auto waiting_task_ids = dependency_manager_.HandleObjectMissing(local_object_id); + ASSERT_TRUE(waiting_task_ids.empty()); + ASSERT_EQ(object_manager_mock_.active_requests.size(), num_objects - 1); + // Cancel the worker's `ray.wait`. + for (const auto &object_id : oids) { + if (object_id != local_object_id) { + EXPECT_CALL(reconstruction_policy_mock_, Cancel(object_id)); + } + } + dependency_manager_.CancelWaitRequest(worker_id); + AssertNoLeaks(); +} + +} // namespace raylet + +} // namespace ray + +int main(int argc, char **argv) { + ::testing::InitGoogleTest(&argc, argv); + return RUN_ALL_TESTS(); +} diff --git a/src/ray/raylet/node_manager.cc b/src/ray/raylet/node_manager.cc index 7f207c4fb..bcfc4e32f 100644 --- a/src/ray/raylet/node_manager.cc +++ b/src/ray/raylet/node_manager.cc @@ -158,7 +158,7 @@ NodeManager::NodeManager(boost::asio::io_service &io_service, const NodeID &self }, RayConfig::instance().object_timeout_milliseconds(), self_node_id_, gcs_client_, object_directory_), - task_dependency_manager_(object_manager, reconstruction_policy_), + dependency_manager_(object_manager, reconstruction_policy_), node_manager_server_("NodeManager", config.node_manager_port), node_manager_service_(io_service, *this), agent_manager_service_handler_( @@ -218,7 +218,7 @@ NodeManager::NodeManager(boost::asio::io_service &io_service, const NodeID &self PublishInfeasibleTaskError(task); }; cluster_task_manager_ = std::shared_ptr(new ClusterTaskManager( - self_node_id_, new_resource_scheduler_, task_dependency_manager_, is_owner_alive, + self_node_id_, new_resource_scheduler_, dependency_manager_, is_owner_alive, get_node_info_func, announce_infeasible_task)); placement_group_resource_manager_ = std::make_shared(new_resource_scheduler_); @@ -382,7 +382,7 @@ void NodeManager::HandleJobFinished(const JobID &job_id, const JobTableData &job for (const auto &worker : workers) { if (!worker->IsDetachedActor()) { // Clean up any open ray.wait calls that the worker made. - task_dependency_manager_.UnsubscribeWaitDependencies(worker->WorkerId()); + dependency_manager_.CancelWaitRequest(worker->WorkerId()); // Mark the worker as dead so further messages from it are ignored // (except DisconnectClient). worker->MarkDead(); @@ -390,18 +390,6 @@ void NodeManager::HandleJobFinished(const JobID &job_id, const JobTableData &job KillWorker(worker); } } - - if (!new_scheduler_enabled_) { - // Remove all tasks for this job from the scheduling queues, mark - // the results for these tasks as not required, cancel any attempts - // at reconstruction. Note that at this time the workers are likely - // alive because of the delay in killing workers. - auto tasks_to_remove = local_queues_.GetTaskIdsForJob(job_id); - task_dependency_manager_.RemoveTasksAndRelatedObjects(tasks_to_remove); - // NOTE(swang): SchedulingQueue::RemoveTasks modifies its argument so we must - // call it last. - local_queues_.RemoveTasks(tasks_to_remove); - } } void NodeManager::Heartbeat() { @@ -1034,7 +1022,7 @@ void NodeManager::ResourceUsageAdded(const NodeID &node_id, if (state != TaskState::INFEASIBLE) { // Don't unsubscribe for infeasible tasks because we never subscribed in // the first place. - RAY_CHECK(task_dependency_manager_.UnsubscribeGetDependencies(task_id)); + dependency_manager_.RemoveTaskDependencies(task_id); } // Attempt to forward the task. If this fails to forward the task, // the task will be resubmit locally. @@ -1403,7 +1391,7 @@ void NodeManager::ProcessDisconnectClientMessage( AsyncResolveObjectsFinish(client, task_id, true); } // Clean up any open ray.wait calls that the worker made. - task_dependency_manager_.UnsubscribeWaitDependencies(worker->WorkerId()); + dependency_manager_.CancelWaitRequest(worker->WorkerId()); } // Erase any lease metadata. @@ -1425,9 +1413,7 @@ void NodeManager::ProcessDisconnectClientMessage( // If the worker was an actor, it'll be cleaned by GCS. if (actor_id.IsNil()) { Task task; - if (local_queues_.RemoveTask(task_id, &task)) { - TreatTaskAsFailed(task, ErrorType::WORKER_DIED); - } + static_cast(local_queues_.RemoveTask(task_id, &task)); } if (!intentional_disconnect) { @@ -1501,14 +1487,14 @@ void NodeManager::ProcessFetchOrReconstructMessage( const auto refs = FlatbufferToObjectReference(*message->object_ids(), *message->owner_addresses()); if (message->fetch_only()) { - for (const auto &ref : refs) { - ObjectID object_id = ObjectID::FromBinary(ref.object_id()); - // If only a fetch is required, then do not subscribe to the - // dependencies to the task dependency manager. - if (!task_dependency_manager_.CheckObjectLocal(object_id)) { - // Fetch the object if it's not already local. - RAY_CHECK_OK(object_manager_.Pull(object_id, ref.owner_address())); - } + std::shared_ptr worker = worker_pool_.GetRegisteredWorker(client); + if (!worker) { + worker = worker_pool_.GetRegisteredDriver(client); + } + if (worker) { + // This will start a fetch for the objects that gets canceled once the + // objects are local, or if the worker dies. + dependency_manager_.StartOrUpdateWaitRequest(worker->WorkerId(), refs); } } else { // The values are needed. Add all requested objects to the list to @@ -1544,7 +1530,7 @@ void NodeManager::ProcessWaitRequestMessage( bool resolve_objects = false; for (auto const &object_id : object_ids) { - if (!task_dependency_manager_.CheckObjectLocal(object_id)) { + if (!dependency_manager_.CheckObjectLocal(object_id)) { // At least one object requires resolution. resolve_objects = true; } @@ -1904,11 +1890,7 @@ void NodeManager::HandleCancelWorkerLease(const rpc::CancelWorkerLeaseRequest &r bool canceled; if (new_scheduler_enabled_) { canceled = cluster_task_manager_->CancelTask(task_id); - if (canceled) { - // We have not yet granted the worker lease. Cancel it now. - task_dependency_manager_.TaskCanceled(task_id); - task_dependency_manager_.UnsubscribeGetDependencies(task_id); - } else { + if (!canceled) { // There are 2 cases here. // 1. We haven't received the lease request yet. It's the caller's job to // retry the cancellation once we've received the request. @@ -1925,8 +1907,9 @@ void NodeManager::HandleCancelWorkerLease(const rpc::CancelWorkerLeaseRequest &r if (removed_task.OnDispatch()) { // We have not yet granted the worker lease. Cancel it now. removed_task.OnCancellation()(); - task_dependency_manager_.TaskCanceled(task_id); - task_dependency_manager_.UnsubscribeGetDependencies(task_id); + if (removed_task_state == TaskState::WAITING) { + dependency_manager_.RemoveTaskDependencies(task_id); + } } else { // We already granted the worker lease and sent the reply. Re-queue the // task and wait for the requester to return the leased worker. @@ -2035,7 +2018,6 @@ void NodeManager::ScheduleTasks( // submission vs. registering remaining queued placeable tasks here. std::unordered_set move_task_set; for (const auto &task : local_queues_.GetTasks(TaskState::PLACEABLE)) { - task_dependency_manager_.TaskPending(task); move_task_set.insert(task.GetTaskSpecification().TaskId()); PublishInfeasibleTaskError(task); // Assert that this placeable task is not feasible locally (necessary but not @@ -2051,37 +2033,6 @@ void NodeManager::ScheduleTasks( RAY_CHECK(local_queues_.GetTasks(TaskState::PLACEABLE).size() == 0); } -void NodeManager::TreatTaskAsFailed(const Task &task, const ErrorType &error_type) { - const TaskSpecification &spec = task.GetTaskSpecification(); - RAY_LOG(DEBUG) << "Treating task " << spec.TaskId() << " as failed because of error " - << ErrorType_Name(error_type) << "."; - // Loop over the return IDs (except the dummy ID) and store a fake object in - // the object store. - int64_t num_returns = spec.NumReturns(); - if (spec.IsActorCreationTask()) { - // TODO(rkn): We subtract 1 to avoid the dummy ID. However, this leaks - // information about the TaskSpecification implementation. - num_returns -= 1; - } - // Determine which IDs should be marked as failed. - std::vector objects_to_fail; - for (int64_t i = 0; i < num_returns; i++) { - rpc::ObjectReference ref; - ref.set_object_id(spec.ReturnId(i).Binary()); - ref.mutable_owner_address()->CopyFrom(spec.CallerAddress()); - objects_to_fail.push_back(ref); - } - const JobID job_id = task.GetTaskSpecification().JobId(); - MarkObjectsAsFailed(error_type, objects_to_fail, job_id); - task_dependency_manager_.TaskCanceled(spec.TaskId()); - // Notify the task dependency manager that we no longer need this task's - // object dependencies. TODO(swang): Ideally, we would check the return value - // here. However, we don't know at this point if the task was in the WAITING - // or READY queue before, in which case we would not have been subscribed to - // its dependencies. - task_dependency_manager_.UnsubscribeGetDependencies(spec.TaskId()); -} - void NodeManager::MarkObjectsAsFailed( const ErrorType &error_type, const std::vector objects_to_fail, const JobID &job_id) { @@ -2189,7 +2140,7 @@ void NodeManager::HandleDirectCallTaskUnblocked( // First, always release task dependencies. This ensures we don't leak resources even // if we don't need to unblock the worker below. - task_dependency_manager_.UnsubscribeGetDependencies(task_id); + dependency_manager_.CancelGetRequest(worker->WorkerId()); if (new_scheduler_enabled_) { // Important: avoid double unblocking if the unblock RPC finishes after task end. @@ -2281,15 +2232,10 @@ void NodeManager::AsyncResolveObjects( // fetched and/or restarted as necessary, until the objects become local // or are unsubscribed. if (ray_get) { - // TODO(ekl) using the assigned task id is a hack to handle unsubscription for - // HandleDirectCallUnblocked. - auto &task_id = mark_worker_blocked ? current_task_id : worker->GetAssignedTaskId(); - if (!task_id.IsNil()) { - task_dependency_manager_.SubscribeGetDependencies(task_id, required_object_refs); - } + dependency_manager_.StartOrUpdateGetRequest(worker->WorkerId(), required_object_refs); } else { - task_dependency_manager_.SubscribeWaitDependencies(worker->WorkerId(), - required_object_refs); + dependency_manager_.StartOrUpdateWaitRequest(worker->WorkerId(), + required_object_refs); } } @@ -2341,13 +2287,13 @@ void NodeManager::AsyncResolveObjectsFinish( worker = worker_pool_.GetRegisteredDriver(client); } + RAY_CHECK(worker); // Unsubscribe from any `ray.get` objects that the task was blocked on. Any // fetch or reconstruction operations to make the objects local are canceled. // `ray.wait` calls will stay active until the objects become local, or the // task/actor that called `ray.wait` exits. - task_dependency_manager_.UnsubscribeGetDependencies(current_task_id); + dependency_manager_.CancelGetRequest(worker->WorkerId()); // Mark the task as unblocked. - RAY_CHECK(worker); if (was_blocked) { worker->RemoveBlockedTaskId(current_task_id); local_queues_.RemoveBlockedTaskId(current_task_id); @@ -2358,7 +2304,7 @@ void NodeManager::EnqueuePlaceableTask(const Task &task) { // TODO(atumanov): add task lookup hashmap and change EnqueuePlaceableTask to take // a vector of TaskIDs. Trigger MoveTask internally. // Subscribe to the task's dependencies. - bool args_ready = task_dependency_manager_.SubscribeGetDependencies( + bool args_ready = dependency_manager_.RequestTaskDependencies( task.GetTaskSpecification().TaskId(), task.GetDependencies()); // Enqueue the task. If all dependencies are available, then the task is queued // in the READY state, else the WAITING state. @@ -2369,10 +2315,6 @@ void NodeManager::EnqueuePlaceableTask(const Task &task) { } else { local_queues_.QueueTasks({task}, TaskState::WAITING); } - // Mark the task as pending. Once the task has finished execution, or once it - // has been forwarded to another node, the task must be marked as canceled in - // the TaskDependencyManager. - task_dependency_manager_.TaskPending(task); } void NodeManager::AssignTask(const std::shared_ptr &worker, @@ -2470,12 +2412,11 @@ bool NodeManager::FinishAssignedTask(const std::shared_ptr &wor } else { // If this was a non-actor task, then cancel any ray.wait calls that were // made during the task execution. - task_dependency_manager_.UnsubscribeWaitDependencies(worker.WorkerId()); + dependency_manager_.CancelWaitRequest(worker.WorkerId()); } // Notify the task dependency manager that this task has finished execution. - task_dependency_manager_.UnsubscribeGetDependencies(spec.TaskId()); - task_dependency_manager_.TaskCanceled(task_id); + dependency_manager_.CancelGetRequest(worker.WorkerId()); if (!spec.IsActorCreationTask()) { // Unset the worker's assigned task. We keep the assigned task ID for @@ -2507,8 +2448,7 @@ void NodeManager::HandleTaskReconstruction(const TaskID &task_id, const ObjectID &required_object_id) { // Get the owner's address. rpc::Address owner_addr; - bool has_owner = - task_dependency_manager_.GetOwnerAddress(required_object_id, &owner_addr); + bool has_owner = dependency_manager_.GetOwnerAddress(required_object_id, &owner_addr); if (has_owner) { if (!RayConfig::instance().object_pinning_enabled()) { // LRU eviction is enabled. The object may still be in scope, but we @@ -2573,7 +2513,7 @@ void NodeManager::HandleTaskReconstruction(const TaskID &task_id, void NodeManager::HandleObjectLocal(const ObjectID &object_id) { // Notify the task dependency manager that this object is local. - const auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(object_id); + const auto ready_task_ids = dependency_manager_.HandleObjectLocal(object_id); RAY_LOG(DEBUG) << "Object local " << object_id << ", " << " on " << self_node_id_ << ", " << ready_task_ids.size() << " tasks ready"; @@ -2621,7 +2561,7 @@ bool NodeManager::IsActorCreationTask(const TaskID &task_id) { void NodeManager::HandleObjectMissing(const ObjectID &object_id) { // Notify the task dependency manager that this object is no longer local. - const auto waiting_task_ids = task_dependency_manager_.HandleObjectMissing(object_id); + const auto waiting_task_ids = dependency_manager_.HandleObjectMissing(object_id); std::stringstream result; result << "Object missing " << object_id << ", " << " on " << self_node_id_ << ", " << waiting_task_ids.size() @@ -2689,10 +2629,6 @@ void NodeManager::ForwardTaskOrResubmit(const Task &task, const NodeID &node_man RAY_LOG(INFO) << "Failed to forward task " << task_id << " to node manager " << node_manager_id; - // Mark the failed task as pending to let other raylets know that we still - // have the task. TaskDependencyManager::TaskPending() is assumed to be - // idempotent. - task_dependency_manager_.TaskPending(task); // The task is not for an actor and may therefore be placed on another // node immediately. Send it to the scheduling policy to be placed again. local_queues_.QueueTasks({task}, TaskState::PLACEABLE); @@ -2743,7 +2679,7 @@ void NodeManager::FinishAssignTask(const std::shared_ptr &worke local_queues_.QueueTasks({assigned_task}, TaskState::RUNNING); // Notify the task dependency manager that we no longer need this task's // object dependencies. - RAY_CHECK(task_dependency_manager_.UnsubscribeGetDependencies(spec.TaskId())); + dependency_manager_.RemoveTaskDependencies(spec.TaskId()); } else { RAY_LOG(WARNING) << "Failed to send task to worker, disconnecting client"; // We failed to send the task to the worker, so disconnect the worker. @@ -2770,7 +2706,7 @@ void NodeManager::ProcessSubscribePlasmaReady( auto message = flatbuffers::GetRoot(message_data); ObjectID id = from_flatbuf(*message->object_id()); - if (task_dependency_manager_.CheckObjectLocal(id)) { + if (dependency_manager_.CheckObjectLocal(id)) { // Object is already local, so we directly fire the callback to tell the core worker // that the plasma object is ready. rpc::PlasmaObjectReadyRequest request; @@ -2797,8 +2733,7 @@ void NodeManager::ProcessSubscribePlasmaReady( // is local at this time but when the core worker was notified, the object is // is evicted. The core worker should be able to handle evicted object in this // case. - task_dependency_manager_.SubscribeWaitDependencies(associated_worker->WorkerId(), - refs); + dependency_manager_.StartOrUpdateWaitRequest(associated_worker->WorkerId(), refs); // Add this worker to the listeners for the object ID. { @@ -2864,7 +2799,7 @@ std::string NodeManager::DebugString() const { result << "\n" << worker_pool_.DebugString(); result << "\n" << local_queues_.DebugString(); result << "\n" << reconstruction_policy_.DebugString(); - result << "\n" << task_dependency_manager_.DebugString(); + result << "\n" << dependency_manager_.DebugString(); { absl::MutexLock guard(&plasma_object_notification_lock_); result << "\nnum async plasma notifications: " diff --git a/src/ray/raylet/node_manager.h b/src/ray/raylet/node_manager.h index a734ebaec..63f741af0 100644 --- a/src/ray/raylet/node_manager.h +++ b/src/ray/raylet/node_manager.h @@ -35,7 +35,7 @@ #include "ray/raylet/scheduling_policy.h" #include "ray/raylet/scheduling_queue.h" #include "ray/raylet/reconstruction_policy.h" -#include "ray/raylet/task_dependency_manager.h" +#include "ray/raylet/dependency_manager.h" #include "ray/raylet/worker_pool.h" #include "ray/rpc/worker/core_worker_client_pool.h" #include "ray/util/ordered_set.h" @@ -239,18 +239,6 @@ class NodeManager : public rpc::NodeManagerServiceHandler { /// \param task The task in question. /// \return Void. void EnqueuePlaceableTask(const Task &task); - /// This will treat a task removed from the local queue as if it had been - /// executed and failed. This is done by looping over the task return IDs and - /// for each ID storing an object that represents a failure in the object - /// store. When clients retrieve these objects, they will raise - /// application-level exceptions. State for the task will be cleaned up as if - /// it were any other task that had been assigned, executed, and removed from - /// the local queue. - /// - /// \param task The task to fail. - /// \param error_type The type of the error that caused this task to fail. - /// \return Void. - void TreatTaskAsFailed(const Task &task, const ErrorType &error_type); /// Mark the specified objects as failed with the given error type. /// /// \param error_type The type of the error that caused this task to fail. @@ -707,8 +695,9 @@ class NodeManager : public rpc::NodeManagerServiceHandler { SchedulingPolicy scheduling_policy_; /// The reconstruction policy for deciding when to re-execute a task. ReconstructionPolicy reconstruction_policy_; - /// A manager to make waiting tasks's missing object dependencies available. - TaskDependencyManager task_dependency_manager_; + /// A manager to resolve objects needed by queued tasks and workers that + /// called `ray.get` or `ray.wait`. + DependencyManager dependency_manager_; std::unique_ptr agent_manager_; diff --git a/src/ray/raylet/scheduling/cluster_task_manager.cc b/src/ray/raylet/scheduling/cluster_task_manager.cc index ab3b46ea7..fc03a2a77 100644 --- a/src/ray/raylet/scheduling/cluster_task_manager.cc +++ b/src/ray/raylet/scheduling/cluster_task_manager.cc @@ -102,7 +102,7 @@ bool ClusterTaskManager::WaitForTaskArgsRequests(Work work) { auto object_ids = task.GetTaskSpecification().GetDependencies(); bool can_dispatch = true; if (object_ids.size() > 0) { - bool args_ready = task_dependency_manager_.SubscribeGetDependencies( + bool args_ready = task_dependency_manager_.RequestTaskDependencies( task.GetTaskSpecification().TaskId(), task.GetDependencies()); if (args_ready) { RAY_LOG(DEBUG) << "Args already ready, task can be dispatched " @@ -164,7 +164,8 @@ void ClusterTaskManager::DispatchScheduledTasksToWorkers( << "'s caller is no longer running. Cancelling task."; worker_pool.PushWorker(worker); if (!spec.GetDependencies().empty()) { - RAY_CHECK(task_dependency_manager_.UnsubscribeGetDependencies(spec.TaskId())); + task_dependency_manager_.RemoveTaskDependencies( + task.GetTaskSpecification().TaskId()); } work_it = dispatch_queue.erase(work_it); } else { @@ -179,7 +180,8 @@ void ClusterTaskManager::DispatchScheduledTasksToWorkers( } if (remove) { if (!spec.GetDependencies().empty()) { - RAY_CHECK(task_dependency_manager_.UnsubscribeGetDependencies(spec.TaskId())); + task_dependency_manager_.RemoveTaskDependencies( + task.GetTaskSpecification().TaskId()); } work_it = dispatch_queue.erase(work_it); } else { @@ -313,7 +315,8 @@ bool ClusterTaskManager::CancelTask(const TaskID &task_id) { RemoveFromBacklogTracker(task); ReplyCancelled(*work_it); if (!task.GetTaskSpecification().GetDependencies().empty()) { - RAY_CHECK(task_dependency_manager_.UnsubscribeGetDependencies(task_id)); + task_dependency_manager_.RemoveTaskDependencies( + task.GetTaskSpecification().TaskId()); } work_queue.erase(work_it); if (work_queue.empty()) { @@ -347,9 +350,11 @@ bool ClusterTaskManager::CancelTask(const TaskID &task_id) { RemoveFromBacklogTracker(task); ReplyCancelled(iter->second); if (!task.GetTaskSpecification().GetDependencies().empty()) { - task_dependency_manager_.UnsubscribeGetDependencies(task_id); + task_dependency_manager_.RemoveTaskDependencies(task_id); } waiting_tasks_.erase(iter); + + task_dependency_manager_.RemoveTaskDependencies(task_id); return true; } diff --git a/src/ray/raylet/scheduling/cluster_task_manager.h b/src/ray/raylet/scheduling/cluster_task_manager.h index 61cfce031..269954eb9 100644 --- a/src/ray/raylet/scheduling/cluster_task_manager.h +++ b/src/ray/raylet/scheduling/cluster_task_manager.h @@ -4,8 +4,8 @@ #include "absl/container/flat_hash_set.h" #include "ray/common/task/task.h" #include "ray/common/task/task_common.h" +#include "ray/raylet/dependency_manager.h" #include "ray/raylet/scheduling/cluster_resource_scheduler.h" -#include "ray/raylet/task_dependency_manager.h" #include "ray/raylet/worker.h" #include "ray/raylet/worker_pool.h" #include "ray/rpc/grpc_client.h" @@ -53,7 +53,7 @@ class ClusterTaskManager { /// \param gcs_client: A gcs client. ClusterTaskManager(const NodeID &self_node_id, std::shared_ptr cluster_resource_scheduler, - TaskDependencyManagerInterface &task_dependency_manager_, + TaskDependencyManagerInterface &task_dependency_manager, std::function is_owner_alive, NodeInfoGetter get_node_info, std::function announce_infeasible_task); diff --git a/src/ray/raylet/scheduling/cluster_task_manager_test.cc b/src/ray/raylet/scheduling/cluster_task_manager_test.cc index 3f33cbf06..e2ca4c5b7 100644 --- a/src/ray/raylet/scheduling/cluster_task_manager_test.cc +++ b/src/ray/raylet/scheduling/cluster_task_manager_test.cc @@ -96,14 +96,14 @@ Task CreateTask(const std::unordered_map &required_resource class MockTaskDependencyManager : public TaskDependencyManagerInterface { public: - bool SubscribeGetDependencies( + bool RequestTaskDependencies( const TaskID &task_id, const std::vector &required_objects) { RAY_CHECK(subscribed_tasks.insert(task_id).second); return task_ready_; } - bool UnsubscribeGetDependencies(const TaskID &task_id) { - return subscribed_tasks.erase(task_id); + void RemoveTaskDependencies(const TaskID &task_id) { + RAY_CHECK(subscribed_tasks.erase(task_id)); } bool IsTaskReady(const TaskID &task_id) const { return task_ready_; } diff --git a/src/ray/raylet/task_dependency_manager.cc b/src/ray/raylet/task_dependency_manager.cc deleted file mode 100644 index 74c3d8c7a..000000000 --- a/src/ray/raylet/task_dependency_manager.cc +++ /dev/null @@ -1,474 +0,0 @@ -// Copyright 2017 The Ray Authors. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -#include "ray/raylet/task_dependency_manager.h" - -#include "absl/time/clock.h" -#include "ray/stats/stats.h" - -namespace ray { - -namespace raylet { - -TaskDependencyManager::TaskDependencyManager( - ObjectManagerInterface &object_manager, - ReconstructionPolicyInterface &reconstruction_policy) - : object_manager_(object_manager), reconstruction_policy_(reconstruction_policy) {} - -bool TaskDependencyManager::CheckObjectLocal(const ObjectID &object_id) const { - return local_objects_.count(object_id) == 1; -} - -bool TaskDependencyManager::CheckObjectRequired(const ObjectID &object_id, - rpc::Address *owner_address) const { - const TaskID task_id = object_id.TaskId(); - auto task_entry = required_tasks_.find(task_id); - // If there are no subscribed tasks that are dependent on the object, then do - // nothing. - if (task_entry == required_tasks_.end()) { - return false; - } - if (task_entry->second.count(object_id) == 0) { - return false; - } - // If the object is already local, then the dependency is fulfilled. Do - // nothing. - if (local_objects_.count(object_id) == 1) { - return false; - } - // If the task that creates the object is pending execution, then the - // dependency will be fulfilled locally. Do nothing. - if (pending_tasks_.count(task_id) == 1) { - return false; - } - if (owner_address != nullptr) { - *owner_address = task_entry->second.at(object_id).owner_address; - } - return true; -} - -void TaskDependencyManager::HandleRemoteDependencyRequired(const ObjectID &object_id) { - rpc::Address owner_address; - bool required = CheckObjectRequired(object_id, &owner_address); - // If the object is required, then try to make the object available locally. - if (required) { - auto inserted = required_objects_.insert(object_id); - if (inserted.second) { - // If we haven't already, request the object manager to pull it from a - // remote node. - RAY_CHECK_OK(object_manager_.Pull(object_id, owner_address)); - reconstruction_policy_.ListenAndMaybeReconstruct(object_id, owner_address); - } - } -} - -void TaskDependencyManager::HandleRemoteDependencyCanceled(const ObjectID &object_id) { - bool required = CheckObjectRequired(object_id, nullptr); - // If the object is no longer required, then cancel the object. - if (!required) { - auto it = required_objects_.find(object_id); - if (it != required_objects_.end()) { - object_manager_.CancelPull(object_id); - reconstruction_policy_.Cancel(object_id); - required_objects_.erase(it); - } - } -} - -std::vector TaskDependencyManager::HandleObjectLocal( - const ray::ObjectID &object_id) { - // Add the object to the table of locally available objects. - auto inserted = local_objects_.insert(object_id); - RAY_CHECK(inserted.second) << object_id; - - // Find all tasks and workers that depend on the newly available object. - std::vector ready_task_ids; - auto creating_task_entry = required_tasks_.find(object_id.TaskId()); - if (creating_task_entry != required_tasks_.end()) { - auto object_entry = creating_task_entry->second.find(object_id); - if (object_entry != creating_task_entry->second.end()) { - // Loop through all tasks that depend on the newly available object. - for (const auto &dependent_task_id : object_entry->second.dependent_tasks) { - auto &task_entry = task_dependencies_[dependent_task_id]; - task_entry.num_missing_get_dependencies--; - // If the dependent task now has all of its arguments ready, it's ready - // to run. - if (task_entry.num_missing_get_dependencies == 0) { - ready_task_ids.push_back(dependent_task_id); - } - } - // Remove the dependency from all workers that called `ray.wait` on the - // newly available object. - for (const auto &worker_id : object_entry->second.dependent_workers) { - RAY_CHECK(worker_dependencies_[worker_id].erase(object_id) > 0); - } - // Clear all workers that called `ray.wait` on this object, since the - // `ray.wait` calls can now return the object as ready. - object_entry->second.dependent_workers.clear(); - - // If there are no more tasks or workers dependent on the local object or - // the task that created it, then remove the entry completely. - if (object_entry->second.Empty()) { - creating_task_entry->second.erase(object_entry); - if (creating_task_entry->second.empty()) { - required_tasks_.erase(creating_task_entry); - } - } - } - } - - // The object is now local, so cancel any in-progress operations to make the - // object local. - HandleRemoteDependencyCanceled(object_id); - - return ready_task_ids; -} - -std::vector TaskDependencyManager::HandleObjectMissing( - const ray::ObjectID &object_id) { - // Remove the object from the table of locally available objects. - auto erased = local_objects_.erase(object_id); - RAY_CHECK(erased == 1); - - // Find any tasks that are dependent on the missing object. - std::vector waiting_task_ids; - TaskID creating_task_id = object_id.TaskId(); - auto creating_task_entry = required_tasks_.find(creating_task_id); - if (creating_task_entry != required_tasks_.end()) { - auto object_entry = creating_task_entry->second.find(object_id); - if (object_entry != creating_task_entry->second.end()) { - for (auto &dependent_task_id : object_entry->second.dependent_tasks) { - auto &task_entry = task_dependencies_[dependent_task_id]; - // If the dependent task had all of its arguments ready, it was ready to - // run but must be switched to waiting since one of its arguments is now - // missing. - if (task_entry.num_missing_get_dependencies == 0) { - waiting_task_ids.push_back(dependent_task_id); - // During normal execution we should be able to include the check - // RAY_CHECK(pending_tasks_.count(dependent_task_id) == 1); - // However, this invariant will not hold during unit test execution. - } - task_entry.num_missing_get_dependencies++; - } - } - } - // The object is no longer local. Try to make the object local if necessary. - HandleRemoteDependencyRequired(object_id); - // Process callbacks for all of the tasks dependent on the object that are - // now ready to run. - return waiting_task_ids; -} - -bool TaskDependencyManager::SubscribeGetDependencies( - const TaskID &task_id, const std::vector &required_objects) { - auto &task_entry = task_dependencies_[task_id]; - - // Record the task's dependencies. - for (const auto &object : required_objects) { - const auto &object_id = ObjectID::FromBinary(object.object_id()); - auto inserted = task_entry.get_dependencies.insert(object_id); - if (inserted.second) { - RAY_LOG(DEBUG) << "Task " << task_id << " blocked on object " << object_id; - // Get the ID of the task that creates the dependency. - TaskID creating_task_id = object_id.TaskId(); - // Determine whether the dependency can be fulfilled by the local node. - if (local_objects_.count(object_id) == 0) { - // The object is not local. - task_entry.num_missing_get_dependencies++; - } - - auto it = required_tasks_[creating_task_id].find(object_id); - if (it == required_tasks_[creating_task_id].end()) { - it = required_tasks_[creating_task_id] - .emplace(object_id, ObjectDependencies(object)) - .first; - } - // Add the subscribed task to the mapping from object ID to list of - // dependent tasks. - it->second.dependent_tasks.insert(task_id); - } - } - - // These dependencies are required by the given task. Try to make them local - // if necessary. - for (const auto &object : required_objects) { - const auto &object_id = ObjectID::FromBinary(object.object_id()); - HandleRemoteDependencyRequired(object_id); - } - - // Return whether all dependencies are local. - return (task_entry.num_missing_get_dependencies == 0); -} - -bool TaskDependencyManager::IsTaskReady(const TaskID &task_id) const { - auto task_entry = task_dependencies_.find(task_id); - RAY_CHECK(task_entry != task_dependencies_.end()); - return task_entry->second.num_missing_get_dependencies == 0; -} - -void TaskDependencyManager::SubscribeWaitDependencies( - const WorkerID &worker_id, - const std::vector &required_objects) { - auto &worker_entry = worker_dependencies_[worker_id]; - - // Record the worker's dependencies. - for (const auto &object : required_objects) { - const auto &object_id = ObjectID::FromBinary(object.object_id()); - if (local_objects_.count(object_id) == 0) { - RAY_LOG(DEBUG) << "Worker " << worker_id << " called ray.wait on remote object " - << object_id; - // Only add the dependency if the object is not local. If the object is - // local, then the `ray.wait` call can already return it. - auto inserted = worker_entry.insert(object_id); - if (inserted.second) { - // Get the ID of the task that creates the dependency. - TaskID creating_task_id = object_id.TaskId(); - auto it = required_tasks_[creating_task_id].find(object_id); - if (it == required_tasks_[creating_task_id].end()) { - it = required_tasks_[creating_task_id] - .emplace(object_id, ObjectDependencies(object)) - .first; - } - // Add the subscribed worker to the mapping from object ID to list of - // dependent workers. - it->second.dependent_workers.insert(worker_id); - } - } - } - - // These dependencies are required by the given worker. Try to make them - // local if necessary. - for (const auto &object : required_objects) { - const auto &object_id = ObjectID::FromBinary(object.object_id()); - HandleRemoteDependencyRequired(object_id); - } -} - -bool TaskDependencyManager::UnsubscribeGetDependencies(const TaskID &task_id) { - RAY_LOG(DEBUG) << "Task " << task_id << " no longer blocked"; - // Remove the task from the table of subscribed tasks. - auto it = task_dependencies_.find(task_id); - if (it == task_dependencies_.end()) { - return false; - } - const TaskDependencies task_entry = std::move(it->second); - task_dependencies_.erase(it); - - // Remove the task's dependencies. - for (const auto &object_id : task_entry.get_dependencies) { - // Get the ID of the task that creates the dependency. - TaskID creating_task_id = object_id.TaskId(); - auto creating_task_entry = required_tasks_.find(creating_task_id); - // Remove the task from the list of tasks that are dependent on this - // object. - auto it = creating_task_entry->second.find(object_id); - RAY_CHECK(it != creating_task_entry->second.end()); - RAY_CHECK(it->second.dependent_tasks.erase(task_id) > 0); - // If nothing else depends on the object, then erase the object entry. - if (it->second.Empty()) { - creating_task_entry->second.erase(it); - // Remove the task that creates this object if there are no more object - // dependencies created by the task. - if (creating_task_entry->second.empty()) { - required_tasks_.erase(creating_task_entry); - } - } - } - - // These dependencies are no longer required by the given task. Cancel any - // in-progress operations to make them local. - for (const auto &object_id : task_entry.get_dependencies) { - HandleRemoteDependencyCanceled(object_id); - } - - return true; -} - -void TaskDependencyManager::UnsubscribeWaitDependencies(const WorkerID &worker_id) { - RAY_LOG(DEBUG) << "Worker " << worker_id << " no longer blocked"; - // Remove the task from the table of subscribed tasks. - auto it = worker_dependencies_.find(worker_id); - if (it == worker_dependencies_.end()) { - return; - } - const WorkerDependencies worker_entry = std::move(it->second); - worker_dependencies_.erase(it); - - // Remove the task's dependencies. - for (const auto &object_id : worker_entry) { - // Get the ID of the task that creates the dependency. - TaskID creating_task_id = object_id.TaskId(); - auto creating_task_entry = required_tasks_.find(creating_task_id); - // Remove the worker from the list of workers that are dependent on this - // object. - auto it = creating_task_entry->second.find(object_id); - RAY_CHECK(it != creating_task_entry->second.end()); - RAY_CHECK(it->second.dependent_workers.erase(worker_id) > 0); - // If nothing else depends on the object, then erase the object entry. - if (it->second.Empty()) { - creating_task_entry->second.erase(it); - // Remove the task that creates this object if there are no more object - // dependencies created by the task. - if (creating_task_entry->second.empty()) { - required_tasks_.erase(creating_task_entry); - } - } - } - - // These dependencies are no longer required by the given task. Cancel any - // in-progress operations to make them local. - for (const auto &object_id : worker_entry) { - HandleRemoteDependencyCanceled(object_id); - } -} - -void TaskDependencyManager::TaskPending(const Task &task) { - // Direct tasks are not tracked by the raylet. - // NOTE(zhijunfu): Direct tasks are not tracked by the raylet, - // but we still need raylet to reconstruct the actors. - // For direct actor creation task: - // - Initially the caller leases a worker from raylet and - // then pushes actor creation task directly to the worker, - // thus it doesn't need task lease. And actually if we - // acquire a lease in this case and forget to cancel it, - // the lease would never expire which will prevent the - // actor from being restarted; - // - When a direct actor is restarted, raylet resubmits - // the task, and the task can be forwarded to another raylet, - // and eventually assigned to a worker. In this case we need - // the task lease to make sure there's only one raylet can - // resubmit the task. - // - // We can use `OnDispatch` to differeniate whether this task is - // a worker lease request. - // For direct actor creation task: - // - when it's submitted by core worker, we guarantee that - // we always request a new worker lease, in that case - // `OnDispatch` is overridden to an actual callback. - // - when it's resubmitted by raylet because of reconstruction, - // `OnDispatch` will not be overridden and thus is nullptr. - if (task.GetTaskSpecification().IsActorCreationTask() && task.OnDispatch() == nullptr) { - // This is an actor creation task, and it's being restarted, - // in this case we still need the task lease. Note that we don't - // require task lease for direct actor creation task. - } else { - return; - } - - TaskID task_id = task.GetTaskSpecification().TaskId(); - RAY_LOG(DEBUG) << "Task execution " << task_id << " pending"; - - // Record that the task is pending execution. - auto inserted = pending_tasks_.insert(task_id); - if (inserted.second) { - // This is the first time we've heard that this task is pending. Find any - // subscribed tasks that are dependent on objects created by the pending - // task. - auto remote_task_entry = required_tasks_.find(task_id); - if (remote_task_entry != required_tasks_.end()) { - for (const auto &object_entry : remote_task_entry->second) { - // This object created by the pending task will appear locally once the - // task completes execution. Cancel any in-progress operations to make - // the object local. - HandleRemoteDependencyCanceled(object_entry.first); - } - } - } -} - -void TaskDependencyManager::TaskCanceled(const TaskID &task_id) { - RAY_LOG(DEBUG) << "Task execution " << task_id << " canceled"; - // Record that the task is no longer pending execution. - auto it = pending_tasks_.find(task_id); - if (it == pending_tasks_.end()) { - return; - } - pending_tasks_.erase(it); - - // Find any subscribed tasks that are dependent on objects created by the - // canceled task. - auto remote_task_entry = required_tasks_.find(task_id); - if (remote_task_entry != required_tasks_.end()) { - for (const auto &object_entry : remote_task_entry->second) { - // This object created by the task will no longer appear locally since - // the task is canceled. Try to make the object local if necessary. - HandleRemoteDependencyRequired(object_entry.first); - } - } -} - -void TaskDependencyManager::RemoveTasksAndRelatedObjects( - const std::unordered_set &task_ids) { - // Collect a list of all the unique objects that these tasks were subscribed - // to. - std::unordered_set required_objects; - for (auto it = task_ids.begin(); it != task_ids.end(); it++) { - auto task_it = task_dependencies_.find(*it); - if (task_it != task_dependencies_.end()) { - // Add the objects that this task was subscribed to. - required_objects.insert(task_it->second.get_dependencies.begin(), - task_it->second.get_dependencies.end()); - } - // The task no longer depends on anything. - task_dependencies_.erase(*it); - // The task is no longer pending execution. - pending_tasks_.erase(*it); - } - - // Cancel all of the objects that were required by the removed tasks. - for (const auto &object_id : required_objects) { - TaskID creating_task_id = object_id.TaskId(); - required_tasks_.erase(creating_task_id); - HandleRemoteDependencyCanceled(object_id); - } - - // Make sure that the tasks in task_ids no longer have tasks dependent on - // them. - for (const auto &task_id : task_ids) { - RAY_CHECK(required_tasks_.find(task_id) == required_tasks_.end()) - << "RemoveTasksAndRelatedObjects was called on " << task_id - << ", but another task depends on it that was not included in the argument"; - } -} - -std::string TaskDependencyManager::DebugString() const { - std::stringstream result; - result << "TaskDependencyManager:"; - result << "\n- task dep map size: " << task_dependencies_.size(); - result << "\n- task req map size: " << required_tasks_.size(); - result << "\n- req objects map size: " << required_objects_.size(); - result << "\n- local objects map size: " << local_objects_.size(); - result << "\n- pending tasks map size: " << pending_tasks_.size(); - return result.str(); -} - -bool TaskDependencyManager::GetOwnerAddress(const ObjectID &object_id, - rpc::Address *owner_address) const { - const auto creating_task_entry = required_tasks_.find(object_id.TaskId()); - if (creating_task_entry == required_tasks_.end()) { - return false; - } - - const auto it = creating_task_entry->second.find(object_id); - if (it == creating_task_entry->second.end()) { - return false; - } - - *owner_address = it->second.owner_address; - return !owner_address->worker_id().empty(); -} - -} // namespace raylet - -} // namespace ray diff --git a/src/ray/raylet/task_dependency_manager.h b/src/ray/raylet/task_dependency_manager.h deleted file mode 100644 index eb2c53ee9..000000000 --- a/src/ray/raylet/task_dependency_manager.h +++ /dev/null @@ -1,260 +0,0 @@ -// Copyright 2017 The Ray Authors. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -#pragma once - -// clang-format off -#include "ray/common/id.h" -#include "ray/common/task/task.h" -#include "ray/object_manager/object_manager.h" -#include "ray/raylet/reconstruction_policy.h" -// clang-format on - -namespace ray { - -namespace raylet { - -using rpc::TaskLeaseData; - -class ReconstructionPolicy; - -/// Used for unit-testing the ClusterTaskManager, which calls these methods for -/// locally queued tasks that have dependencies. -class TaskDependencyManagerInterface { - public: - virtual bool SubscribeGetDependencies( - const TaskID &task_id, - const std::vector &required_objects) = 0; - virtual bool IsTaskReady(const TaskID &task_id) const = 0; - virtual bool UnsubscribeGetDependencies(const TaskID &task_id) = 0; - virtual ~TaskDependencyManagerInterface() {} -}; - -/// \class TaskDependencyManager -/// -/// Responsible for managing object dependencies for tasks. The caller can -/// subscribe to object dependencies for a task. The task manager will -/// determine which object dependencies are remote. These are the objects that -/// are neither in the local object store, nor will they be created by a -/// locally queued task. The task manager will request that these objects be -/// made available locally, either by object transfer from a remote node or -/// reconstruction. The task manager will also cancel these objects if they are -/// no longer needed by any task. -class TaskDependencyManager : public TaskDependencyManagerInterface { - public: - /// Create a task dependency manager. - TaskDependencyManager(ObjectManagerInterface &object_manager, - ReconstructionPolicyInterface &reconstruction_policy); - - /// Check whether an object is locally available. - /// - /// \param object_id The object to check for. - /// \return Whether the object is local. - bool CheckObjectLocal(const ObjectID &object_id) const; - - /// Subscribe to object depedencies required by the task and check whether - /// all dependencies are fulfilled. This should be called for task arguments and - /// `ray.get` calls during task execution. - /// - /// The TaskDependencyManager will track the task's dependencies - /// until UnsubscribeGetDependencies is called on the same task ID. If any - /// dependencies are remote, then they will be requested. When the last - /// remote dependency later appears locally via a call to HandleObjectLocal, - /// the subscribed task will be returned by the HandleObjectLocal call, - /// signifying that it is ready to run. This method may be called multiple - /// times per task. - /// - /// \param task_id The ID of the task whose dependencies to subscribe to. - /// \param required_objects The objects required by the task. - /// \return Whether all of the given dependencies for the given task are - /// local. - bool SubscribeGetDependencies( - const TaskID &task_id, const std::vector &required_objects); - - /// Check whether a task is ready to run. The task ID must - /// have been previously subscribed by the caller. - /// - /// \param task_id The ID of the task to check. - /// \return Whether all of the dependencies for the task are - /// local. - bool IsTaskReady(const TaskID &task_id) const; - - /// Subscribe to object depedencies required by the worker. This should be called for - /// ray.wait calls during task execution. - /// - /// The TaskDependencyManager will track all remote dependencies until the - /// dependencies are local, or until UnsubscribeWaitDependencies is called - /// with the same worker ID, whichever occurs first. Remote dependencies will - /// be requested. This method may be called multiple times per worker on the - /// same objects. - /// - /// \param worker_id The ID of the worker that called `ray.wait`. - /// \param required_objects The objects required by the worker. - /// \return Void. - void SubscribeWaitDependencies( - const WorkerID &worker_id, - const std::vector &required_objects); - - /// Unsubscribe from the object dependencies required by this task through the task - /// arguments or `ray.get`. If the objects were remote and are no longer required by any - /// subscribed task, then they will be canceled. - /// - /// \param task_id The ID of the task whose dependencies we should unsubscribe from. - /// \return Whether the task was subscribed before. - bool UnsubscribeGetDependencies(const TaskID &task_id); - - /// Unsubscribe from the object dependencies required by this worker through `ray.wait`. - /// If the objects were remote and are no longer required by any subscribed task, then - /// they will be canceled. - /// - /// \param worker_id The ID of the worker whose dependencies we should unsubscribe from. - /// \return The objects that the worker was waiting on. - void UnsubscribeWaitDependencies(const WorkerID &worker_id); - - /// Mark that the given task is pending execution. Any objects that it creates - /// are now considered to be pending creation. If there are any subscribed - /// tasks that depend on these objects, then the objects will be canceled. - /// - /// \param task The task that is pending execution. - void TaskPending(const Task &task); - - /// Mark that the given task is no longer pending execution. Any objects that - /// it creates that are not already local are now considered to be remote. If - /// there are any subscribed tasks that depend on these objects, then the - /// objects will be requested. - /// - /// \param task_id The ID of the task to cancel. - void TaskCanceled(const TaskID &task_id); - - /// Handle an object becoming locally available. If there are any subscribed - /// tasks that depend on this object, then the object will be canceled. - /// - /// \param object_id The object ID of the object to mark as locally - /// available. - /// \return A list of task IDs. This contains all subscribed tasks that now - /// have all of their dependencies fulfilled, once this object was made - /// local. - std::vector HandleObjectLocal(const ray::ObjectID &object_id); - - /// Handle an object that is no longer locally available. If there are any - /// subscribed tasks that depend on this object, then the object will be - /// requested. - /// - /// \param object_id The object ID of the object that was previously locally - /// available. - /// \return A list of task IDs. This contains all subscribed tasks that - /// previously had all of their dependencies fulfilled, but are now missing - /// this object dependency. - std::vector HandleObjectMissing(const ray::ObjectID &object_id); - - /// Remove all of the tasks specified. These tasks will no longer be - /// considered pending and the objects they depend on will no longer be - /// required. - /// - /// \param task_ids The collection of task IDs. For a given task in this set, - /// all tasks that depend on the task must also be included in the set. - void RemoveTasksAndRelatedObjects(const std::unordered_set &task_ids); - - /// Returns debug string for class. - /// - /// \return string. - std::string DebugString() const; - - /// Get the address of the owner of this object. An address will only be - /// returned if the caller previously specified that this object is required - /// on this node, through a call to SubscribeGetDependencies or - /// SubscribeWaitDependencies. - /// - /// \param[in] object_id The object whose owner to get. - /// \param[out] owner_address The address of the object's owner, if - /// available. - /// \return True if we have owner information for the object. - bool GetOwnerAddress(const ObjectID &object_id, rpc::Address *owner_address) const; - - private: - struct ObjectDependencies { - ObjectDependencies(const rpc::ObjectReference &ref) - : owner_address(ref.owner_address()) {} - /// The tasks that depend on this object, either because the object is a task argument - /// or because the task called `ray.get` on the object. - std::unordered_set dependent_tasks; - /// The workers that depend on this object because they called `ray.wait` on the - /// object. - std::unordered_set dependent_workers; - /// The address of the worker that owns this object. - rpc::Address owner_address; - - bool Empty() const { return dependent_tasks.empty() && dependent_workers.empty(); } - }; - - /// A struct to represent the object dependencies of a task. - struct TaskDependencies { - /// The objects that the task depends on. These are either the arguments to - /// the task or objects that the task calls `ray.get` on. These must be - /// local before the task is ready to execute. Objects are removed from - /// this set once UnsubscribeGetDependencies is called. - std::unordered_set get_dependencies; - /// The number of object arguments that are not available locally. This - /// must be zero before the task is ready to execute. - int64_t num_missing_get_dependencies; - }; - - /// The objects that the worker is fetching. These are objects that a task that executed - /// or is executing on the worker called `ray.wait` on that are not yet local. An object - /// will be automatically removed from this set once it becomes local. - using WorkerDependencies = std::unordered_set; - - /// Check whether the given object needs to be made available through object - /// transfer or reconstruction. These are objects for which: (1) there is a - /// subscribed task dependent on it, (2) the object is not local, and (3) the - /// task that creates the object is not pending execution locally. - bool CheckObjectRequired(const ObjectID &object_id, rpc::Address *owner_address) const; - /// If the given object is required, then request that the object be made - /// available through object transfer or reconstruction. - void HandleRemoteDependencyRequired(const ObjectID &object_id); - /// If the given object is no longer required, then cancel any in-progress - /// operations to make the object available through object transfer or - /// reconstruction. - void HandleRemoteDependencyCanceled(const ObjectID &object_id); - - /// The object manager, used to fetch required objects from remote nodes. - ObjectManagerInterface &object_manager_; - /// The reconstruction policy, used to reconstruct required objects that no - /// longer exist on any live nodes. - ReconstructionPolicyInterface &reconstruction_policy_; - /// A mapping from task ID of each subscribed task to its list of object - /// dependencies, either task arguments or objects passed into `ray.get`. - std::unordered_map task_dependencies_; - /// A mapping from worker ID to each object that the worker called `ray.wait` on. - std::unordered_map worker_dependencies_; - /// All tasks whose outputs are required by a subscribed task. This is a - /// mapping from task ID to information about the objects that the task - /// creates, either by return value or by `ray.put`. For each object, we - /// store the IDs of the subscribed tasks that are dependent on the object. - std::unordered_map> - required_tasks_; - /// Objects that are required by a subscribed task, are not local, and are - /// not created by a pending task. For these objects, there are pending - /// operations to make the object available. - std::unordered_set required_objects_; - /// The set of locally available objects. - std::unordered_set local_objects_; - /// The set of tasks that are pending execution. Any objects created by these - /// tasks that are not already local are pending creation. - std::unordered_set pending_tasks_; -}; - -} // namespace raylet - -} // namespace ray diff --git a/src/ray/raylet/task_dependency_manager_test.cc b/src/ray/raylet/task_dependency_manager_test.cc deleted file mode 100644 index d65b0aced..000000000 --- a/src/ray/raylet/task_dependency_manager_test.cc +++ /dev/null @@ -1,559 +0,0 @@ -// Copyright 2017 The Ray Authors. -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -#include "ray/raylet/task_dependency_manager.h" - -#include -#include - -#include "gmock/gmock.h" -#include "gtest/gtest.h" -#include "ray/common/task/task_util.h" -#include "ray/common/test_util.h" - -namespace ray { - -namespace raylet { - -using ::testing::_; - -const static JobID kDefaultJobId = JobID::FromInt(1); - -const static TaskID kDefaultDriverTaskId = TaskID::ForDriverTask(kDefaultJobId); - -class MockObjectManager : public ObjectManagerInterface { - public: - MOCK_METHOD2(Pull, - ray::Status(const ObjectID &object_id, const rpc::Address &owner_address)); - MOCK_METHOD1(CancelPull, void(const ObjectID &object_id)); -}; - -class MockReconstructionPolicy : public ReconstructionPolicyInterface { - public: - MOCK_METHOD2(ListenAndMaybeReconstruct, - void(const ObjectID &object_id, const rpc::Address &owner_address)); - MOCK_METHOD1(Cancel, void(const ObjectID &object_id)); -}; - -class TaskDependencyManagerTest : public ::testing::Test { - public: - TaskDependencyManagerTest() - : object_manager_mock_(), - reconstruction_policy_mock_(), - task_dependency_manager_(object_manager_mock_, reconstruction_policy_mock_) {} - - protected: - MockObjectManager object_manager_mock_; - MockReconstructionPolicy reconstruction_policy_mock_; - TaskDependencyManager task_dependency_manager_; -}; - -static inline Task ExampleTask(const std::vector &arguments, - uint64_t num_returns) { - TaskSpecBuilder builder; - rpc::Address address; - builder.SetCommonTaskSpec(RandomTaskId(), "example_task", Language::PYTHON, - FunctionDescriptorBuilder::BuildPython("", "", "", ""), - JobID::Nil(), RandomTaskId(), 0, RandomTaskId(), address, - num_returns, {}, {}, - std::make_pair(PlacementGroupID::Nil(), -1), true, ""); - builder.SetActorCreationTaskSpec(ActorID::Nil(), 1, 1, {}, 1, false, "", false); - for (const auto &arg : arguments) { - builder.AddArg(TaskArgByReference(arg, rpc::Address())); - } - rpc::TaskExecutionSpec execution_spec_message; - execution_spec_message.set_num_forwards(1); - return Task(builder.Build(), TaskExecutionSpecification(execution_spec_message)); -} - -std::vector MakeTaskChain(int chain_size, - const std::vector &initial_arguments, - int64_t num_returns) { - std::vector task_chain; - std::vector arguments = initial_arguments; - for (int i = 0; i < chain_size; i++) { - auto task = ExampleTask(arguments, num_returns); - task_chain.push_back(task); - arguments.clear(); - for (size_t j = 0; j < task.GetTaskSpecification().NumReturns(); j++) { - arguments.push_back(task.GetTaskSpecification().ReturnId(j)); - } - } - return task_chain; -} - -TEST_F(TaskDependencyManagerTest, TestSimpleTask) { - // Create a task with 3 arguments. - int num_arguments = 3; - std::vector arguments; - for (int i = 0; i < num_arguments; i++) { - arguments.push_back(ObjectID::FromRandom()); - } - TaskID task_id = RandomTaskId(); - // No objects have been registered in the task dependency manager, so all - // arguments should be remote. - for (const auto &argument_id : arguments) { - EXPECT_CALL(object_manager_mock_, Pull(argument_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); - } - // Subscribe to the task's dependencies. - bool ready = task_dependency_manager_.SubscribeGetDependencies( - task_id, ObjectIdsToRefs(arguments)); - ASSERT_FALSE(ready); - - // All arguments should be canceled as they become available locally. - for (const auto &argument_id : arguments) { - EXPECT_CALL(object_manager_mock_, CancelPull(argument_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); - } - // For each argument except the last, tell the task dependency manager that - // the argument is local. - int i = 0; - for (; i < num_arguments - 1; i++) { - auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(arguments[i]); - ASSERT_TRUE(ready_task_ids.empty()); - } - // Tell the task dependency manager that the last argument is local. Now the - // task should be ready to run. - auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(arguments[i]); - ASSERT_EQ(ready_task_ids.size(), 1); - ASSERT_EQ(ready_task_ids.front(), task_id); -} - -TEST_F(TaskDependencyManagerTest, TestDuplicateSubscribeGetDependencies) { - // Create a task with 3 arguments. - TaskID task_id = RandomTaskId(); - int num_arguments = 3; - std::vector arguments; - for (int i = 0; i < num_arguments; i++) { - // Add the new argument to the list of dependencies to subscribe to. - ObjectID argument_id = ObjectID::FromRandom(); - arguments.push_back(argument_id); - // Subscribe to the task's dependencies. All arguments except the last are - // duplicates of previous subscription calls. Each argument should only be - // requested from the node manager once. - EXPECT_CALL(object_manager_mock_, Pull(argument_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); - bool ready = task_dependency_manager_.SubscribeGetDependencies( - task_id, ObjectIdsToRefs(arguments)); - ASSERT_FALSE(ready); - } - - // All arguments should be canceled as they become available locally. - for (const auto &argument_id : arguments) { - EXPECT_CALL(object_manager_mock_, CancelPull(argument_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); - } - // For each argument except the last, tell the task dependency manager that - // the argument is local. - int i = 0; - for (; i < num_arguments - 1; i++) { - auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(arguments[i]); - ASSERT_TRUE(ready_task_ids.empty()); - } - // Tell the task dependency manager that the last argument is local. Now the - // task should be ready to run. - auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(arguments[i]); - ASSERT_EQ(ready_task_ids.size(), 1); - ASSERT_EQ(ready_task_ids.front(), task_id); -} - -TEST_F(TaskDependencyManagerTest, TestMultipleTasks) { - // Create 3 tasks that are dependent on the same object. - ObjectID argument_id = ObjectID::FromRandom(); - std::vector dependent_tasks; - int num_dependent_tasks = 3; - // The object should only be requested from the object manager once for all - // three tasks. - EXPECT_CALL(object_manager_mock_, Pull(argument_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); - for (int i = 0; i < num_dependent_tasks; i++) { - TaskID task_id = RandomTaskId(); - dependent_tasks.push_back(task_id); - // Subscribe to each of the task's dependencies. - bool ready = task_dependency_manager_.SubscribeGetDependencies( - task_id, ObjectIdsToRefs({argument_id})); - ASSERT_FALSE(ready); - } - - // Tell the task dependency manager that the object is local. - EXPECT_CALL(object_manager_mock_, CancelPull(argument_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); - auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(argument_id); - // Check that all tasks are now ready to run. - ASSERT_EQ(ready_task_ids.size(), dependent_tasks.size()); - for (const auto &task_id : ready_task_ids) { - ASSERT_NE(std::find(dependent_tasks.begin(), dependent_tasks.end(), task_id), - dependent_tasks.end()); - } -} - -TEST_F(TaskDependencyManagerTest, TestTaskChain) { - // Create 3 tasks, each dependent on the previous. The first task has no - // arguments. - int num_tasks = 3; - auto tasks = MakeTaskChain(num_tasks, {}, 1); - int num_ready_tasks = 1; - int i = 0; - // No objects should be remote or canceled since each task depends on a - // locally queued task. - EXPECT_CALL(object_manager_mock_, Pull(_, _)).Times(0); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(_, _)).Times(0); - EXPECT_CALL(object_manager_mock_, CancelPull(_)).Times(0); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(_)).Times(0); - for (const auto &task : tasks) { - // Subscribe to each of the tasks' arguments. - const auto &arguments = task.GetDependencies(); - bool ready = task_dependency_manager_.SubscribeGetDependencies( - task.GetTaskSpecification().TaskId(), arguments); - if (i < num_ready_tasks) { - // The first task should be ready to run since it has no arguments. - ASSERT_TRUE(ready); - } else { - // All remaining tasks depend on the previous task. - ASSERT_FALSE(ready); - } - - // Mark each task as pending. - task_dependency_manager_.TaskPending(task); - - i++; - } - - // Simulate executing each task. Each task's completion should make the next - // task runnable. - while (!tasks.empty()) { - auto task = tasks.front(); - tasks.erase(tasks.begin()); - TaskID task_id = task.GetTaskSpecification().TaskId(); - auto return_id = task.GetTaskSpecification().ReturnId(0); - - task_dependency_manager_.UnsubscribeGetDependencies(task_id); - // Simulate the object notifications for the task's return values. - auto ready_tasks = task_dependency_manager_.HandleObjectLocal(return_id); - if (tasks.empty()) { - // If there are no more tasks, then there should be no more tasks that - // become ready to run. - ASSERT_TRUE(ready_tasks.empty()); - } else { - // If there are more tasks to run, then the next task in the chain should - // now be ready to run. - ASSERT_EQ(ready_tasks.size(), 1); - ASSERT_EQ(ready_tasks.front(), tasks.front().GetTaskSpecification().TaskId()); - } - // Simulate the task finishing execution. - task_dependency_manager_.TaskCanceled(task_id); - } -} - -TEST_F(TaskDependencyManagerTest, TestDependentPut) { - // Create a task with 3 arguments. - auto task1 = ExampleTask({}, 0); - ObjectID put_id = - ObjectID::FromIndex(task1.GetTaskSpecification().TaskId(), /*index=*/1); - auto task2 = ExampleTask({put_id}, 0); - - // No objects have been registered in the task dependency manager, so the put - // object should be remote. - EXPECT_CALL(object_manager_mock_, Pull(put_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(put_id, _)); - // Subscribe to the task's dependencies. - bool ready = task_dependency_manager_.SubscribeGetDependencies( - task2.GetTaskSpecification().TaskId(), ObjectIdsToRefs({put_id})); - ASSERT_FALSE(ready); - - // The put object should be considered local as soon as the task that creates - // it is pending execution. - EXPECT_CALL(object_manager_mock_, CancelPull(put_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(put_id)); - task_dependency_manager_.TaskPending(task1); -} - -TEST_F(TaskDependencyManagerTest, TestTaskForwarding) { - // Create 2 tasks, one dependent on the other. The first has no arguments. - int num_tasks = 2; - auto tasks = MakeTaskChain(num_tasks, {}, 1); - for (const auto &task : tasks) { - // Subscribe to each of the tasks' arguments. - const auto &arguments = task.GetDependencies(); - static_cast(task_dependency_manager_.SubscribeGetDependencies( - task.GetTaskSpecification().TaskId(), arguments)); - task_dependency_manager_.TaskPending(task); - } - - // Get the first task. - const auto task = tasks.front(); - TaskID task_id = task.GetTaskSpecification().TaskId(); - ObjectID return_id = task.GetTaskSpecification().ReturnId(0); - // Simulate forwarding the first task to a remote node. - task_dependency_manager_.UnsubscribeGetDependencies(task_id); - // The object returned by the first task should be considered remote once we - // cancel the forwarded task, since the second task depends on it. - EXPECT_CALL(object_manager_mock_, Pull(return_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(return_id, _)); - task_dependency_manager_.TaskCanceled(task_id); - - // Simulate the task executing on a remote node and its return value - // appearing locally. - EXPECT_CALL(object_manager_mock_, CancelPull(return_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(return_id)); - auto ready_tasks = task_dependency_manager_.HandleObjectLocal(return_id); - // Check that the task that we kept is now ready to run. - ASSERT_EQ(ready_tasks.size(), 1); - ASSERT_EQ(ready_tasks.front(), tasks.back().GetTaskSpecification().TaskId()); -} - -TEST_F(TaskDependencyManagerTest, TestEviction) { - // Create a task with 3 arguments. - int num_arguments = 3; - std::vector arguments; - for (int i = 0; i < num_arguments; i++) { - arguments.push_back(ObjectID::FromRandom()); - } - TaskID task_id = RandomTaskId(); - // No objects have been registered in the task dependency manager, so all - // arguments should be remote. - for (const auto &argument_id : arguments) { - EXPECT_CALL(object_manager_mock_, Pull(argument_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); - } - // Subscribe to the task's dependencies. - bool ready = task_dependency_manager_.SubscribeGetDependencies( - task_id, ObjectIdsToRefs(arguments)); - ASSERT_FALSE(ready); - - // Tell the task dependency manager that each of the arguments is now - // available. - for (const auto &argument_id : arguments) { - EXPECT_CALL(object_manager_mock_, CancelPull(argument_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); - } - for (size_t i = 0; i < arguments.size(); i++) { - std::vector ready_tasks; - ready_tasks = task_dependency_manager_.HandleObjectLocal(arguments[i]); - if (i == arguments.size() - 1) { - ASSERT_EQ(ready_tasks.size(), 1); - ASSERT_EQ(ready_tasks.front(), task_id); - } else { - ASSERT_TRUE(ready_tasks.empty()); - } - } - - // Simulate each of the arguments getting evicted. Each object should now be - // considered remote. - for (const auto &argument_id : arguments) { - EXPECT_CALL(object_manager_mock_, Pull(argument_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(argument_id, _)); - } - for (size_t i = 0; i < arguments.size(); i++) { - std::vector waiting_tasks; - waiting_tasks = task_dependency_manager_.HandleObjectMissing(arguments[i]); - if (i == 0) { - // The first eviction should cause the task to go back to the waiting - // state. - ASSERT_EQ(waiting_tasks.size(), 1); - ASSERT_EQ(waiting_tasks.front(), task_id); - } else { - // The subsequent evictions shouldn't cause any more tasks to go back to - // the waiting state. - ASSERT_TRUE(waiting_tasks.empty()); - } - } - - // Tell the task dependency manager that each of the arguments is available - // again. - for (const auto &argument_id : arguments) { - EXPECT_CALL(object_manager_mock_, CancelPull(argument_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(argument_id)); - } - for (size_t i = 0; i < arguments.size(); i++) { - std::vector ready_tasks; - ready_tasks = task_dependency_manager_.HandleObjectLocal(arguments[i]); - if (i == arguments.size() - 1) { - ASSERT_EQ(ready_tasks.size(), 1); - ASSERT_EQ(ready_tasks.front(), task_id); - } else { - ASSERT_TRUE(ready_tasks.empty()); - } - } -} - -TEST_F(TaskDependencyManagerTest, TestRemoveTasksAndRelatedObjects) { - // Create 3 tasks, each dependent on the previous. The first task has no - // arguments. - int num_tasks = 3; - auto tasks = MakeTaskChain(num_tasks, {}, 1); - // No objects should be remote or canceled since each task depends on a - // locally queued task. - EXPECT_CALL(object_manager_mock_, Pull(_, _)).Times(0); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(_, _)).Times(0); - EXPECT_CALL(object_manager_mock_, CancelPull(_)).Times(0); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(_)).Times(0); - for (const auto &task : tasks) { - // Subscribe to each of the tasks' arguments. - const auto &arguments = task.GetDependencies(); - task_dependency_manager_.SubscribeGetDependencies( - task.GetTaskSpecification().TaskId(), arguments); - // Mark each task as pending. - task_dependency_manager_.TaskPending(task); - } - - // Simulate executing the first task. This should make the second task - // runnable. - auto task = tasks.front(); - TaskID task_id = task.GetTaskSpecification().TaskId(); - auto return_id = task.GetTaskSpecification().ReturnId(0); - task_dependency_manager_.UnsubscribeGetDependencies(task_id); - // Simulate the object notifications for the task's return values. - auto ready_tasks = task_dependency_manager_.HandleObjectLocal(return_id); - // The second task should be ready to run. - ASSERT_EQ(ready_tasks.size(), 1); - // Simulate the task finishing execution. - task_dependency_manager_.TaskCanceled(task_id); - - // Remove all tasks from the manager except the first task, which already - // finished executing. - std::unordered_set task_ids; - for (const auto &task : tasks) { - task_ids.insert(task.GetTaskSpecification().TaskId()); - } - task_ids.erase(task_id); - task_dependency_manager_.RemoveTasksAndRelatedObjects(task_ids); - // Simulate evicting the return value of the first task. Make sure that this - // does not return the second task, which should have been removed. - auto waiting_tasks = task_dependency_manager_.HandleObjectMissing(return_id); - ASSERT_TRUE(waiting_tasks.empty()); - - // Simulate the object notifications for the second task's return values. - // Make sure that this does not return the third task, which should have been - // removed. - return_id = tasks[1].GetTaskSpecification().ReturnId(0); - ready_tasks = task_dependency_manager_.HandleObjectLocal(return_id); - ASSERT_TRUE(ready_tasks.empty()); -} - -/// Test that when no objects are locally available, a `ray.wait` call makes -/// the correct requests to remote nodes and correctly cancels the requests -/// when the `ray.wait` call is canceled. -TEST_F(TaskDependencyManagerTest, TestWaitDependencies) { - // Generate a random worker and objects to wait on. - WorkerID worker_id = WorkerID::FromRandom(); - int num_objects = 3; - std::vector wait_object_ids; - for (int i = 0; i < num_objects; i++) { - wait_object_ids.push_back(ObjectID::FromRandom()); - } - // Simulate a worker calling `ray.wait` on some objects. - EXPECT_CALL(object_manager_mock_, Pull(_, _)).Times(num_objects); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(_, _)) - .Times(num_objects); - task_dependency_manager_.SubscribeWaitDependencies(worker_id, - ObjectIdsToRefs(wait_object_ids)); - // Check that it's okay to call `ray.wait` on the same objects again. No new - // calls should be made to try and make the objects local. - task_dependency_manager_.SubscribeWaitDependencies(worker_id, - ObjectIdsToRefs(wait_object_ids)); - // Cancel the worker's `ray.wait`. calls. - EXPECT_CALL(object_manager_mock_, CancelPull(_)).Times(num_objects); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(_)).Times(num_objects); - task_dependency_manager_.UnsubscribeWaitDependencies(worker_id); -} - -/// Test that when one of the objects is already local at the time of the -/// `ray.wait` call, the `ray.wait` call does not trigger any requests to -/// remote nodes for that object. -TEST_F(TaskDependencyManagerTest, TestWaitDependenciesObjectLocal) { - // Generate a random worker and objects to wait on. - WorkerID worker_id = WorkerID::FromRandom(); - int num_objects = 3; - std::vector wait_object_ids; - for (int i = 0; i < num_objects; i++) { - wait_object_ids.push_back(ObjectID::FromRandom()); - } - // Simulate one of the objects becoming local. The later `ray.wait` call - // should have no effect because the object is already local. - const ObjectID local_object_id = std::move(wait_object_ids.back()); - auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(local_object_id); - ASSERT_TRUE(ready_task_ids.empty()); - - // Simulate a worker calling `ray.wait` on the objects. It should only make - // requests for the objects that are not local. - for (const auto &object_id : wait_object_ids) { - if (object_id != local_object_id) { - EXPECT_CALL(object_manager_mock_, Pull(object_id, _)); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(object_id, _)); - } - } - task_dependency_manager_.SubscribeWaitDependencies(worker_id, - ObjectIdsToRefs(wait_object_ids)); - // Simulate the local object getting evicted. The `ray.wait` call should not - // be reactivated. - auto waiting_task_ids = task_dependency_manager_.HandleObjectMissing(local_object_id); - ASSERT_TRUE(waiting_task_ids.empty()); - // Simulate a worker calling `ray.wait` on the objects. It should only make - // requests for the objects that are not local. - for (const auto &object_id : wait_object_ids) { - if (object_id != local_object_id) { - EXPECT_CALL(object_manager_mock_, CancelPull(object_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(object_id)); - } - } - task_dependency_manager_.UnsubscribeWaitDependencies(worker_id); -} - -/// Test that when one of the objects becomes local after a `ray.wait` call, -/// all requests to remote nodes associated with the object are canceled. -TEST_F(TaskDependencyManagerTest, TestWaitDependenciesHandleObjectLocal) { - // Generate a random worker and objects to wait on. - WorkerID worker_id = WorkerID::FromRandom(); - int num_objects = 3; - std::vector wait_object_ids; - for (int i = 0; i < num_objects; i++) { - wait_object_ids.push_back(ObjectID::FromRandom()); - } - // Simulate a worker calling `ray.wait` on some objects. - EXPECT_CALL(object_manager_mock_, Pull(_, _)).Times(num_objects); - EXPECT_CALL(reconstruction_policy_mock_, ListenAndMaybeReconstruct(_, _)) - .Times(num_objects); - task_dependency_manager_.SubscribeWaitDependencies(worker_id, - ObjectIdsToRefs(wait_object_ids)); - // Simulate one of the objects becoming local while the `ray.wait` calls is - // active. The `ray.wait` call should be canceled. - const ObjectID local_object_id = std::move(wait_object_ids.back()); - wait_object_ids.pop_back(); - EXPECT_CALL(object_manager_mock_, CancelPull(local_object_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(local_object_id)); - auto ready_task_ids = task_dependency_manager_.HandleObjectLocal(local_object_id); - ASSERT_TRUE(ready_task_ids.empty()); - // Simulate the local object getting evicted. The `ray.wait` call should not - // be reactivated. - auto waiting_task_ids = task_dependency_manager_.HandleObjectMissing(local_object_id); - ASSERT_TRUE(waiting_task_ids.empty()); - // Cancel the worker's `ray.wait` calls. Only the objects that are still not - // local should be canceled. - for (const auto &object_id : wait_object_ids) { - EXPECT_CALL(object_manager_mock_, CancelPull(object_id)); - EXPECT_CALL(reconstruction_policy_mock_, Cancel(object_id)); - } - task_dependency_manager_.UnsubscribeWaitDependencies(worker_id); -} - -} // namespace raylet - -} // namespace ray - -int main(int argc, char **argv) { - ::testing::InitGoogleTest(&argc, argv); - return RUN_ALL_TESTS(); -}