diff --git a/src/ray/core_worker/core_worker.cc b/src/ray/core_worker/core_worker.cc index f76539f38..a072d330d 100644 --- a/src/ray/core_worker/core_worker.cc +++ b/src/ray/core_worker/core_worker.cc @@ -106,8 +106,9 @@ CoreWorker::CoreWorker(const WorkerType worker_type, const Language language, RAY_CHECK(task_execution_callback_ != nullptr); auto execute_task = std::bind(&CoreWorker::ExecuteTask, this, std::placeholders::_1, std::placeholders::_2, std::placeholders::_3); - raylet_task_receiver_ = std::unique_ptr( - new CoreWorkerRayletTaskReceiver(raylet_client_, execute_task, exit_handler)); + raylet_task_receiver_ = + std::unique_ptr(new CoreWorkerRayletTaskReceiver( + local_raylet_client_, execute_task, exit_handler)); direct_task_receiver_ = std::unique_ptr(new CoreWorkerDirectTaskReceiver( worker_context_, task_execution_service_, execute_task, exit_handler)); @@ -124,22 +125,22 @@ CoreWorker::CoreWorker(const WorkerType worker_type, const Language language, // instead of crashing. auto grpc_client = rpc::NodeManagerWorkerClient::make( node_ip_address, node_manager_port, *client_call_manager_); - ClientID raylet_id; - raylet_client_ = std::shared_ptr(new RayletClient( + ClientID local_raylet_id; + local_raylet_client_ = std::shared_ptr(new RayletClient( std::move(grpc_client), raylet_socket, WorkerID::FromBinary(worker_context_.GetWorkerID().Binary()), (worker_type_ == ray::WorkerType::WORKER), worker_context_.GetCurrentJobID(), - language_, &raylet_id, core_worker_server_.GetPort())); + language_, &local_raylet_id, core_worker_server_.GetPort())); // Unfortunately the raylet client has to be constructed after the receivers. if (direct_task_receiver_ != nullptr) { - direct_task_receiver_->Init(*raylet_client_); + direct_task_receiver_->Init(*local_raylet_client_); } // Set our own address. - RAY_CHECK(!raylet_id.IsNil()); + RAY_CHECK(!local_raylet_id.IsNil()); rpc_address_.set_ip_address(node_ip_address); rpc_address_.set_port(core_worker_server_.GetPort()); - rpc_address_.set_raylet_id(raylet_id.Binary()); + rpc_address_.set_raylet_id(local_raylet_id.Binary()); // Set timer to periodically send heartbeats containing active object IDs to the raylet. // If the heartbeat timeout is < 0, the heartbeats are disabled. @@ -157,13 +158,13 @@ CoreWorker::CoreWorker(const WorkerType worker_type, const Language language, io_thread_ = std::thread(&CoreWorker::RunIOService, this); - plasma_store_provider_.reset( - new CoreWorkerPlasmaStoreProvider(store_socket, raylet_client_, check_signals_)); + plasma_store_provider_.reset(new CoreWorkerPlasmaStoreProvider( + store_socket, local_raylet_client_, check_signals_)); memory_store_.reset(new CoreWorkerMemoryStore( [this](const RayObject &obj, const ObjectID &obj_id) { RAY_CHECK_OK(plasma_store_provider_->Put(obj, obj_id)); }, - ref_counting_enabled ? reference_counter_ : nullptr, raylet_client_)); + ref_counting_enabled ? reference_counter_ : nullptr, local_raylet_client_)); task_manager_.reset( new TaskManager(memory_store_, [this](const TaskSpecification &spec) { @@ -203,14 +204,14 @@ CoreWorker::CoreWorker(const WorkerType worker_type, const Language language, direct_task_submitter_ = std::unique_ptr(new CoreWorkerDirectTaskSubmitter( - raylet_client_, client_factory, + local_raylet_client_, client_factory, [this](const rpc::Address &address) { auto grpc_client = rpc::NodeManagerWorkerClient::make( address.ip_address(), address.port(), *client_call_manager_); return std::shared_ptr( new RayletClient(std::move(grpc_client))); }, - memory_store_, task_manager_, + memory_store_, task_manager_, local_raylet_id, RayConfig::instance().worker_lease_timeout_milliseconds())); future_resolver_.reset(new FutureResolver(memory_store_, client_factory, io_service_)); } @@ -236,8 +237,8 @@ void CoreWorker::Shutdown() { void CoreWorker::Disconnect() { io_service_.stop(); gcs_client_->Disconnect(); - if (raylet_client_) { - RAY_IGNORE_EXPR(raylet_client_->Disconnect()); + if (local_raylet_client_) { + RAY_IGNORE_EXPR(local_raylet_client_->Disconnect()); } } @@ -273,7 +274,7 @@ void CoreWorker::ReportActiveObjectIDs() { RAY_LOG(INFO) << active_object_ids.size() << " object IDs are currently in scope."; } - if (!raylet_client_->ReportActiveObjectIDs(active_object_ids).ok()) { + if (!local_raylet_client_->ReportActiveObjectIDs(active_object_ids).ok()) { RAY_LOG(ERROR) << "Raylet connection failed. Shutting down."; Shutdown(); } @@ -613,7 +614,7 @@ Status CoreWorker::SubmitTask(const RayFunction &function, return direct_task_submitter_->SubmitTask(task_spec); } else { PinObjectReferences(task_spec, TaskTransportType::RAYLET); - return raylet_client_->SubmitTask(task_spec); + return local_raylet_client_->SubmitTask(task_spec); } } @@ -652,7 +653,7 @@ Status CoreWorker::CreateActor(const RayFunction &function, // TODO(ekl) if we moved actor creation to use direct call tasks, then we won't // need to manually resolve direct call args here. resolver_->ResolveDependencies(task_spec, [this, task_spec]() { - RAY_CHECK_OK(raylet_client_->SubmitTask(task_spec)); + RAY_CHECK_OK(local_raylet_client_->SubmitTask(task_spec)); }); return Status::OK(); } @@ -696,7 +697,7 @@ Status CoreWorker::SubmitActorTask(const ActorID &actor_id, const RayFunction &f status = direct_actor_submitter_->SubmitTask(task_spec); } else { PinObjectReferences(task_spec, TaskTransportType::RAYLET); - RAY_CHECK_OK(raylet_client_->SubmitTask(task_spec)); + RAY_CHECK_OK(local_raylet_client_->SubmitTask(task_spec)); } return status; } diff --git a/src/ray/core_worker/core_worker.h b/src/ray/core_worker/core_worker.h index 31553405d..e7c731402 100644 --- a/src/ray/core_worker/core_worker.h +++ b/src/ray/core_worker/core_worker.h @@ -89,7 +89,7 @@ class CoreWorker { WorkerContext &GetWorkerContext() { return worker_context_; } - RayletClient &GetRayletClient() { return *raylet_client_; } + RayletClient &GetRayletClient() { return *local_raylet_client_; } const TaskID &GetCurrentTaskId() const { return worker_context_.GetCurrentTaskID(); } @@ -525,7 +525,7 @@ class CoreWorker { // shared_ptr for direct calls because we can lease multiple workers through // one client, and we need to keep the connection alive until we return all // of the workers. - std::shared_ptr raylet_client_; + std::shared_ptr local_raylet_client_; // Thread that runs a boost::asio service to process IO events. std::thread io_thread_; diff --git a/src/ray/core_worker/test/direct_task_transport_test.cc b/src/ray/core_worker/test/direct_task_transport_test.cc index 04da447ad..78790cda5 100644 --- a/src/ray/core_worker/test/direct_task_transport_test.cc +++ b/src/ray/core_worker/test/direct_task_transport_test.cc @@ -234,7 +234,7 @@ TEST(DirectTaskTransportTest, TestSubmitOneTask) { auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store, - task_finisher, kLongTimeout); + task_finisher, ClientID::Nil(), kLongTimeout); std::unordered_map empty_resources; std::vector empty_descriptor; @@ -264,7 +264,7 @@ TEST(DirectTaskTransportTest, TestHandleTaskFailure) { auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store, - task_finisher, kLongTimeout); + task_finisher, ClientID::Nil(), kLongTimeout); std::unordered_map empty_resources; std::vector empty_descriptor; TaskSpecification task = BuildTaskSpec(empty_resources, empty_descriptor); @@ -287,7 +287,7 @@ TEST(DirectTaskTransportTest, TestConcurrentWorkerLeases) { auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store, - task_finisher, kLongTimeout); + task_finisher, ClientID::Nil(), kLongTimeout); std::unordered_map empty_resources; std::vector empty_descriptor; TaskSpecification task1 = BuildTaskSpec(empty_resources, empty_descriptor); @@ -331,7 +331,7 @@ TEST(DirectTaskTransportTest, TestReuseWorkerLease) { auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store, - task_finisher, kLongTimeout); + task_finisher, ClientID::Nil(), kLongTimeout); std::unordered_map empty_resources; std::vector empty_descriptor; TaskSpecification task1 = BuildTaskSpec(empty_resources, empty_descriptor); @@ -378,7 +378,7 @@ TEST(DirectTaskTransportTest, TestWorkerNotReusedOnError) { auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store, - task_finisher, kLongTimeout); + task_finisher, ClientID::Nil(), kLongTimeout); std::unordered_map empty_resources; std::vector empty_descriptor; TaskSpecification task1 = BuildTaskSpec(empty_resources, empty_descriptor); @@ -425,7 +425,8 @@ TEST(DirectTaskTransportTest, TestSpillback) { }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, lease_client_factory, - store, task_finisher, kLongTimeout); + store, task_finisher, ClientID::Nil(), + kLongTimeout); std::unordered_map empty_resources; std::vector empty_descriptor; TaskSpecification task = BuildTaskSpec(empty_resources, empty_descriptor); @@ -456,6 +457,61 @@ TEST(DirectTaskTransportTest, TestSpillback) { ASSERT_EQ(task_finisher->num_tasks_failed, 0); } +TEST(DirectTaskTransportTest, TestSpillbackRoundTrip) { + auto raylet_client = std::make_shared(); + auto worker_client = std::make_shared(); + auto store = std::make_shared(); + auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; + + std::unordered_map> remote_lease_clients; + auto lease_client_factory = [&](const rpc::Address &addr) { + ClientID raylet_id = ClientID::FromBinary(addr.raylet_id()); + // We should not create a connection to the same raylet more than once. + RAY_CHECK(remote_lease_clients.count(raylet_id) == 0); + auto client = std::make_shared(); + remote_lease_clients[raylet_id] = client; + return client; + }; + auto task_finisher = std::make_shared(); + auto local_raylet_id = ClientID::FromRandom(); + CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, lease_client_factory, + store, task_finisher, local_raylet_id, + kLongTimeout); + std::unordered_map empty_resources; + std::vector empty_descriptor; + TaskSpecification task = BuildTaskSpec(empty_resources, empty_descriptor); + + ASSERT_TRUE(submitter.SubmitTask(task).ok()); + ASSERT_EQ(raylet_client->num_workers_requested, 1); + ASSERT_EQ(raylet_client->num_workers_returned, 0); + ASSERT_EQ(worker_client->callbacks.size(), 0); + ASSERT_EQ(remote_lease_clients.size(), 0); + + // Spillback to a remote node. + auto remote_raylet_id = ClientID::FromRandom(); + ASSERT_TRUE(raylet_client->GrantWorkerLease("localhost", 1234, remote_raylet_id)); + ASSERT_EQ(remote_lease_clients.count(remote_raylet_id), 1); + ASSERT_FALSE(raylet_client->GrantWorkerLease("remote", 1234, ClientID::Nil())); + // Trigger a spillback back to the local node. + ASSERT_TRUE(remote_lease_clients[remote_raylet_id]->GrantWorkerLease("local", 1234, + local_raylet_id)); + // We should not have created another lease client to the local raylet. + ASSERT_EQ(remote_lease_clients.size(), 1); + // There should be no more callbacks on the remote node. + ASSERT_FALSE(remote_lease_clients[remote_raylet_id]->GrantWorkerLease("remote", 1234, + ClientID::Nil())); + + // The worker is returned to the local node. + ASSERT_TRUE(raylet_client->GrantWorkerLease("local", 1234, ClientID::Nil())); + ASSERT_TRUE(worker_client->ReplyPushTask()); + ASSERT_EQ(raylet_client->num_workers_returned, 1); + ASSERT_EQ(remote_lease_clients[remote_raylet_id]->num_workers_returned, 0); + ASSERT_EQ(raylet_client->num_workers_disconnected, 0); + ASSERT_EQ(remote_lease_clients[remote_raylet_id]->num_workers_disconnected, 0); + ASSERT_EQ(task_finisher->num_tasks_complete, 1); + ASSERT_EQ(task_finisher->num_tasks_failed, 0); +} + // Helper to run a test that checks that 'same1' and 'same2' are treated as the same // resource shape, while 'different' is treated as a separate shape. void TestSchedulingKey(const std::shared_ptr store, @@ -466,7 +522,7 @@ void TestSchedulingKey(const std::shared_ptr store, auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store, - task_finisher, kLongTimeout); + task_finisher, ClientID::Nil(), kLongTimeout); ASSERT_TRUE(submitter.SubmitTask(same1).ok()); ASSERT_TRUE(submitter.SubmitTask(same2).ok()); @@ -565,7 +621,7 @@ TEST(DirectTaskTransportTest, TestWorkerLeaseTimeout) { auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; }; auto task_finisher = std::make_shared(); CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store, - task_finisher, + task_finisher, ClientID::Nil(), /*lease_timeout_ms=*/5); std::unordered_map empty_resources; std::vector empty_descriptor; diff --git a/src/ray/core_worker/transport/direct_task_transport.cc b/src/ray/core_worker/transport/direct_task_transport.cc index 171c9bbb6..54760be6c 100644 --- a/src/ray/core_worker/transport/direct_task_transport.cc +++ b/src/ray/core_worker/transport/direct_task_transport.cc @@ -64,8 +64,9 @@ std::shared_ptr CoreWorkerDirectTaskSubmitter::GetOrConnectLeaseClient( const rpc::Address *raylet_address) { std::shared_ptr lease_client; - if (raylet_address) { - // Connect to raylet. + if (raylet_address && + ClientID::FromBinary(raylet_address->raylet_id()) != local_raylet_id_) { + // A remote raylet was specified. Connect to the raylet if needed. ClientID raylet_id = ClientID::FromBinary(raylet_address->raylet_id()); auto it = remote_lease_clients_.find(raylet_id); if (it == remote_lease_clients_.end()) { diff --git a/src/ray/core_worker/transport/direct_task_transport.h b/src/ray/core_worker/transport/direct_task_transport.h index 6650b188f..d8bdcc23a 100644 --- a/src/ray/core_worker/transport/direct_task_transport.h +++ b/src/ray/core_worker/transport/direct_task_transport.h @@ -34,12 +34,13 @@ class CoreWorkerDirectTaskSubmitter { LeaseClientFactoryFn lease_client_factory, std::shared_ptr store, std::shared_ptr task_finisher, - int64_t lease_timeout_ms) + ClientID local_raylet_id, int64_t lease_timeout_ms) : local_lease_client_(lease_client), client_factory_(client_factory), lease_client_factory_(lease_client_factory), resolver_(store), task_finisher_(task_finisher), + local_raylet_id_(local_raylet_id), lease_timeout_ms_(lease_timeout_ms) {} /// Schedule a task for direct submission to a worker. @@ -102,6 +103,10 @@ class CoreWorkerDirectTaskSubmitter { /// to the raylet. int64_t lease_timeout_ms_; + /// The local raylet ID. Used to make sure that we use the local lease client + /// if a remote raylet tells us to spill the task back to the local raylet. + const ClientID local_raylet_id_; + // Protects task submission state below. absl::Mutex mu_;