[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:
Zhijun Fu
2019-07-04 20:16:42 +08:00
committed by Qing Wang
co-authored by Hao Chen Stephanie Wang
parent 41a16c55ef
commit 54d5969cea
29 changed files with 351 additions and 62 deletions
+36
View File
@@ -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",
],
)
+18 -10
View File
@@ -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
+7 -3
View File
@@ -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
+3 -3
View File
@@ -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),
+5 -3
View File
@@ -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;
};
+13 -6
View File
@@ -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());
+16 -2
View File
@@ -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
+2 -2
View File
@@ -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>(
+2 -1
View File
@@ -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;
};
+23
View File
@@ -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);
}
+4
View File
@@ -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 {
+2 -2
View File
@@ -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));
+8 -3
View File
@@ -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.
+3 -1
View File
@@ -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.
+4 -1
View File
@@ -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_; }
+5 -1
View File
@@ -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.
+2 -1
View File
@@ -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));
+1 -1
View File
@@ -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(
+3
View File
@@ -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_);
+3 -1
View File
@@ -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) {
+1 -1
View File
@@ -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(); }
+58
View File
@@ -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
+68
View File
@@ -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