mirror of
https://github.com/wassname/ray.git
synced 2026-07-23 13:10:11 +08:00
Add gcs server task info handler (#6695)
This commit is contained in:
@@ -3,6 +3,7 @@
|
||||
#include "job_info_handler_impl.h"
|
||||
#include "node_info_handler_impl.h"
|
||||
#include "object_info_handler_impl.h"
|
||||
#include "task_info_handler_impl.h"
|
||||
|
||||
namespace ray {
|
||||
namespace gcs {
|
||||
@@ -38,6 +39,11 @@ void GcsServer::Start() {
|
||||
new rpc::ObjectInfoGrpcService(main_service_, *object_info_handler_));
|
||||
rpc_server_.RegisterService(*object_info_service_);
|
||||
|
||||
task_info_handler_ = InitTaskInfoHandler();
|
||||
task_info_service_.reset(
|
||||
new rpc::TaskInfoGrpcService(main_service_, *task_info_handler_));
|
||||
rpc_server_.RegisterService(*task_info_service_);
|
||||
|
||||
// Run rpc server.
|
||||
rpc_server_.Run();
|
||||
|
||||
@@ -84,5 +90,10 @@ std::unique_ptr<rpc::ObjectInfoHandler> GcsServer::InitObjectInfoHandler() {
|
||||
new rpc::DefaultObjectInfoHandler(*redis_gcs_client_));
|
||||
}
|
||||
|
||||
std::unique_ptr<rpc::TaskInfoHandler> GcsServer::InitTaskInfoHandler() {
|
||||
return std::unique_ptr<rpc::DefaultTaskInfoHandler>(
|
||||
new rpc::DefaultTaskInfoHandler(*redis_gcs_client_));
|
||||
}
|
||||
|
||||
} // namespace gcs
|
||||
} // namespace ray
|
||||
|
||||
@@ -3,8 +3,6 @@
|
||||
|
||||
#include <ray/gcs/redis_gcs_client.h>
|
||||
#include <ray/rpc/gcs_server/gcs_rpc_server.h>
|
||||
#include <ray/rpc/grpc_server.h>
|
||||
#include <string>
|
||||
|
||||
namespace ray {
|
||||
namespace gcs {
|
||||
@@ -58,6 +56,9 @@ class GcsServer {
|
||||
/// The object info handler
|
||||
virtual std::unique_ptr<rpc::ObjectInfoHandler> InitObjectInfoHandler();
|
||||
|
||||
/// The task info handler
|
||||
virtual std::unique_ptr<rpc::TaskInfoHandler> InitTaskInfoHandler();
|
||||
|
||||
private:
|
||||
/// Gcs server configuration
|
||||
GcsServerConfig config_;
|
||||
@@ -77,6 +78,9 @@ class GcsServer {
|
||||
/// Object info handler and service
|
||||
std::unique_ptr<rpc::ObjectInfoHandler> object_info_handler_;
|
||||
std::unique_ptr<rpc::ObjectInfoGrpcService> object_info_service_;
|
||||
/// Task info handler and service
|
||||
std::unique_ptr<rpc::TaskInfoHandler> task_info_handler_;
|
||||
std::unique_ptr<rpc::TaskInfoGrpcService> task_info_service_;
|
||||
/// Backend client
|
||||
std::shared_ptr<RedisGcsClient> redis_gcs_client_;
|
||||
};
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
#include "task_info_handler_impl.h"
|
||||
|
||||
namespace ray {
|
||||
namespace rpc {
|
||||
|
||||
void DefaultTaskInfoHandler::HandleAddTask(const AddTaskRequest &request,
|
||||
AddTaskReply *reply,
|
||||
SendReplyCallback send_reply_callback) {
|
||||
JobID job_id = JobID::FromBinary(request.task_data().task().task_spec().job_id());
|
||||
TaskID task_id = TaskID::FromBinary(request.task_data().task().task_spec().task_id());
|
||||
RAY_LOG(DEBUG) << "Adding task, task id = " << task_id << ", job id = " << job_id;
|
||||
auto task_table_data = std::make_shared<TaskTableData>();
|
||||
task_table_data->CopyFrom(request.task_data());
|
||||
auto on_done = [job_id, task_id, request, send_reply_callback](Status status) {
|
||||
if (!status.ok()) {
|
||||
RAY_LOG(ERROR) << "Failed to add task, task id = " << task_id
|
||||
<< ", job id = " << job_id;
|
||||
}
|
||||
send_reply_callback(status, nullptr, nullptr);
|
||||
};
|
||||
|
||||
Status status = gcs_client_.Tasks().AsyncAdd(task_table_data, on_done);
|
||||
if (!status.ok()) {
|
||||
on_done(status);
|
||||
}
|
||||
RAY_LOG(DEBUG) << "Finished adding task, task id = " << task_id
|
||||
<< ", job id = " << job_id;
|
||||
}
|
||||
|
||||
void DefaultTaskInfoHandler::HandleGetTask(const GetTaskRequest &request,
|
||||
GetTaskReply *reply,
|
||||
SendReplyCallback send_reply_callback) {
|
||||
TaskID task_id = TaskID::FromBinary(request.task_id());
|
||||
RAY_LOG(DEBUG) << "Getting task, task id = " << task_id;
|
||||
auto on_done = [task_id, request, reply, send_reply_callback](
|
||||
Status status, const boost::optional<TaskTableData> &result) {
|
||||
if (status.ok()) {
|
||||
RAY_DCHECK(result);
|
||||
reply->mutable_task_data()->CopyFrom(*result);
|
||||
} else {
|
||||
RAY_LOG(ERROR) << "Failed to get task, task id = " << task_id;
|
||||
}
|
||||
send_reply_callback(status, nullptr, nullptr);
|
||||
};
|
||||
|
||||
Status status = gcs_client_.Tasks().AsyncGet(task_id, on_done);
|
||||
if (!status.ok()) {
|
||||
on_done(status, boost::none);
|
||||
}
|
||||
RAY_LOG(DEBUG) << "Finished getting task, task id = " << task_id;
|
||||
}
|
||||
|
||||
void DefaultTaskInfoHandler::HandleDeleteTasks(const DeleteTasksRequest &request,
|
||||
DeleteTasksReply *reply,
|
||||
SendReplyCallback send_reply_callback) {
|
||||
std::vector<TaskID> task_ids = IdVectorFromProtobuf<TaskID>(request.task_id_list());
|
||||
RAY_LOG(DEBUG) << "Deleting tasks, task id list size = " << task_ids.size();
|
||||
auto on_done = [task_ids, request, send_reply_callback](Status status) {
|
||||
if (!status.ok()) {
|
||||
RAY_LOG(ERROR) << "Failed to delete tasks, task id list size = " << task_ids.size();
|
||||
}
|
||||
send_reply_callback(status, nullptr, nullptr);
|
||||
};
|
||||
|
||||
Status status = gcs_client_.Tasks().AsyncDelete(task_ids, on_done);
|
||||
if (!status.ok()) {
|
||||
on_done(status);
|
||||
}
|
||||
RAY_LOG(DEBUG) << "Finished deleting tasks, task id list size = " << task_ids.size();
|
||||
}
|
||||
|
||||
} // namespace rpc
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,32 @@
|
||||
#ifndef RAY_GCS_TASK_INFO_HANDLER_IMPL_H
|
||||
#define RAY_GCS_TASK_INFO_HANDLER_IMPL_H
|
||||
|
||||
#include "ray/gcs/redis_gcs_client.h"
|
||||
#include "ray/rpc/gcs_server/gcs_rpc_server.h"
|
||||
|
||||
namespace ray {
|
||||
namespace rpc {
|
||||
|
||||
/// This implementation class of `TaskInfoHandler`.
|
||||
class DefaultTaskInfoHandler : public rpc::TaskInfoHandler {
|
||||
public:
|
||||
explicit DefaultTaskInfoHandler(gcs::RedisGcsClient &gcs_client)
|
||||
: gcs_client_(gcs_client) {}
|
||||
|
||||
void HandleAddTask(const AddTaskRequest &request, AddTaskReply *reply,
|
||||
SendReplyCallback send_reply_callback) override;
|
||||
|
||||
void HandleGetTask(const GetTaskRequest &request, GetTaskReply *reply,
|
||||
SendReplyCallback send_reply_callback) override;
|
||||
|
||||
void HandleDeleteTasks(const DeleteTasksRequest &request, DeleteTasksReply *reply,
|
||||
SendReplyCallback send_reply_callback) override;
|
||||
|
||||
private:
|
||||
gcs::RedisGcsClient &gcs_client_;
|
||||
};
|
||||
|
||||
} // namespace rpc
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_GCS_TASK_INFO_HANDLER_IMPL_H
|
||||
@@ -1,7 +1,5 @@
|
||||
#include "gtest/gtest.h"
|
||||
#include "ray/gcs/gcs_server/actor_info_handler_impl.h"
|
||||
#include "ray/gcs/gcs_server/gcs_server.h"
|
||||
#include "ray/gcs/gcs_server/job_info_handler_impl.h"
|
||||
#include "ray/rpc/gcs_server/gcs_rpc_client.h"
|
||||
#include "ray/util/test_util.h"
|
||||
|
||||
@@ -289,6 +287,43 @@ class GcsServerTest : public RedisServiceManagerForTest {
|
||||
return object_locations;
|
||||
}
|
||||
|
||||
bool AddTask(const rpc::AddTaskRequest &request) {
|
||||
std::promise<bool> promise;
|
||||
client_->AddTask(request,
|
||||
[&promise](const Status &status, const rpc::AddTaskReply &reply) {
|
||||
RAY_CHECK_OK(status);
|
||||
promise.set_value(true);
|
||||
});
|
||||
return WaitReady(promise.get_future(), timeout_ms_);
|
||||
}
|
||||
|
||||
rpc::TaskTableData GetTask(const std::string &task_id) {
|
||||
rpc::TaskTableData task_data;
|
||||
rpc::GetTaskRequest request;
|
||||
request.set_task_id(task_id);
|
||||
std::promise<bool> promise;
|
||||
client_->GetTask(request, [&task_data, &promise](const Status &status,
|
||||
const rpc::GetTaskReply &reply) {
|
||||
if (status.ok()) {
|
||||
task_data.CopyFrom(reply.task_data());
|
||||
}
|
||||
promise.set_value(true);
|
||||
});
|
||||
|
||||
EXPECT_TRUE(WaitReady(promise.get_future(), timeout_ms_));
|
||||
return task_data;
|
||||
}
|
||||
|
||||
bool DeleteTasks(const rpc::DeleteTasksRequest &request) {
|
||||
std::promise<bool> promise;
|
||||
client_->DeleteTasks(
|
||||
request, [&promise](const Status &status, const rpc::DeleteTasksReply &reply) {
|
||||
RAY_CHECK_OK(status);
|
||||
promise.set_value(true);
|
||||
});
|
||||
return WaitReady(promise.get_future(), timeout_ms_);
|
||||
}
|
||||
|
||||
bool WaitReady(const std::future<bool> &future, uint64_t timeout_ms) {
|
||||
auto status = future.wait_for(std::chrono::milliseconds(timeout_ms));
|
||||
return status == std::future_status::ready;
|
||||
@@ -323,6 +358,18 @@ class GcsServerTest : public RedisServiceManagerForTest {
|
||||
return gcs_node_info;
|
||||
}
|
||||
|
||||
rpc::TaskTableData GenTaskTableData(const std::string &job_id,
|
||||
const std::string &task_id) {
|
||||
rpc::TaskTableData task_table_data;
|
||||
rpc::Task task;
|
||||
rpc::TaskSpec task_spec;
|
||||
task_spec.set_job_id(job_id);
|
||||
task_spec.set_task_id(task_id);
|
||||
task.mutable_task_spec()->CopyFrom(task_spec);
|
||||
task_table_data.mutable_task()->CopyFrom(task);
|
||||
return task_table_data;
|
||||
}
|
||||
|
||||
protected:
|
||||
// Gcs server
|
||||
std::unique_ptr<gcs::GcsServer> gcs_server_;
|
||||
@@ -480,6 +527,27 @@ TEST_F(GcsServerTest, TestObjectInfo) {
|
||||
ASSERT_TRUE(object_locations[0].manager() == node2_id.Binary());
|
||||
}
|
||||
|
||||
TEST_F(GcsServerTest, TestTaskInfo) {
|
||||
// Create task_table_data
|
||||
JobID job_id = JobID::FromInt(1);
|
||||
TaskID task_id = TaskID::ForDriverTask(job_id);
|
||||
rpc::TaskTableData job_table_data = GenTaskTableData(job_id.Binary(), task_id.Binary());
|
||||
|
||||
// Add task
|
||||
rpc::AddTaskRequest add_task_request;
|
||||
add_task_request.mutable_task_data()->CopyFrom(job_table_data);
|
||||
ASSERT_TRUE(AddTask(add_task_request));
|
||||
rpc::TaskTableData result = GetTask(task_id.Binary());
|
||||
ASSERT_TRUE(result.task().task_spec().job_id() == job_id.Binary());
|
||||
|
||||
// Delete task
|
||||
rpc::DeleteTasksRequest delete_tasks_request;
|
||||
delete_tasks_request.add_task_id_list(task_id.Binary());
|
||||
ASSERT_TRUE(DeleteTasks(delete_tasks_request));
|
||||
result = GetTask(task_id.Binary());
|
||||
ASSERT_TRUE(!result.has_task());
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
|
||||
int main(int argc, char **argv) {
|
||||
|
||||
@@ -286,6 +286,9 @@ Status RedisTaskInfoAccessor::AsyncDelete(const std::vector<TaskID> &task_ids,
|
||||
const StatusCallback &callback) {
|
||||
raylet::TaskTable &task_table = client_impl_->raylet_task_table();
|
||||
task_table.Delete(JobID::Nil(), task_ids);
|
||||
if (callback) {
|
||||
callback(Status::OK());
|
||||
}
|
||||
// TODO(micafan) Always return OK here.
|
||||
// Confirm if we need to handle the deletion failure and how to handle it.
|
||||
return Status::OK();
|
||||
|
||||
@@ -218,3 +218,35 @@ service ObjectInfoGcsService {
|
||||
rpc RemoveObjectLocation(RemoveObjectLocationRequest)
|
||||
returns (RemoveObjectLocationReply);
|
||||
}
|
||||
|
||||
message AddTaskRequest {
|
||||
TaskTableData task_data = 1;
|
||||
}
|
||||
|
||||
message AddTaskReply {
|
||||
}
|
||||
|
||||
message GetTaskRequest {
|
||||
bytes task_id = 1;
|
||||
}
|
||||
|
||||
message GetTaskReply {
|
||||
TaskTableData task_data = 1;
|
||||
}
|
||||
|
||||
message DeleteTasksRequest {
|
||||
repeated bytes task_id_list = 1;
|
||||
}
|
||||
|
||||
message DeleteTasksReply {
|
||||
}
|
||||
|
||||
// Service for task info access.
|
||||
service TaskInfoGcsService {
|
||||
// Add a task to GCS Service.
|
||||
rpc AddTask(AddTaskRequest) returns (AddTaskReply);
|
||||
// Get task information from GCS Service.
|
||||
rpc GetTask(GetTaskRequest) returns (GetTaskReply);
|
||||
// Delete tasks from GCS Service.
|
||||
rpc DeleteTasks(DeleteTasksRequest) returns (DeleteTasksReply);
|
||||
}
|
||||
|
||||
@@ -1,11 +1,6 @@
|
||||
#ifndef RAY_RPC_GCS_RPC_CLIENT_H
|
||||
#define RAY_RPC_GCS_RPC_CLIENT_H
|
||||
|
||||
#include <thread>
|
||||
|
||||
#include <grpcpp/grpcpp.h>
|
||||
|
||||
#include "src/ray/protobuf/gcs_service.pb.h"
|
||||
#include "src/ray/rpc/grpc_client.h"
|
||||
|
||||
namespace ray {
|
||||
@@ -20,8 +15,7 @@ class GcsRpcClient {
|
||||
/// \param[in] port Port of the gcs server.
|
||||
/// \param[in] client_call_manager The `ClientCallManager` used for managing requests.
|
||||
GcsRpcClient(const std::string &address, const int port,
|
||||
ClientCallManager &client_call_manager)
|
||||
: client_call_manager_(client_call_manager) {
|
||||
ClientCallManager &client_call_manager) {
|
||||
job_info_grpc_client_ = std::unique_ptr<GrpcClient<JobInfoGcsService>>(
|
||||
new GrpcClient<JobInfoGcsService>(address, port, client_call_manager));
|
||||
actor_info_grpc_client_ = std::unique_ptr<GrpcClient<ActorInfoGcsService>>(
|
||||
@@ -30,6 +24,8 @@ class GcsRpcClient {
|
||||
new GrpcClient<NodeInfoGcsService>(address, port, client_call_manager));
|
||||
object_info_grpc_client_ = std::unique_ptr<GrpcClient<ObjectInfoGcsService>>(
|
||||
new GrpcClient<ObjectInfoGcsService>(address, port, client_call_manager));
|
||||
task_info_grpc_client_ = std::unique_ptr<GrpcClient<TaskInfoGcsService>>(
|
||||
new GrpcClient<TaskInfoGcsService>(address, port, client_call_manager));
|
||||
};
|
||||
|
||||
/// Add job info to gcs server.
|
||||
@@ -108,15 +104,25 @@ class GcsRpcClient {
|
||||
VOID_RPC_CLIENT_METHOD(ObjectInfoGcsService, RemoveObjectLocation, request, callback,
|
||||
object_info_grpc_client_, )
|
||||
|
||||
/// Add a task to GCS Service.
|
||||
VOID_RPC_CLIENT_METHOD(TaskInfoGcsService, AddTask, request, callback,
|
||||
task_info_grpc_client_)
|
||||
|
||||
/// Get task information from GCS Service.
|
||||
VOID_RPC_CLIENT_METHOD(TaskInfoGcsService, GetTask, request, callback,
|
||||
task_info_grpc_client_)
|
||||
|
||||
/// Delete tasks from GCS Service.
|
||||
VOID_RPC_CLIENT_METHOD(TaskInfoGcsService, DeleteTasks, request, callback,
|
||||
task_info_grpc_client_)
|
||||
|
||||
private:
|
||||
/// The gRPC-generated stub.
|
||||
std::unique_ptr<GrpcClient<JobInfoGcsService>> job_info_grpc_client_;
|
||||
std::unique_ptr<GrpcClient<ActorInfoGcsService>> actor_info_grpc_client_;
|
||||
std::unique_ptr<GrpcClient<NodeInfoGcsService>> node_info_grpc_client_;
|
||||
std::unique_ptr<GrpcClient<ObjectInfoGcsService>> object_info_grpc_client_;
|
||||
|
||||
/// The `ClientCallManager` used for managing requests.
|
||||
ClientCallManager &client_call_manager_;
|
||||
std::unique_ptr<GrpcClient<TaskInfoGcsService>> task_info_grpc_client_;
|
||||
};
|
||||
|
||||
} // namespace rpc
|
||||
|
||||
@@ -45,6 +45,15 @@ namespace rpc {
|
||||
server_call_factories_and_concurrencies->emplace_back( \
|
||||
std::move(HANDLER##_call_factory), CONCURRENCY);
|
||||
|
||||
#define TASK_INFO_SERVICE_RPC_HANDLER(HANDLER, CONCURRENCY) \
|
||||
std::unique_ptr<ServerCallFactory> HANDLER##_call_factory( \
|
||||
new ServerCallFactoryImpl<TaskInfoGcsService, TaskInfoHandler, HANDLER##Request, \
|
||||
HANDLER##Reply>( \
|
||||
service_, &TaskInfoGcsService::AsyncService::Request##HANDLER, \
|
||||
service_handler_, &TaskInfoHandler::Handle##HANDLER, cq, main_service_)); \
|
||||
server_call_factories_and_concurrencies->emplace_back( \
|
||||
std::move(HANDLER##_call_factory), CONCURRENCY);
|
||||
|
||||
class JobInfoHandler {
|
||||
public:
|
||||
virtual ~JobInfoHandler() = default;
|
||||
@@ -263,6 +272,50 @@ class ObjectInfoGrpcService : public GrpcService {
|
||||
ObjectInfoHandler &service_handler_;
|
||||
};
|
||||
|
||||
class TaskInfoHandler {
|
||||
public:
|
||||
virtual ~TaskInfoHandler() = default;
|
||||
|
||||
virtual void HandleAddTask(const AddTaskRequest &request, AddTaskReply *reply,
|
||||
SendReplyCallback send_reply_callback) = 0;
|
||||
|
||||
virtual void HandleGetTask(const GetTaskRequest &request, GetTaskReply *reply,
|
||||
SendReplyCallback send_reply_callback) = 0;
|
||||
|
||||
virtual void HandleDeleteTasks(const DeleteTasksRequest &request,
|
||||
DeleteTasksReply *reply,
|
||||
SendReplyCallback send_reply_callback) = 0;
|
||||
};
|
||||
|
||||
/// The `GrpcService` for `TaskInfoGcsService`.
|
||||
class TaskInfoGrpcService : public GrpcService {
|
||||
public:
|
||||
/// Constructor.
|
||||
///
|
||||
/// \param[in] handler The service handler that actually handle the requests.
|
||||
explicit TaskInfoGrpcService(boost::asio::io_service &io_service,
|
||||
TaskInfoHandler &handler)
|
||||
: GrpcService(io_service), service_handler_(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 {
|
||||
TASK_INFO_SERVICE_RPC_HANDLER(AddTask, 1);
|
||||
TASK_INFO_SERVICE_RPC_HANDLER(GetTask, 1);
|
||||
TASK_INFO_SERVICE_RPC_HANDLER(DeleteTasks, 1);
|
||||
}
|
||||
|
||||
private:
|
||||
/// The grpc async service object.
|
||||
TaskInfoGcsService::AsyncService service_;
|
||||
/// The service handler that actually handle the requests.
|
||||
TaskInfoHandler &service_handler_;
|
||||
};
|
||||
|
||||
} // namespace rpc
|
||||
} // namespace ray
|
||||
|
||||
|
||||
Reference in New Issue
Block a user