[direct task] Fix bug that starts duplicate connections from the worker to the local raylet (#6307)

* Fix bug and add unit test

* rename
This commit is contained in:
Stephanie Wang
2019-12-02 10:25:05 -08:00
committed by GitHub
parent da41180dc0
commit 69dd5c9319
5 changed files with 95 additions and 32 deletions
+20 -19
View File
@@ -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<CoreWorkerRayletTaskReceiver>(
new CoreWorkerRayletTaskReceiver(raylet_client_, execute_task, exit_handler));
raylet_task_receiver_ =
std::unique_ptr<CoreWorkerRayletTaskReceiver>(new CoreWorkerRayletTaskReceiver(
local_raylet_client_, execute_task, exit_handler));
direct_task_receiver_ =
std::unique_ptr<CoreWorkerDirectTaskReceiver>(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<RayletClient>(new RayletClient(
ClientID local_raylet_id;
local_raylet_client_ = std::shared_ptr<RayletClient>(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<CoreWorkerDirectTaskSubmitter>(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<RayletClient>(
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;
}
+2 -2
View File
@@ -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<RayletClient> raylet_client_;
std::shared_ptr<RayletClient> local_raylet_client_;
// Thread that runs a boost::asio service to process IO events.
std::thread io_thread_;
@@ -234,7 +234,7 @@ TEST(DirectTaskTransportTest, TestSubmitOneTask) {
auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; };
auto task_finisher = std::make_shared<MockTaskFinisher>();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store,
task_finisher, kLongTimeout);
task_finisher, ClientID::Nil(), kLongTimeout);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> empty_descriptor;
@@ -264,7 +264,7 @@ TEST(DirectTaskTransportTest, TestHandleTaskFailure) {
auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; };
auto task_finisher = std::make_shared<MockTaskFinisher>();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store,
task_finisher, kLongTimeout);
task_finisher, ClientID::Nil(), kLongTimeout);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> 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<MockTaskFinisher>();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store,
task_finisher, kLongTimeout);
task_finisher, ClientID::Nil(), kLongTimeout);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> 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<MockTaskFinisher>();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store,
task_finisher, kLongTimeout);
task_finisher, ClientID::Nil(), kLongTimeout);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> 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<MockTaskFinisher>();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store,
task_finisher, kLongTimeout);
task_finisher, ClientID::Nil(), kLongTimeout);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> empty_descriptor;
TaskSpecification task1 = BuildTaskSpec(empty_resources, empty_descriptor);
@@ -425,7 +425,8 @@ TEST(DirectTaskTransportTest, TestSpillback) {
};
auto task_finisher = std::make_shared<MockTaskFinisher>();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, lease_client_factory,
store, task_finisher, kLongTimeout);
store, task_finisher, ClientID::Nil(),
kLongTimeout);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> 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<MockRayletClient>();
auto worker_client = std::make_shared<MockWorkerClient>();
auto store = std::make_shared<CoreWorkerMemoryStore>();
auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; };
std::unordered_map<ClientID, std::shared_ptr<MockRayletClient>> 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<MockRayletClient>();
remote_lease_clients[raylet_id] = client;
return client;
};
auto task_finisher = std::make_shared<MockTaskFinisher>();
auto local_raylet_id = ClientID::FromRandom();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, lease_client_factory,
store, task_finisher, local_raylet_id,
kLongTimeout);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> 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<CoreWorkerMemoryStore> store,
@@ -466,7 +522,7 @@ void TestSchedulingKey(const std::shared_ptr<CoreWorkerMemoryStore> store,
auto factory = [&](const rpc::WorkerAddress &addr) { return worker_client; };
auto task_finisher = std::make_shared<MockTaskFinisher>();
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<MockTaskFinisher>();
CoreWorkerDirectTaskSubmitter submitter(raylet_client, factory, nullptr, store,
task_finisher,
task_finisher, ClientID::Nil(),
/*lease_timeout_ms=*/5);
std::unordered_map<std::string, double> empty_resources;
std::vector<std::string> empty_descriptor;
@@ -64,8 +64,9 @@ std::shared_ptr<WorkerLeaseInterface>
CoreWorkerDirectTaskSubmitter::GetOrConnectLeaseClient(
const rpc::Address *raylet_address) {
std::shared_ptr<WorkerLeaseInterface> 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()) {
@@ -34,12 +34,13 @@ class CoreWorkerDirectTaskSubmitter {
LeaseClientFactoryFn lease_client_factory,
std::shared_ptr<CoreWorkerMemoryStore> store,
std::shared_ptr<TaskFinisherInterface> 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_;