mirror of
https://github.com/wassname/ray.git
synced 2026-08-19 12:30:27 +08:00
[grpc] Add grpc server to worker (#5054)
* refactor grpc server * format * change GetTask() to PushTask() * change PushTask to AssignTask * format * update * fix test * format * Update src/ray/rpc/worker_client.h Co-Authored-By: Hao Chen <chenh1024@gmail.com> * Update BUILD.bazel * Update src/ray/core_worker/task_execution.cc Co-Authored-By: Stephanie Wang <swang@cs.berkeley.edu> * update * format * address comments * format * Update src/ray/rpc/worker/worker_server.h Co-Authored-By: Stephanie Wang <swang@cs.berkeley.edu> * Update src/ray/protobuf/worker.proto Co-Authored-By: Stephanie Wang <swang@cs.berkeley.edu> * format * fix * format
This commit is contained in:
committed by
Qing Wang
co-authored by
Hao Chen
Stephanie Wang
parent
41a16c55ef
commit
54d5969cea
+36
@@ -47,6 +47,16 @@ cc_proto_library(
|
||||
deps = ["object_manager_proto"],
|
||||
)
|
||||
|
||||
proto_library(
|
||||
name = "worker_proto",
|
||||
srcs = ["src/ray/protobuf/worker.proto"],
|
||||
)
|
||||
|
||||
cc_proto_library(
|
||||
name = "worker_cc_proto",
|
||||
deps = ["worker_proto"],
|
||||
)
|
||||
|
||||
proto_library(
|
||||
name = "core_worker_proto",
|
||||
srcs = ["src/ray/protobuf/core_worker.proto"],
|
||||
@@ -127,6 +137,30 @@ cc_library(
|
||||
],
|
||||
)
|
||||
|
||||
# worker gRPC lib.
|
||||
cc_grpc_library(
|
||||
name = "worker_cc_grpc",
|
||||
srcs = [":worker_proto"],
|
||||
grpc_only = True,
|
||||
deps = [":worker_cc_proto"],
|
||||
)
|
||||
|
||||
# worker server and client.
|
||||
cc_library(
|
||||
name = "worker_rpc",
|
||||
hdrs = glob([
|
||||
"src/ray/rpc/worker/*.h",
|
||||
]),
|
||||
copts = COPTS,
|
||||
deps = [
|
||||
":grpc_common_lib",
|
||||
":ray_common",
|
||||
":worker_cc_grpc",
|
||||
"@boost//:asio",
|
||||
"@com_github_grpc_grpc//:grpc++",
|
||||
],
|
||||
)
|
||||
|
||||
# === End of rpc definitions ===
|
||||
|
||||
cc_binary(
|
||||
@@ -216,6 +250,7 @@ cc_library(
|
||||
":ray_common",
|
||||
":ray_util",
|
||||
":stats_lib",
|
||||
":worker_rpc",
|
||||
"@boost//:asio",
|
||||
"@com_github_jupp0r_prometheus_cpp//pull",
|
||||
"@com_google_absl//absl/base:core_headers",
|
||||
@@ -253,6 +288,7 @@ cc_library(
|
||||
":ray_common",
|
||||
":ray_util",
|
||||
":raylet_lib",
|
||||
":worker_rpc",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@@ -10,15 +10,23 @@ CoreWorker::CoreWorker(const enum WorkerType worker_type, const ::Language langu
|
||||
language_(language),
|
||||
raylet_socket_(raylet_socket),
|
||||
worker_context_(worker_type, job_id),
|
||||
// TODO(zhijunfu): currently RayletClient would crash in its constructor
|
||||
// if it cannot connect to Raylet after a number of retries, this needs
|
||||
// to be changed so that the worker (java/python .etc) can retrieve and
|
||||
// handle the error instead of crashing.
|
||||
raylet_client_(raylet_socket_,
|
||||
ClientID::FromBinary(worker_context_.GetWorkerID().Binary()),
|
||||
(worker_type_ == ray::WorkerType::WORKER),
|
||||
worker_context_.GetCurrentJobID(), language_),
|
||||
task_interface_(worker_context_, raylet_client_),
|
||||
object_interface_(worker_context_, raylet_client_, store_socket),
|
||||
task_execution_interface_(worker_context_, raylet_client_, object_interface_) {}
|
||||
object_interface_(worker_context_, raylet_client_, store_socket) {
|
||||
int rpc_server_port = 0;
|
||||
if (worker_type_ == ray::WorkerType::WORKER) {
|
||||
task_execution_interface_ = std::unique_ptr<CoreWorkerTaskExecutionInterface>(
|
||||
new CoreWorkerTaskExecutionInterface(worker_context_, raylet_client_,
|
||||
object_interface_));
|
||||
rpc_server_port = task_execution_interface_->worker_server_.GetPort();
|
||||
}
|
||||
// TODO(zhijunfu): currently RayletClient would crash in its constructor if it cannot
|
||||
// connect to Raylet after a number of retries, this can be changed later
|
||||
// so that the worker (java/python .etc) can retrieve and handle the error
|
||||
// instead of crashing.
|
||||
raylet_client_ = std::unique_ptr<RayletClient>(new RayletClient(
|
||||
raylet_socket_, ClientID::FromBinary(worker_context_.GetWorkerID().Binary()),
|
||||
(worker_type_ == ray::WorkerType::WORKER), worker_context_.GetCurrentJobID(),
|
||||
language_, rpc_server_port));
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
|
||||
@@ -43,7 +43,10 @@ class CoreWorker {
|
||||
|
||||
/// Return the `CoreWorkerTaskExecutionInterface` that contains methods related to
|
||||
/// task execution.
|
||||
CoreWorkerTaskExecutionInterface &Execution() { return task_execution_interface_; }
|
||||
CoreWorkerTaskExecutionInterface &Execution() {
|
||||
RAY_CHECK(task_execution_interface_ != nullptr);
|
||||
return *task_execution_interface_;
|
||||
}
|
||||
|
||||
private:
|
||||
/// Type of this worker.
|
||||
@@ -59,7 +62,7 @@ class CoreWorker {
|
||||
WorkerContext worker_context_;
|
||||
|
||||
/// Raylet client.
|
||||
RayletClient raylet_client_;
|
||||
std::unique_ptr<RayletClient> raylet_client_;
|
||||
|
||||
/// The `CoreWorkerTaskInterface` instance.
|
||||
CoreWorkerTaskInterface task_interface_;
|
||||
@@ -68,7 +71,8 @@ class CoreWorker {
|
||||
CoreWorkerObjectInterface object_interface_;
|
||||
|
||||
/// The `CoreWorkerTaskExecutionInterface` instance.
|
||||
CoreWorkerTaskExecutionInterface task_execution_interface_;
|
||||
/// This is only available if it's not a driver.
|
||||
std::unique_ptr<CoreWorkerTaskExecutionInterface> task_execution_interface_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
@@ -4,9 +4,9 @@
|
||||
|
||||
namespace ray {
|
||||
|
||||
CoreWorkerObjectInterface::CoreWorkerObjectInterface(WorkerContext &worker_context,
|
||||
RayletClient &raylet_client,
|
||||
const std::string &store_socket)
|
||||
CoreWorkerObjectInterface::CoreWorkerObjectInterface(
|
||||
WorkerContext &worker_context, std::unique_ptr<RayletClient> &raylet_client,
|
||||
const std::string &store_socket)
|
||||
: worker_context_(worker_context), raylet_client_(raylet_client) {
|
||||
store_providers_.emplace(
|
||||
static_cast<int>(StoreProviderType::PLASMA),
|
||||
|
||||
@@ -17,7 +17,8 @@ class CoreWorkerStoreProvider;
|
||||
/// The interface that contains all `CoreWorker` methods that are related to object store.
|
||||
class CoreWorkerObjectInterface {
|
||||
public:
|
||||
CoreWorkerObjectInterface(WorkerContext &worker_context, RayletClient &raylet_client,
|
||||
CoreWorkerObjectInterface(WorkerContext &worker_context,
|
||||
std::unique_ptr<RayletClient> &raylet_client,
|
||||
const std::string &store_socket);
|
||||
|
||||
/// Put an object into object store.
|
||||
@@ -59,7 +60,8 @@ class CoreWorkerObjectInterface {
|
||||
/// \param[in] local_only Whether only delete the objects in local node, or all nodes in
|
||||
/// the cluster.
|
||||
/// \param[in] delete_creating_tasks Whether also delete the tasks that
|
||||
/// created these objects. \return Status.
|
||||
/// created these objects.
|
||||
/// \return Status.
|
||||
Status Delete(const std::vector<ObjectID> &object_ids, bool local_only,
|
||||
bool delete_creating_tasks);
|
||||
|
||||
@@ -67,7 +69,7 @@ class CoreWorkerObjectInterface {
|
||||
/// Reference to the parent CoreWorker's context.
|
||||
WorkerContext &worker_context_;
|
||||
/// Reference to the parent CoreWorker's raylet client.
|
||||
RayletClient &raylet_client_;
|
||||
std::unique_ptr<RayletClient> &raylet_client_;
|
||||
|
||||
/// All the store providers supported.
|
||||
std::unordered_map<int, std::unique_ptr<CoreWorkerStoreProvider>> store_providers_;
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
namespace ray {
|
||||
|
||||
CoreWorkerPlasmaStoreProvider::CoreWorkerPlasmaStoreProvider(
|
||||
const std::string &store_socket, RayletClient &raylet_client)
|
||||
const std::string &store_socket, std::unique_ptr<RayletClient> &raylet_client)
|
||||
: raylet_client_(raylet_client) {
|
||||
auto status = store_client_.Connect(store_socket);
|
||||
if (!status.ok()) {
|
||||
@@ -70,7 +70,7 @@ Status CoreWorkerPlasmaStoreProvider::Get(
|
||||
}
|
||||
|
||||
// TODO(zhijunfu): can call `fetchOrReconstruct` in batches as an optimization.
|
||||
RAY_CHECK_OK(raylet_client_.FetchOrReconstruct(unready_ids, fetch_only, task_id));
|
||||
RAY_CHECK_OK(raylet_client_->FetchOrReconstruct(unready_ids, fetch_only, task_id));
|
||||
|
||||
// Get the objects from the object store, and parse the result.
|
||||
int64_t get_timeout;
|
||||
@@ -109,7 +109,7 @@ Status CoreWorkerPlasmaStoreProvider::Get(
|
||||
}
|
||||
|
||||
if (was_blocked) {
|
||||
RAY_CHECK_OK(raylet_client_.NotifyUnblocked(task_id));
|
||||
RAY_CHECK_OK(raylet_client_->NotifyUnblocked(task_id));
|
||||
}
|
||||
|
||||
return Status::OK();
|
||||
@@ -120,8 +120,8 @@ Status CoreWorkerPlasmaStoreProvider::Wait(const std::vector<ObjectID> &object_i
|
||||
const TaskID &task_id,
|
||||
std::vector<bool> *results) {
|
||||
WaitResultPair result_pair;
|
||||
auto status = raylet_client_.Wait(object_ids, num_objects, timeout_ms, false, task_id,
|
||||
&result_pair);
|
||||
auto status = raylet_client_->Wait(object_ids, num_objects, timeout_ms, false, task_id,
|
||||
&result_pair);
|
||||
std::unordered_set<ObjectID> ready_ids;
|
||||
for (const auto &entry : result_pair.first) {
|
||||
ready_ids.insert(entry);
|
||||
@@ -141,7 +141,7 @@ Status CoreWorkerPlasmaStoreProvider::Wait(const std::vector<ObjectID> &object_i
|
||||
Status CoreWorkerPlasmaStoreProvider::Delete(const std::vector<ObjectID> &object_ids,
|
||||
bool local_only,
|
||||
bool delete_creating_tasks) {
|
||||
return raylet_client_.FreeObjects(object_ids, local_only, delete_creating_tasks);
|
||||
return raylet_client_->FreeObjects(object_ids, local_only, delete_creating_tasks);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
|
||||
@@ -18,7 +18,7 @@ class CoreWorker;
|
||||
class CoreWorkerPlasmaStoreProvider : public CoreWorkerStoreProvider {
|
||||
public:
|
||||
CoreWorkerPlasmaStoreProvider(const std::string &store_socket,
|
||||
RayletClient &raylet_client);
|
||||
std::unique_ptr<RayletClient> &raylet_client);
|
||||
|
||||
/// Put an object with specified ID into object store.
|
||||
///
|
||||
@@ -67,7 +67,7 @@ class CoreWorkerPlasmaStoreProvider : public CoreWorkerStoreProvider {
|
||||
std::mutex store_client_mutex_;
|
||||
|
||||
/// Raylet client.
|
||||
RayletClient &raylet_client_;
|
||||
std::unique_ptr<RayletClient> &raylet_client_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
@@ -77,7 +77,8 @@ class CoreWorkerStoreProvider {
|
||||
/// \param[in] local_only Whether only delete the objects in local node, or all nodes in
|
||||
/// the cluster.
|
||||
/// \param[in] delete_creating_tasks Whether also delete the tasks that
|
||||
/// created these objects. \return Status.
|
||||
/// created these objects.
|
||||
/// \return Status.
|
||||
virtual Status Delete(const std::vector<ObjectID> &object_ids, bool local_only,
|
||||
bool delete_creating_tasks) = 0;
|
||||
};
|
||||
|
||||
@@ -6,19 +6,26 @@
|
||||
namespace ray {
|
||||
|
||||
CoreWorkerTaskExecutionInterface::CoreWorkerTaskExecutionInterface(
|
||||
WorkerContext &worker_context, RayletClient &raylet_client,
|
||||
WorkerContext &worker_context, std::unique_ptr<RayletClient> &raylet_client,
|
||||
CoreWorkerObjectInterface &object_interface)
|
||||
: worker_context_(worker_context), object_interface_(object_interface) {
|
||||
task_receivers.emplace(static_cast<int>(TaskTransportType::RAYLET),
|
||||
std::unique_ptr<CoreWorkerRayletTaskReceiver>(
|
||||
new CoreWorkerRayletTaskReceiver(raylet_client)));
|
||||
: worker_context_(worker_context),
|
||||
object_interface_(object_interface),
|
||||
worker_server_("Worker", 0 /* let grpc choose port */),
|
||||
main_work_(main_service_) {
|
||||
task_receivers_.emplace(
|
||||
static_cast<int>(TaskTransportType::RAYLET),
|
||||
std::unique_ptr<CoreWorkerRayletTaskReceiver>(new CoreWorkerRayletTaskReceiver(
|
||||
raylet_client, main_service_, worker_server_)));
|
||||
|
||||
// Start RPC server after all the task receivers are properly initialized.
|
||||
worker_server_.Run();
|
||||
}
|
||||
|
||||
Status CoreWorkerTaskExecutionInterface::Run(const TaskExecutor &executor) {
|
||||
while (true) {
|
||||
std::vector<TaskSpec> tasks;
|
||||
auto status =
|
||||
task_receivers[static_cast<int>(TaskTransportType::RAYLET)]->GetTasks(&tasks);
|
||||
task_receivers_[static_cast<int>(TaskTransportType::RAYLET)]->GetTasks(&tasks);
|
||||
if (!status.ok()) {
|
||||
RAY_LOG(ERROR) << "Getting task failed with error: "
|
||||
<< ray::Status::IOError(status.message());
|
||||
|
||||
@@ -7,6 +7,9 @@
|
||||
#include "ray/core_worker/context.h"
|
||||
#include "ray/core_worker/object_interface.h"
|
||||
#include "ray/core_worker/transport/transport.h"
|
||||
#include "ray/rpc/client_call.h"
|
||||
#include "ray/rpc/worker/worker_client.h"
|
||||
#include "ray/rpc/worker/worker_server.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
@@ -21,7 +24,7 @@ class TaskSpecification;
|
||||
class CoreWorkerTaskExecutionInterface {
|
||||
public:
|
||||
CoreWorkerTaskExecutionInterface(WorkerContext &worker_context,
|
||||
RayletClient &raylet_client,
|
||||
std::unique_ptr<RayletClient> &raylet_client,
|
||||
CoreWorkerObjectInterface &object_interface);
|
||||
|
||||
/// The callback provided app-language workers that executes tasks.
|
||||
@@ -56,7 +59,18 @@ class CoreWorkerTaskExecutionInterface {
|
||||
CoreWorkerObjectInterface &object_interface_;
|
||||
|
||||
/// All the task task receivers supported.
|
||||
std::unordered_map<int, std::unique_ptr<CoreWorkerTaskReceiver>> task_receivers;
|
||||
std::unordered_map<int, std::unique_ptr<CoreWorkerTaskReceiver>> task_receivers_;
|
||||
|
||||
/// The RPC server.
|
||||
rpc::GrpcServer worker_server_;
|
||||
|
||||
/// Event loop where tasks are processed.
|
||||
boost::asio::io_service main_service_;
|
||||
|
||||
/// The asio work to keep main_service_ alive.
|
||||
boost::asio::io_service::work main_work_;
|
||||
|
||||
friend class CoreWorker;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
@@ -92,8 +92,8 @@ std::vector<ray::ActorHandleID> ActorHandle::NewActorHandles() const {
|
||||
|
||||
void ActorHandle::ClearNewActorHandles() { new_actor_handles_.clear(); }
|
||||
|
||||
CoreWorkerTaskInterface::CoreWorkerTaskInterface(WorkerContext &worker_context,
|
||||
RayletClient &raylet_client)
|
||||
CoreWorkerTaskInterface::CoreWorkerTaskInterface(
|
||||
WorkerContext &worker_context, std::unique_ptr<RayletClient> &raylet_client)
|
||||
: worker_context_(worker_context) {
|
||||
task_submitters_.emplace(static_cast<int>(TaskTransportType::RAYLET),
|
||||
std::unique_ptr<CoreWorkerRayletTaskSubmitter>(
|
||||
|
||||
@@ -111,7 +111,8 @@ class ActorHandle {
|
||||
/// submission.
|
||||
class CoreWorkerTaskInterface {
|
||||
public:
|
||||
CoreWorkerTaskInterface(WorkerContext &worker_context, RayletClient &raylet_client);
|
||||
CoreWorkerTaskInterface(WorkerContext &worker_context,
|
||||
std::unique_ptr<RayletClient> &raylet_client);
|
||||
|
||||
/// Submit a normal task.
|
||||
///
|
||||
|
||||
@@ -1,21 +1,20 @@
|
||||
|
||||
#include "ray/core_worker/transport/raylet_transport.h"
|
||||
#include "ray/raylet/task.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
CoreWorkerRayletTaskSubmitter::CoreWorkerRayletTaskSubmitter(RayletClient &raylet_client)
|
||||
CoreWorkerRayletTaskSubmitter::CoreWorkerRayletTaskSubmitter(
|
||||
std::unique_ptr<RayletClient> &raylet_client)
|
||||
: raylet_client_(raylet_client) {}
|
||||
|
||||
Status CoreWorkerRayletTaskSubmitter::SubmitTask(const TaskSpec &task) {
|
||||
return raylet_client_.SubmitTask(task.GetDependencies(), task.GetTaskSpecification());
|
||||
return raylet_client_->SubmitTask(task.GetDependencies(), task.GetTaskSpecification());
|
||||
}
|
||||
|
||||
CoreWorkerRayletTaskReceiver::CoreWorkerRayletTaskReceiver(RayletClient &raylet_client)
|
||||
: raylet_client_(raylet_client) {}
|
||||
|
||||
Status CoreWorkerRayletTaskReceiver::GetTasks(std::vector<TaskSpec> *tasks) {
|
||||
std::unique_ptr<raylet::TaskSpecification> task_spec;
|
||||
auto status = raylet_client_.GetTask(&task_spec);
|
||||
auto status = raylet_client_->GetTask(&task_spec);
|
||||
if (!status.ok()) {
|
||||
RAY_LOG(ERROR) << "Get task from raylet failed with error: "
|
||||
<< ray::Status::IOError(status.message());
|
||||
@@ -29,4 +28,23 @@ Status CoreWorkerRayletTaskReceiver::GetTasks(std::vector<TaskSpec> *tasks) {
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
CoreWorkerRayletTaskReceiver::CoreWorkerRayletTaskReceiver(
|
||||
std::unique_ptr<RayletClient> &raylet_client, boost::asio::io_service &io_service,
|
||||
rpc::GrpcServer &server)
|
||||
: raylet_client_(raylet_client), task_service_(io_service, *this) {
|
||||
server.RegisterService(task_service_);
|
||||
}
|
||||
|
||||
void CoreWorkerRayletTaskReceiver::HandleAssignTask(
|
||||
const rpc::AssignTaskRequest &request, rpc::AssignTaskReply *reply,
|
||||
rpc::RequestDoneCallback done_callback) {
|
||||
const std::string &task_message = request.task_spec();
|
||||
const raylet::Task task(*flatbuffers::GetRoot<protocol::Task>(
|
||||
reinterpret_cast<const uint8_t *>(task_message.data())));
|
||||
const auto &spec = task.GetTaskSpecification();
|
||||
|
||||
auto status = task_handler_(spec);
|
||||
done_callback(status);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
|
||||
#include "ray/core_worker/transport/transport.h"
|
||||
#include "ray/raylet/raylet_client.h"
|
||||
#include "ray/rpc/worker/worker_server.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
@@ -14,7 +15,7 @@ namespace ray {
|
||||
|
||||
class CoreWorkerRayletTaskSubmitter : public CoreWorkerTaskSubmitter {
|
||||
public:
|
||||
CoreWorkerRayletTaskSubmitter(RayletClient &raylet_client);
|
||||
CoreWorkerRayletTaskSubmitter(std::unique_ptr<RayletClient> &raylet_client);
|
||||
|
||||
/// Submit a task for execution to raylet.
|
||||
///
|
||||
@@ -24,19 +25,41 @@ class CoreWorkerRayletTaskSubmitter : public CoreWorkerTaskSubmitter {
|
||||
|
||||
private:
|
||||
/// Raylet client.
|
||||
RayletClient &raylet_client_;
|
||||
std::unique_ptr<RayletClient> &raylet_client_;
|
||||
};
|
||||
|
||||
class CoreWorkerRayletTaskReceiver : public CoreWorkerTaskReceiver {
|
||||
class CoreWorkerRayletTaskReceiver : public CoreWorkerTaskReceiver,
|
||||
public rpc::WorkerTaskHandler {
|
||||
public:
|
||||
CoreWorkerRayletTaskReceiver(RayletClient &raylet_client);
|
||||
CoreWorkerRayletTaskReceiver(std::unique_ptr<RayletClient> &raylet_client,
|
||||
boost::asio::io_service &io_service,
|
||||
rpc::GrpcServer &server);
|
||||
|
||||
// Get tasks for execution from raylet.
|
||||
virtual Status GetTasks(std::vector<TaskSpec> *tasks) override;
|
||||
|
||||
/// TODO(zhijunfu): This is currently unused. Later when we migrate from worker "get
|
||||
/// task" to raylet "assign task", this method will be used and the `GetTask` above will
|
||||
/// be removed.
|
||||
///
|
||||
/// Handle a `AssignTask` request.
|
||||
/// The implementation can handle this request asynchronously. When hanling is done, the
|
||||
/// `done_callback` should be called.
|
||||
///
|
||||
/// \param[in] request The request message.
|
||||
/// \param[out] reply The reply message.
|
||||
/// \param[in] done_callback The callback to be called when the request is done.
|
||||
void HandleAssignTask(const rpc::AssignTaskRequest &request,
|
||||
rpc::AssignTaskReply *reply,
|
||||
rpc::RequestDoneCallback done_callback) override;
|
||||
|
||||
private:
|
||||
/// Raylet client.
|
||||
RayletClient &raylet_client_;
|
||||
std::unique_ptr<RayletClient> &raylet_client_;
|
||||
/// The callback function to process a task.
|
||||
TaskHandler task_handler_;
|
||||
/// The rpc service for `WorkerTaskService`.
|
||||
rpc::WorkerTaskGrpcService task_service_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
@@ -32,6 +32,8 @@ class CoreWorkerTaskSubmitter {
|
||||
/// This class receives tasks for execution.
|
||||
class CoreWorkerTaskReceiver {
|
||||
public:
|
||||
using TaskHandler = std::function<Status(const raylet::TaskSpecification &task_spec)>;
|
||||
|
||||
// Get tasks for execution.
|
||||
virtual Status GetTasks(std::vector<TaskSpec> *tasks) = 0;
|
||||
};
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
syntax = "proto3";
|
||||
|
||||
package ray.rpc;
|
||||
|
||||
message AssignTaskRequest {
|
||||
// The ID of the task to be pushed.
|
||||
bytes task_id = 1;
|
||||
// The task to be pushed. This should include task_id.
|
||||
// TODO(hchen): Currently, `task_spec` are represented as
|
||||
// flatbutters-serialized bytes. This is because the flatbuffers-defined Task data
|
||||
// structure is being used in many places. We should move Task and all related data
|
||||
// structures to protobuf.
|
||||
bytes task_spec = 2;
|
||||
}
|
||||
|
||||
message AssignTaskReply {
|
||||
}
|
||||
|
||||
// Service for worker.
|
||||
service WorkerTaskService {
|
||||
// Push a task to a worker.
|
||||
rpc AssignTask(AssignTaskRequest) returns (AssignTaskReply);
|
||||
}
|
||||
@@ -138,6 +138,10 @@ table RegisterClientRequest {
|
||||
job_id: string;
|
||||
// Language of this worker.
|
||||
language: Language;
|
||||
// Port that this worker is listening on.
|
||||
// If port > 0, then worker will listen to this port and wait for
|
||||
// raylet to push tasks, instead of invoking GetTask().
|
||||
port: int;
|
||||
}
|
||||
|
||||
table RegisterClientReply {
|
||||
|
||||
@@ -845,8 +845,8 @@ void NodeManager::ProcessRegisterClientRequestMessage(
|
||||
const std::shared_ptr<LocalClientConnection> &client, const uint8_t *message_data) {
|
||||
auto message = flatbuffers::GetRoot<protocol::RegisterClientRequest>(message_data);
|
||||
client->SetClientID(from_flatbuf<ClientID>(*message->worker_id()));
|
||||
auto worker =
|
||||
std::make_shared<Worker>(message->worker_pid(), message->language(), client);
|
||||
auto worker = std::make_shared<Worker>(message->worker_pid(), message->language(),
|
||||
message->port(), client);
|
||||
if (message->is_worker()) {
|
||||
// Register the new worker.
|
||||
worker_pool_.RegisterWorker(std::move(worker));
|
||||
|
||||
@@ -202,15 +202,20 @@ ray::Status RayletConnection::AtomicRequestReply(
|
||||
}
|
||||
|
||||
RayletClient::RayletClient(const std::string &raylet_socket, const ClientID &client_id,
|
||||
bool is_worker, const JobID &job_id, const Language &language)
|
||||
: client_id_(client_id), is_worker_(is_worker), job_id_(job_id), language_(language) {
|
||||
bool is_worker, const JobID &job_id, const Language &language,
|
||||
int port)
|
||||
: client_id_(client_id),
|
||||
is_worker_(is_worker),
|
||||
job_id_(job_id),
|
||||
language_(language),
|
||||
port_(port) {
|
||||
// For C++14, we could use std::make_unique
|
||||
conn_ = std::unique_ptr<RayletConnection>(new RayletConnection(raylet_socket, -1, -1));
|
||||
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
auto message = ray::protocol::CreateRegisterClientRequest(
|
||||
fbb, is_worker, to_flatbuf(fbb, client_id), getpid(), to_flatbuf(fbb, job_id),
|
||||
language);
|
||||
language, port);
|
||||
fbb.Finish(message);
|
||||
// Register the process ID with the raylet.
|
||||
// NOTE(swang): If raylet exits and we are registered as a worker, we will get killed.
|
||||
|
||||
@@ -69,7 +69,8 @@ class RayletClient {
|
||||
/// \param job_id The ID of the driver. This is non-nil if the client is a driver.
|
||||
/// \return The connection information.
|
||||
RayletClient(const std::string &raylet_socket, const ClientID &client_id,
|
||||
bool is_worker, const JobID &job_id, const Language &language);
|
||||
bool is_worker, const JobID &job_id, const Language &language,
|
||||
int port = -1);
|
||||
|
||||
ray::Status Disconnect() { return conn_->Disconnect(); };
|
||||
|
||||
@@ -188,6 +189,7 @@ class RayletClient {
|
||||
const bool is_worker_;
|
||||
const JobID job_id_;
|
||||
const Language language_;
|
||||
const int port_;
|
||||
/// A map from resource name to the resource IDs that are currently reserved
|
||||
/// for this worker. Each pair consists of the resource ID and the fraction
|
||||
/// of that resource allocated for this worker.
|
||||
|
||||
@@ -10,10 +10,11 @@ namespace ray {
|
||||
namespace raylet {
|
||||
|
||||
/// A constructor responsible for initializing the state of a worker.
|
||||
Worker::Worker(pid_t pid, const Language &language,
|
||||
Worker::Worker(pid_t pid, const Language &language, int port,
|
||||
std::shared_ptr<LocalClientConnection> connection)
|
||||
: pid_(pid),
|
||||
language_(language),
|
||||
port_(port),
|
||||
connection_(connection),
|
||||
dead_(false),
|
||||
blocked_(false) {}
|
||||
@@ -32,6 +33,8 @@ pid_t Worker::Pid() const { return pid_; }
|
||||
|
||||
Language Worker::GetLanguage() const { return language_; }
|
||||
|
||||
int Worker::Port() const { return port_; }
|
||||
|
||||
void Worker::AssignTaskId(const TaskID &task_id) { assigned_task_id_ = task_id; }
|
||||
|
||||
const TaskID &Worker::GetAssignedTaskId() const { return assigned_task_id_; }
|
||||
|
||||
@@ -17,7 +17,7 @@ namespace raylet {
|
||||
class Worker {
|
||||
public:
|
||||
/// A constructor that initializes a worker object.
|
||||
Worker(pid_t pid, const Language &language,
|
||||
Worker(pid_t pid, const Language &language, int port,
|
||||
std::shared_ptr<LocalClientConnection> connection);
|
||||
/// A destructor responsible for freeing all worker state.
|
||||
~Worker() {}
|
||||
@@ -29,6 +29,7 @@ class Worker {
|
||||
/// Return the worker's PID.
|
||||
pid_t Pid() const;
|
||||
Language GetLanguage() const;
|
||||
int Port() const;
|
||||
void AssignTaskId(const TaskID &task_id);
|
||||
const TaskID &GetAssignedTaskId() const;
|
||||
bool AddBlockedTaskId(const TaskID &task_id);
|
||||
@@ -56,6 +57,9 @@ class Worker {
|
||||
pid_t pid_;
|
||||
/// The language type of this worker.
|
||||
Language language_;
|
||||
/// Port that this worker listens on.
|
||||
/// If port <= 0, this indicates that the worker will not listen to a port.
|
||||
int port_;
|
||||
/// Connection state of a worker.
|
||||
std::shared_ptr<LocalClientConnection> connection_;
|
||||
/// The worker's currently assigned task.
|
||||
|
||||
@@ -175,7 +175,8 @@ pid_t WorkerPool::StartProcess(const std::vector<const char *> &worker_command_a
|
||||
|
||||
void WorkerPool::RegisterWorker(const std::shared_ptr<Worker> &worker) {
|
||||
const auto pid = worker->Pid();
|
||||
RAY_LOG(DEBUG) << "Registering worker with pid " << pid;
|
||||
const auto port = worker->Port();
|
||||
RAY_LOG(DEBUG) << "Registering worker with pid " << pid << ", port: " << port;
|
||||
auto &state = GetStateForLanguage(worker->GetLanguage());
|
||||
state.registered_workers.insert(std::move(worker));
|
||||
|
||||
|
||||
@@ -86,7 +86,7 @@ class WorkerPoolTest : public ::testing::Test {
|
||||
auto client =
|
||||
LocalClientConnection::Create(client_handler, message_handler, std::move(socket),
|
||||
"worker", {}, error_message_type_);
|
||||
return std::shared_ptr<Worker>(new Worker(pid, language, client));
|
||||
return std::shared_ptr<Worker>(new Worker(pid, language, -1, client));
|
||||
}
|
||||
|
||||
void SetWorkerCommands(
|
||||
|
||||
@@ -30,6 +30,8 @@ class ClientCall {
|
||||
/// The callback to be called by `ClientCallManager` when the reply of this request is
|
||||
/// received.
|
||||
virtual void OnReplyReceived() = 0;
|
||||
/// Return status.
|
||||
virtual ray::Status GetStatus() = 0;
|
||||
|
||||
virtual ~ClientCall() = default;
|
||||
};
|
||||
@@ -49,6 +51,7 @@ using ClientCallback = std::function<void(const Status &status, const Reply &rep
|
||||
template <class Reply>
|
||||
class ClientCallImpl : public ClientCall {
|
||||
public:
|
||||
Status GetStatus() override { return GrpcStatusToRayStatus(status_); }
|
||||
void OnReplyReceived() override {
|
||||
if (callback_ != nullptr) {
|
||||
callback_(GrpcStatusToRayStatus(status_), reply_);
|
||||
|
||||
@@ -13,7 +13,7 @@ void GrpcServer::Run() {
|
||||
builder.AddListeningPort(server_address, grpc::InsecureServerCredentials(), &port_);
|
||||
// Register all the services to this server.
|
||||
if (services_.size() == 0) {
|
||||
RAY_LOG(WARNING) << "No service is found when start grpc server.";
|
||||
RAY_LOG(WARNING) << "No service is found when start grpc server " << name_;
|
||||
}
|
||||
for (auto &entry : services_) {
|
||||
builder.RegisterService(&entry.get());
|
||||
@@ -34,6 +34,8 @@ void GrpcServer::Run() {
|
||||
}
|
||||
// Start a thread that polls incoming requests.
|
||||
polling_thread_ = std::thread(&GrpcServer::PollEventsFromCompletionQueue, this);
|
||||
// Set the server as running.
|
||||
is_closed_ = false;
|
||||
}
|
||||
|
||||
void GrpcServer::RegisterService(GrpcService &service) {
|
||||
|
||||
@@ -33,7 +33,7 @@ class GrpcServer {
|
||||
/// \param[in] main_service The main event loop, to which service handler functions
|
||||
/// will be posted.
|
||||
GrpcServer(const std::string &name, const uint32_t port)
|
||||
: name_(name), port_(port), is_closed_(false) {}
|
||||
: name_(name), port_(port), is_closed_(true) {}
|
||||
|
||||
/// Destruct this gRPC server.
|
||||
~GrpcServer() { Shutdown(); }
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
#ifndef RAY_RPC_WORKER_CLIENT_H
|
||||
#define RAY_RPC_WORKER_CLIENT_H
|
||||
|
||||
#include <thread>
|
||||
|
||||
#include <grpcpp/grpcpp.h>
|
||||
|
||||
#include "ray/common/status.h"
|
||||
#include "ray/rpc/client_call.h"
|
||||
#include "ray/util/logging.h"
|
||||
#include "src/ray/protobuf/worker.grpc.pb.h"
|
||||
#include "src/ray/protobuf/worker.pb.h"
|
||||
|
||||
namespace ray {
|
||||
namespace rpc {
|
||||
|
||||
/// Client used for communicating with a remote worker server.
|
||||
class WorkerTaskClient {
|
||||
public:
|
||||
/// Constructor.
|
||||
///
|
||||
/// \param[in] address Address of the worker server.
|
||||
/// \param[in] port Port of the worker server.
|
||||
/// \param[in] client_call_manager The `ClientCallManager` used for managing requests.
|
||||
WorkerTaskClient(const std::string &address, const int port,
|
||||
ClientCallManager &client_call_manager)
|
||||
: client_call_manager_(client_call_manager) {
|
||||
std::shared_ptr<grpc::Channel> channel = grpc::CreateChannel(
|
||||
address + ":" + std::to_string(port), grpc::InsecureChannelCredentials());
|
||||
stub_ = WorkerTaskService::NewStub(channel);
|
||||
};
|
||||
|
||||
/// Assign a task to the work.
|
||||
///
|
||||
/// \param[in] request The request message.
|
||||
/// \param[in] callback The callback function that handles reply.
|
||||
/// \return if the rpc call succeeds
|
||||
ray::Status AssignTask(const AssignTaskRequest &request,
|
||||
const ClientCallback<AssignTaskReply> &callback) {
|
||||
auto call = client_call_manager_
|
||||
.CreateCall<WorkerTaskService, AssignTaskRequest, AssignTaskReply>(
|
||||
*stub_, &WorkerTaskService::Stub::PrepareAsyncAssignTask, request,
|
||||
callback);
|
||||
return call->GetStatus();
|
||||
}
|
||||
|
||||
private:
|
||||
/// The gRPC-generated stub.
|
||||
std::unique_ptr<WorkerTaskService::Stub> stub_;
|
||||
|
||||
/// The `ClientCallManager` used for managing requests.
|
||||
ClientCallManager &client_call_manager_;
|
||||
};
|
||||
|
||||
} // namespace rpc
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RPC_WORKER_CLIENT_H
|
||||
@@ -0,0 +1,68 @@
|
||||
#ifndef RAY_RPC_WORKER_SERVER_H
|
||||
#define RAY_RPC_WORKER_SERVER_H
|
||||
|
||||
#include "ray/rpc/grpc_server.h"
|
||||
#include "ray/rpc/server_call.h"
|
||||
|
||||
#include "src/ray/protobuf/worker.grpc.pb.h"
|
||||
#include "src/ray/protobuf/worker.pb.h"
|
||||
|
||||
namespace ray {
|
||||
namespace rpc {
|
||||
|
||||
/// Interface of the `WorkerService`, see `src/ray/protobuf/worker.proto`.
|
||||
class WorkerTaskHandler {
|
||||
public:
|
||||
/// Handle a `AssignTask` request.
|
||||
/// The implementation can handle this request asynchronously. When handling is done,
|
||||
/// the `done_callback` should be called.
|
||||
///
|
||||
/// \param[in] request The request message.
|
||||
/// \param[out] reply The reply message.
|
||||
/// \param[in] done_callback The callback to be called when the request is done.
|
||||
virtual void HandleAssignTask(const AssignTaskRequest &request, AssignTaskReply *reply,
|
||||
RequestDoneCallback done_callback) = 0;
|
||||
};
|
||||
|
||||
/// The `GrpcServer` for `WorkerService`.
|
||||
class WorkerTaskGrpcService : public GrpcService {
|
||||
public:
|
||||
/// Constructor.
|
||||
///
|
||||
/// \param[in] main_service See super class.
|
||||
/// \param[in] handler The service handler that actually handle the requests.
|
||||
WorkerTaskGrpcService(boost::asio::io_service &main_service,
|
||||
WorkerTaskHandler &service_handler)
|
||||
: GrpcService(main_service), service_handler_(service_handler){};
|
||||
|
||||
protected:
|
||||
grpc::Service &GetGrpcService() override { return service_; }
|
||||
|
||||
void InitServerCallFactories(
|
||||
const std::unique_ptr<grpc::ServerCompletionQueue> &cq,
|
||||
std::vector<std::pair<std::unique_ptr<ServerCallFactory>, int>>
|
||||
*server_call_factories_and_concurrencies) override {
|
||||
// Initialize the Factory for `AssignTask` requests.
|
||||
std::unique_ptr<ServerCallFactory> push_task_call_Factory(
|
||||
new ServerCallFactoryImpl<WorkerTaskService, WorkerTaskHandler, AssignTaskRequest,
|
||||
AssignTaskReply>(
|
||||
service_, &WorkerTaskService::AsyncService::RequestAssignTask,
|
||||
service_handler_, &WorkerTaskHandler::HandleAssignTask, cq, main_service_));
|
||||
|
||||
// Set `AssignTask`'s accept concurrency to 5.
|
||||
server_call_factories_and_concurrencies->emplace_back(
|
||||
std::move(push_task_call_Factory), 5);
|
||||
}
|
||||
|
||||
private:
|
||||
/// The grpc async service object.
|
||||
WorkerTaskService::AsyncService service_;
|
||||
|
||||
/// The service handler that actually handle the requests.
|
||||
WorkerTaskHandler &service_handler_;
|
||||
};
|
||||
|
||||
} // namespace rpc
|
||||
} // namespace ray
|
||||
|
||||
#endif
|
||||
Reference in New Issue
Block a user