mirror of
https://github.com/wassname/ray.git
synced 2026-08-06 13:31:10 +08:00
[XRay] Raylet node and object manager unification/backend redesign. (#1640)
* directory for raylet * some initial class scaffolding -- in progress * node_manager build code and test stub files. * class scaffolding for resources, workers, and the worker pool * Node manager server loop * raylet policy and queue - wip checkpoint * fix dependencies * add gen_nm_fbs as target. * object manager build, stub, and test code. * Start integrating WorkerPool into node manager * fix build on mac * tmp * adding LsResources boilerplate * add/build Task spec boilerplate * checkpoint ActorInformation and LsQueue * Worker pool maintains started and removed workers * todos for e2e task assignment * fix build on mac * build/add lsqueue interface * channel resource config through from NodeServer to LsResources; prep LsResources to replace/provide worker_pool * progress on LsResources class: resource availability check implementation * Read task submission messages from a client * Submit tasks from the client to the local scheduler * Assign a task to a worker from the WorkerPool * change the way node_manager is built to prevent build issues for object_manager. * add namespaces. fix build. * Move ClientConnection message handling into server, remove reference to WorkerPool * Add raw constructors for TaskSpecification * Define TaskArgument by reference and by value * Flatbuffer serialization for TaskSpec * expand resource implementation * Start integrating TaskExecutionSpecification into Task * Separate WorkerPool from LsResources, give ownership to NodeServer * checkpoint queue and resource code * resoving merge conflicts * lspolicy::schedule ; adding lsqueue and lspolicy to the nodeserver * Implement LsQueue RemoveTasks and QueueReadyTasks * Fill in some LsQueue code for assigning a task * added suport for test_asio * Implement LsQueue queue tasks methods, queue running tasks * calling into policy from nodeserver; adding cluster resource map * Feedback and Testing. Incorporate Alexey's feedback. Actually test some code. Clean up callback imp. * end to end task assignment * Decouple local scheduler from node server * move TODO * Move local scheduler to separate file * Add scaffolding for reconstruction policy, task dependency manager, and object manager * fix * asio for store client notifications. added asio for plasma store connection. added tests for store notifications. encapsulate store interaction under store_messenger. * Move Worker inside of ClientConnection * Set the assigned task ID in the worker * Several changes toward object manager implementation. Store client integration with asio. Complete OM/OD scaffolding. * simple simulator to estimate number of retry timeouts * changing dbclientid --> clientid * fix build (include sandbox after it's fixed). * changes to object manager, adding lambdas to the interface * changing void * callbacks to std::function typed callbacks * remove use namespace std from .h files. use ray:: for Status everywhere. * minor * lineage cache interfaces * TODO for object IDs * Interface for the GCS client table * Revert "Set the assigned task ID in the worker" This reverts commit a770dd31048a289ef431c56d64e491fa7f9b2737. * Revert "Move Worker inside of ClientConnection" This reverts commit dfaa0d662a76976c05be6d76b214b45d88482818. * OD/OM: ray::Status * mock gcs integration. * gcs mock clientinfo assignment * Allow lookup of a Worker in the WorkerPool * Split out Worker and ClientConnection source files * Allow assignment of a task ID to a worker, skeleton for finishing a task * integrate mock gcs with om tests. * added tcp connection acceptor * integrated OM with NM. integrated GcsClient with NM. Added multi-node integration tests. * OM to receive incoming tcp connections. * implemented object manager connection protocol. * Added todos. * slight adjustment to add/remove handler invocation on object store client. * Simplify Task interface for getting dependencies * Remove unused object manager file * TaskDependencyManager tracks missing task dependencies and processes object add notifications * Local scheduler queues tasks according to argument availability * Fill in TaskSpecification methods to get arguments * Implemented push. * Queue tasks that have been scheduled but that are waiting for a worker * Pull + mock gcs cleanup. * OD/OM/GCS mock code review, fixing unused-result issues, eliminating copy ctor * Remove unique_ptr from object_store_client * Fix object manager Push memory error * Pull task arguments in task dependency manager * Add a demo script for remote task dependencies * Some comments for the TaskDependencyManager * code cleanup; builds on mac * Make ClientConnection a templated type based on the connection protocol * Add gmock to build * Add WorkerPool unit tests * clean up. * clean up connection code. * instantiate a template instance in the module * Virtual destructors * Document public api. * Separate read and write buffers in ClientConnection; documentation * Remove ObjectDirectory from NodeServer constructor, make directory InitGcs call a separate constructor * Convert NodeServer Terminate to a destructor * NodeServer documentation * WorkerPool documentation * TaskDependencyManager doc * unifying naming conventions * unifying naming conventions * Task cleanup and documentation * unifying naming conventions * unifying naming conventions * code cleanup and naming conventions * code cleanup * Rename om --> object_manager * Merge with master * SchedulingQueue doc * Docs and implementation skeleton for ClientTable * Node manager documentation * ReconstructionPolicy doc * Replace std::bind with lambda in TaskDependencyManager * lineage cache doc * Use \param style for doc * documentation for scheduling policy and resources * minor code cleanup * SchedulingResources class documentation + code cleanup * referencing ray/raylet directory; doxygen documentation * updating trivial policy * Fix bug where event loop stops after task submission * Define entry point for ClientManager for handling new connections * Node manager to node manager protocol, heartbeat protocol * Fix flatbuffer * Fix GCS flatbuffer naming conflict * client connection moved to common dir. * rename based on feedback. * Added google style and 90 char lines clang-format file under src/ray. * const ref ClientID. * Incorporated feedback from PR. * raylet: includes and namespaces * raylet/om/gcs logging/using * doxygen style * camel casing, comments, other style; DBClientID -> ClientID * object_manager : naming, defines, style * consistent caps and naming; misc style * cleaning up client connection + other stylistic fixes * cmath, std::nan * more style polish: OM, Raylet, gcs tables * removing sandbox (moved to ray-project/sandbox) * raylet linting * object manager linting * gcs linting * all other linting Co-authored-by: Melih <elibol@gmail.com> Co-authored-by: Stephanie <swang@cs.berkeley.edu>
This commit is contained in:
committed by
Stephanie Wang
co-authored by
Melih
Stephanie
parent
d85274a12e
commit
91464a56dd
@@ -20,6 +20,8 @@
|
||||
/src/plasma/format/*_generated.h
|
||||
/src/local_scheduler/format/*_generated.h
|
||||
/src/ray/gcs/format/*_generated.h
|
||||
/src/ray/object_manager/format/*_generated.h
|
||||
/src/ray/raylet/format/*_generated.h
|
||||
|
||||
# Redis temporary files
|
||||
*dump.rdb
|
||||
|
||||
@@ -19,6 +19,8 @@ if(RAY_BUILD_TESTS OR RAY_BUILD_BENCHMARKS)
|
||||
"${GTEST_PREFIX}/lib/${CMAKE_STATIC_LIBRARY_PREFIX}gtest${CMAKE_STATIC_LIBRARY_SUFFIX}")
|
||||
set(GTEST_MAIN_STATIC_LIB
|
||||
"${GTEST_PREFIX}/lib/${CMAKE_STATIC_LIBRARY_PREFIX}gtest_main${CMAKE_STATIC_LIBRARY_SUFFIX}")
|
||||
set(GMOCK_MAIN_STATIC_LIB
|
||||
"${GTEST_PREFIX}/lib/${CMAKE_STATIC_LIBRARY_PREFIX}gmock_main${CMAKE_STATIC_LIBRARY_SUFFIX}")
|
||||
set(GTEST_CMAKE_ARGS -DCMAKE_BUILD_TYPE=${CMAKE_BUILD_TYPE}
|
||||
-DCMAKE_INSTALL_PREFIX=${GTEST_PREFIX}
|
||||
-DCMAKE_CXX_FLAGS=${GTEST_CMAKE_CXX_FLAGS})
|
||||
@@ -28,7 +30,7 @@ if(RAY_BUILD_TESTS OR RAY_BUILD_BENCHMARKS)
|
||||
|
||||
ExternalProject_Add(googletest_ep
|
||||
URL "https://github.com/google/googletest/archive/release-${GTEST_VERSION}.tar.gz"
|
||||
BUILD_BYPRODUCTS ${GTEST_STATIC_LIB} ${GTEST_MAIN_STATIC_LIB}
|
||||
BUILD_BYPRODUCTS ${GTEST_STATIC_LIB} ${GTEST_MAIN_STATIC_LIB} ${GMOCK_MAIN_STATIC_LIB}
|
||||
CMAKE_ARGS ${GTEST_CMAKE_ARGS}
|
||||
${EP_LOG_OPTIONS})
|
||||
|
||||
@@ -39,9 +41,12 @@ if(RAY_BUILD_TESTS OR RAY_BUILD_BENCHMARKS)
|
||||
STATIC_LIB ${GTEST_STATIC_LIB})
|
||||
ADD_THIRDPARTY_LIB(gtest_main
|
||||
STATIC_LIB ${GTEST_MAIN_STATIC_LIB})
|
||||
ADD_THIRDPARTY_LIB(gmock_main
|
||||
STATIC_LIB ${GMOCK_MAIN_STATIC_LIB})
|
||||
|
||||
add_dependencies(gtest googletest_ep)
|
||||
add_dependencies(gtest_main googletest_ep)
|
||||
add_dependencies(gmock_main googletest_ep)
|
||||
|
||||
set(GFLAGS_CMAKE_CXX_FLAGS ${EP_CXX_FLAGS})
|
||||
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
BasedOnStyle: Google
|
||||
ColumnLimit: 90
|
||||
DerivePointerAlignment: false
|
||||
IndentCaseLabels: false
|
||||
PointerAlignment: Right
|
||||
+30
-3
@@ -5,6 +5,8 @@ include_directories(${CMAKE_CURRENT_LIST_DIR}/../common/thirdparty/ae)
|
||||
|
||||
add_subdirectory(util)
|
||||
add_subdirectory(gcs)
|
||||
add_subdirectory(object_manager)
|
||||
add_subdirectory(raylet)
|
||||
|
||||
include(${CMAKE_CURRENT_LIST_DIR}/../common/cmake/Common.cmake)
|
||||
|
||||
@@ -21,6 +23,11 @@ set(HIREDIS_SRCS
|
||||
${CMAKE_CURRENT_LIST_DIR}/../common/thirdparty/hiredis/sds.c
|
||||
)
|
||||
|
||||
set(UTIL_SRCS
|
||||
${CMAKE_CURRENT_LIST_DIR}/../common/thirdparty/sha256.c
|
||||
${CMAKE_CURRENT_LIST_DIR}/../common/common_protocol.cc
|
||||
)
|
||||
|
||||
set(RAY_SRCS
|
||||
id.cc
|
||||
status.cc
|
||||
@@ -28,6 +35,26 @@ set(RAY_SRCS
|
||||
gcs/tables.cc
|
||||
gcs/redis_context.cc
|
||||
gcs/asio.cc
|
||||
common/client_connection.cc
|
||||
object_manager/object_manager_client_connection.cc
|
||||
object_manager/object_store_client.cc
|
||||
object_manager/object_directory.cc
|
||||
object_manager/object_manager.cc
|
||||
raylet/mock_gcs_client.cc
|
||||
raylet/task.cc
|
||||
raylet/task_execution_spec.cc
|
||||
raylet/task_spec.cc
|
||||
raylet/worker.cc
|
||||
raylet/worker_pool.cc
|
||||
raylet/scheduling_resources.cc
|
||||
raylet/actor.cc
|
||||
raylet/scheduling_queue.cc
|
||||
raylet/scheduling_policy.cc
|
||||
raylet/task_dependency_manager.cc
|
||||
raylet/reconstruction_policy.cc
|
||||
raylet/node_manager.cc
|
||||
raylet/lineage_cache.cc
|
||||
raylet/raylet.cc
|
||||
)
|
||||
|
||||
install(FILES
|
||||
@@ -45,7 +72,7 @@ install(
|
||||
DESTINATION "${CMAKE_INSTALL_LIBDIR}/pkgconfig/")
|
||||
|
||||
ADD_RAY_LIB(ray
|
||||
SOURCES ${RAY_SRCS} ${AE_SRCS} ${HIREDIS_SRCS}
|
||||
DEPENDENCIES gen_gcs_fbs
|
||||
SOURCES ${RAY_SRCS} ${AE_SRCS} ${HIREDIS_SRCS} ${UTIL_SRCS}
|
||||
DEPENDENCIES gen_gcs_fbs gen_object_manager_fbs gen_node_manager_fbs
|
||||
SHARED_LINK_LIBS ""
|
||||
STATIC_LINK_LIBS "")
|
||||
STATIC_LINK_LIBS ${PLASMA_STATIC_LIB} ${ARROW_STATIC_LIB})
|
||||
|
||||
@@ -0,0 +1,108 @@
|
||||
#include "client_connection.h"
|
||||
|
||||
#include <boost/bind.hpp>
|
||||
|
||||
#include "common.h"
|
||||
#include "ray/raylet/format/node_manager_generated.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
template <class T>
|
||||
std::shared_ptr<ClientConnection<T>> ClientConnection<T>::Create(
|
||||
ClientManager<T> &manager, boost::asio::basic_stream_socket<T> &&socket) {
|
||||
std::shared_ptr<ClientConnection<T>> self(
|
||||
new ClientConnection(manager, std::move(socket)));
|
||||
// Let our manager process our new connection.
|
||||
self->manager_.ProcessNewClient(self);
|
||||
return self;
|
||||
}
|
||||
|
||||
template <class T>
|
||||
ClientConnection<T>::ClientConnection(ClientManager<T> &manager,
|
||||
boost::asio::basic_stream_socket<T> &&socket)
|
||||
: socket_(std::move(socket)), manager_(manager) {}
|
||||
|
||||
template <class T>
|
||||
void ClientConnection<T>::ProcessMessages() {
|
||||
// Wait for a message header from the client. The message header includes the
|
||||
// protocol version, the message type, and the length of the message.
|
||||
std::vector<boost::asio::mutable_buffer> header;
|
||||
header.push_back(boost::asio::buffer(&read_version_, sizeof(read_version_)));
|
||||
header.push_back(boost::asio::buffer(&read_type_, sizeof(read_type_)));
|
||||
header.push_back(boost::asio::buffer(&read_length_, sizeof(read_length_)));
|
||||
boost::asio::async_read(
|
||||
socket_, header,
|
||||
boost::bind(&ClientConnection<T>::ProcessMessageHeader, this->shared_from_this(),
|
||||
boost::asio::placeholders::error));
|
||||
}
|
||||
|
||||
template <class T>
|
||||
void ClientConnection<T>::ProcessMessageHeader(const boost::system::error_code &error) {
|
||||
if (error) {
|
||||
// If there was an error, disconnect the client.
|
||||
read_type_ = MessageType_DisconnectClient;
|
||||
read_length_ = 0;
|
||||
ProcessMessage(error);
|
||||
return;
|
||||
}
|
||||
|
||||
// If there was no error, make sure the protocol version matches.
|
||||
RAY_CHECK(read_version_ == RayConfig::instance().ray_protocol_version());
|
||||
// Resize the message buffer to match the received length.
|
||||
read_message_.resize(read_length_);
|
||||
// Wait for the message to be read.
|
||||
boost::asio::async_read(
|
||||
socket_, boost::asio::buffer(read_message_),
|
||||
boost::bind(&ClientConnection<T>::ProcessMessage, this->shared_from_this(),
|
||||
boost::asio::placeholders::error));
|
||||
}
|
||||
|
||||
template <class T>
|
||||
void ClientConnection<T>::WriteMessage(int64_t type, size_t length,
|
||||
const uint8_t *message) {
|
||||
std::vector<boost::asio::const_buffer> message_buffers;
|
||||
write_version_ = RayConfig::instance().ray_protocol_version();
|
||||
write_type_ = type;
|
||||
write_length_ = length;
|
||||
write_message_.assign(message, message + length);
|
||||
message_buffers.push_back(boost::asio::buffer(&write_version_, sizeof(write_version_)));
|
||||
message_buffers.push_back(boost::asio::buffer(&write_type_, sizeof(write_type_)));
|
||||
message_buffers.push_back(boost::asio::buffer(&write_length_, sizeof(write_length_)));
|
||||
message_buffers.push_back(boost::asio::buffer(write_message_));
|
||||
boost::system::error_code error;
|
||||
// Write the message and then wait for more messages.
|
||||
boost::asio::async_write(
|
||||
socket_, message_buffers,
|
||||
boost::bind(&ClientConnection<T>::ProcessMessages, this->shared_from_this(),
|
||||
boost::asio::placeholders::error));
|
||||
}
|
||||
|
||||
template <class T>
|
||||
void ClientConnection<T>::ProcessMessage(const boost::system::error_code &error) {
|
||||
if (error) {
|
||||
// TODO(hme): Disconnect differently & remove dependency on node_manager_generated.h
|
||||
read_type_ = MessageType_DisconnectClient;
|
||||
}
|
||||
manager_.ProcessClientMessage(this->shared_from_this(), read_type_,
|
||||
read_message_.data());
|
||||
}
|
||||
|
||||
template <class T>
|
||||
void ClientConnection<T>::ProcessMessages(const boost::system::error_code &error) {
|
||||
if (error) {
|
||||
ProcessMessage(error);
|
||||
} else {
|
||||
ProcessMessages();
|
||||
}
|
||||
}
|
||||
|
||||
template class ClientConnection<boost::asio::local::stream_protocol>;
|
||||
template class ClientConnection<boost::asio::ip::tcp>;
|
||||
|
||||
template <class T>
|
||||
ClientManager<T>::~ClientManager<T>() {}
|
||||
|
||||
template class ClientManager<boost::asio::local::stream_protocol>;
|
||||
template class ClientManager<boost::asio::ip::tcp>;
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,105 @@
|
||||
#ifndef RAY_COMMON_CLIENT_CONNECTION_H
|
||||
#define RAY_COMMON_CLIENT_CONNECTION_H
|
||||
|
||||
#include <memory>
|
||||
|
||||
#include <boost/asio.hpp>
|
||||
#include <boost/asio/error.hpp>
|
||||
#include <boost/enable_shared_from_this.hpp>
|
||||
|
||||
namespace ray {
|
||||
|
||||
template <class T>
|
||||
class ClientManager;
|
||||
|
||||
/// \class ClientConnection
|
||||
///
|
||||
/// A generic type representing a client connection on a server. This class can
|
||||
/// be used to process and write messages asynchronously from and to the
|
||||
/// client.
|
||||
template <class T>
|
||||
class ClientConnection : public std::enable_shared_from_this<ClientConnection<T>> {
|
||||
public:
|
||||
/// Allocate a new node client connection.
|
||||
///
|
||||
/// \param ClientManager A reference to the manager that will process a
|
||||
/// message from this client.
|
||||
/// \param socket The client socket.
|
||||
/// \return std::shared_ptr<ClientConnection>.
|
||||
static std::shared_ptr<ClientConnection<T>> Create(
|
||||
ClientManager<T> &manager, boost::asio::basic_stream_socket<T> &&socket);
|
||||
|
||||
/// Listen for and process messages from the client connection. Once a
|
||||
/// message has been fully received, the client manager's
|
||||
/// ProcessClientMessage handler will be called.
|
||||
void ProcessMessages();
|
||||
|
||||
/// Write a message to the client and then listen for more messages.
|
||||
///
|
||||
/// \param type The message type (e.g., a flatbuffer enum).
|
||||
/// \param length The size in bytes of the message.
|
||||
/// \param message A pointer to the message buffer. This will be copied into
|
||||
/// the ClientConnection's buffer.
|
||||
void WriteMessage(int64_t type, size_t length, const uint8_t *message);
|
||||
|
||||
private:
|
||||
/// A private constructor for a node client connection.
|
||||
ClientConnection(ClientManager<T> &manager,
|
||||
boost::asio::basic_stream_socket<T> &&socket);
|
||||
/// Process an error from the last operation, then process the message
|
||||
/// header from the client.
|
||||
void ProcessMessageHeader(const boost::system::error_code &error);
|
||||
/// Process an error from reading the message header, then process the
|
||||
/// message from the client.
|
||||
void ProcessMessage(const boost::system::error_code &error);
|
||||
/// Process an error from the last operation and then listen for more
|
||||
/// messages.
|
||||
void ProcessMessages(const boost::system::error_code &error);
|
||||
|
||||
/// The client socket.
|
||||
boost::asio::basic_stream_socket<T> socket_;
|
||||
/// A reference to the manager for this client. The manager exposes a handler
|
||||
/// for all messages processed by this client.
|
||||
ClientManager<T> &manager_;
|
||||
/// Buffers for the current message being read rom the client.
|
||||
int64_t read_version_;
|
||||
int64_t read_type_;
|
||||
uint64_t read_length_;
|
||||
std::vector<uint8_t> read_message_;
|
||||
/// Buffers for the current message being written to the client.
|
||||
int64_t write_version_;
|
||||
int64_t write_type_;
|
||||
uint64_t write_length_;
|
||||
std::vector<uint8_t> write_message_;
|
||||
};
|
||||
|
||||
using LocalClientConnection = ClientConnection<boost::asio::local::stream_protocol>;
|
||||
using TcpClientConnection = ClientConnection<boost::asio::ip::tcp>;
|
||||
|
||||
/// \class ClientManager
|
||||
///
|
||||
/// A virtual cliant manager. Derived classes should define a method for
|
||||
/// processing a message on the server sent by the client.
|
||||
template <class T>
|
||||
class ClientManager {
|
||||
public:
|
||||
/// Process a new client connection.
|
||||
///
|
||||
/// \param client A shared pointer to the client that connected.
|
||||
virtual void ProcessNewClient(std::shared_ptr<ClientConnection<T>> client) = 0;
|
||||
|
||||
/// Process a message from a client, then listen for more messages if the
|
||||
/// client is still alive.
|
||||
///
|
||||
/// \param client A shared pointer to the client that sent the message.
|
||||
/// \param message_type The message type (e.g., a flatbuffer enum).
|
||||
/// \param message A pointer to the message buffer.
|
||||
virtual void ProcessClientMessage(std::shared_ptr<ClientConnection<T>> client,
|
||||
int64_t message_type, const uint8_t *message) = 0;
|
||||
|
||||
virtual ~ClientManager() = 0;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_COMMON_CLIENT_CONNECTION_H
|
||||
+2
-6
@@ -69,18 +69,14 @@ void RedisAsioClient::add_read() {
|
||||
operate();
|
||||
}
|
||||
|
||||
void RedisAsioClient::del_read() {
|
||||
read_requested_ = false;
|
||||
}
|
||||
void RedisAsioClient::del_read() { read_requested_ = false; }
|
||||
|
||||
void RedisAsioClient::add_write() {
|
||||
write_requested_ = true;
|
||||
operate();
|
||||
}
|
||||
|
||||
void RedisAsioClient::del_write() {
|
||||
write_requested_ = false;
|
||||
}
|
||||
void RedisAsioClient::del_write() { write_requested_ = false; }
|
||||
|
||||
void RedisAsioClient::cleanup() {}
|
||||
|
||||
|
||||
+2
-4
@@ -21,9 +21,6 @@
|
||||
#ifndef RAY_GCS_ASIO_H
|
||||
#define RAY_GCS_ASIO_H
|
||||
|
||||
#include "hiredis/async.h"
|
||||
#include "hiredis/hiredis.h"
|
||||
|
||||
#include <stdio.h>
|
||||
#include <iostream>
|
||||
#include <string>
|
||||
@@ -32,7 +29,8 @@
|
||||
#include <boost/asio/error.hpp>
|
||||
#include <boost/bind.hpp>
|
||||
|
||||
using boost::asio::ip::tcp;
|
||||
#include "hiredis/async.h"
|
||||
#include "hiredis/hiredis.h"
|
||||
|
||||
class RedisAsioClient {
|
||||
public:
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
#include <iostream>
|
||||
|
||||
#include "gtest/gtest.h"
|
||||
|
||||
#include "asio.h"
|
||||
#include "gtest/gtest.h"
|
||||
#include "ray/util/logging.h"
|
||||
|
||||
boost::asio::io_service io_service;
|
||||
|
||||
@@ -12,13 +12,11 @@ class TestRedisAsioClient : public ::testing::Test {
|
||||
public:
|
||||
TestRedisAsioClient() {
|
||||
int r = system("redis-server > /dev/null & sleep 1");
|
||||
std::cout << "TestRedisAsioClient: redis-server status code was " << r
|
||||
<< std::endl;
|
||||
RAY_LOG(INFO) << "TestRedisAsioClient: redis-server status code was " << r;
|
||||
}
|
||||
~TestRedisAsioClient() {
|
||||
int r = system("redis-cli -c shutdown");
|
||||
std::cout << "TestRedisAsioClient: redis-cli status code was " << r
|
||||
<< std::endl;
|
||||
RAY_LOG(INFO) << "TestRedisAsioClient: redis-cli status code was " << r;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -25,8 +25,7 @@ Status Attach(plasma::EventLoop &event_loop) {
|
||||
}
|
||||
|
||||
Status AsyncGcsClient::Attach(boost::asio::io_service &io_service) {
|
||||
asio_client_.reset(
|
||||
new RedisAsioClient(io_service, context_->async_context()));
|
||||
asio_client_.reset(new RedisAsioClient(io_service, context_->async_context()));
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
|
||||
@@ -5,10 +5,10 @@
|
||||
#include <string>
|
||||
|
||||
#include "plasma/events.h"
|
||||
#include "ray/id.h"
|
||||
#include "ray/status.h"
|
||||
#include "ray/gcs/asio.h"
|
||||
#include "ray/gcs/tables.h"
|
||||
#include "ray/id.h"
|
||||
#include "ray/status.h"
|
||||
#include "ray/util/logging.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
@@ -75,8 +75,7 @@ void TestObjectTable(const UniqueID &job_id, gcs::AsyncGcsClient &client) {
|
||||
data->managers.push_back("A");
|
||||
data->managers.push_back("B");
|
||||
ObjectID object_id = ObjectID::from_random();
|
||||
RAY_CHECK_OK(
|
||||
client.object_table().Add(job_id, object_id, data, &ObjectAdded));
|
||||
RAY_CHECK_OK(client.object_table().Add(job_id, object_id, data, &ObjectAdded));
|
||||
RAY_CHECK_OK(client.object_table().Lookup(job_id, object_id, &Lookup));
|
||||
// Run the event loop. The loop will only stop if the Lookup callback is
|
||||
// called (or an assertion failure).
|
||||
@@ -123,8 +122,7 @@ void TaskUpdateCallback(gcs::AsyncGcsClient *client,
|
||||
void TestTaskTable(const UniqueID &job_id, gcs::AsyncGcsClient &client) {
|
||||
auto data = std::make_shared<TaskTableDataT>();
|
||||
data->scheduling_state = SchedulingState_SCHEDULED;
|
||||
DBClientID local_scheduler_id =
|
||||
DBClientID::from_binary("abcdefghijklmnopqrst");
|
||||
ClientID local_scheduler_id = ClientID::from_binary("abcdefghijklmnopqrst");
|
||||
data->scheduler_id = local_scheduler_id.binary();
|
||||
TaskID task_id = TaskID::from_random();
|
||||
RAY_CHECK_OK(client.task_table().Add(job_id, task_id, data, &TaskAdded));
|
||||
@@ -135,8 +133,8 @@ void TestTaskTable(const UniqueID &job_id, gcs::AsyncGcsClient &client) {
|
||||
update->update_state = SchedulingState_LOST;
|
||||
// After test-and-setting, the callback will lookup the current state of the
|
||||
// task.
|
||||
RAY_CHECK_OK(client.task_table().TestAndUpdate(job_id, task_id, update,
|
||||
&TaskUpdateCallback));
|
||||
RAY_CHECK_OK(
|
||||
client.task_table().TestAndUpdate(job_id, task_id, update, &TaskUpdateCallback));
|
||||
// Run the event loop. The loop will only stop if the lookup after the
|
||||
// test-and-set succeeds (or an assertion failure).
|
||||
test->Start();
|
||||
@@ -152,8 +150,7 @@ TEST_F(TestGcsWithAsio, TestTaskTable) {
|
||||
TestTaskTable(job_id_, client_);
|
||||
}
|
||||
|
||||
void ObjectTableSubscribed(gcs::AsyncGcsClient *client,
|
||||
const UniqueID &id,
|
||||
void ObjectTableSubscribed(gcs::AsyncGcsClient *client, const UniqueID &id,
|
||||
std::shared_ptr<ObjectTableDataT> data) {
|
||||
test->Stop();
|
||||
}
|
||||
@@ -171,8 +168,7 @@ void TestSubscribeAll(const UniqueID &job_id, gcs::AsyncGcsClient &client) {
|
||||
data->managers.push_back("A");
|
||||
data->managers.push_back("B");
|
||||
ObjectID object_id = ObjectID::from_random();
|
||||
RAY_CHECK_OK(
|
||||
client.object_table().Add(job_id, object_id, data, &ObjectAdded));
|
||||
RAY_CHECK_OK(client.object_table().Add(job_id, object_id, data, &ObjectAdded));
|
||||
// Run the event loop. The loop will only stop if the registered subscription
|
||||
// callback is called (or an assertion failure).
|
||||
test->Start();
|
||||
|
||||
@@ -73,3 +73,36 @@ table CustomSerializerData {
|
||||
|
||||
table ConfigTableData {
|
||||
}
|
||||
|
||||
table Resource {
|
||||
// The type of the resource.
|
||||
resource_name: string;
|
||||
// The total capacity of this resource type.
|
||||
resource_capacity: double;
|
||||
}
|
||||
|
||||
table ClientTableData {
|
||||
// The client ID of the client that the message is about.
|
||||
client_id: string;
|
||||
// The IP address of the client's node manager.
|
||||
node_manager_address: string;
|
||||
// The port at which the client's node manager is listening for TCP
|
||||
// connections from other node managers.
|
||||
node_manager_port: int;
|
||||
// The port at which the client's object manager is listening for TCP
|
||||
// connections from other object managers.
|
||||
object_manager_port: int;
|
||||
// The total resources of this client.
|
||||
resources_total: [Resource];
|
||||
// True if the message is about the addition of a client and false if it is
|
||||
// about the deletion of a client.
|
||||
is_insertion: bool;
|
||||
}
|
||||
|
||||
table NodeManagerHeartbeat {
|
||||
// The available resources on this node manager. This information may be
|
||||
// stale.
|
||||
resources_available: [Resource];
|
||||
// The total resources on this node manager.
|
||||
resources_total: [Resource];
|
||||
}
|
||||
|
||||
@@ -62,8 +62,7 @@ void SubscribeRedisCallback(void *c, void *r, void *privdata) {
|
||||
data = std::string(message->str, message->len);
|
||||
RAY_CHECK(!data.empty()) << "Empty message received on subscribe channel";
|
||||
} else {
|
||||
RAY_LOG(FATAL) << "Fatal redis error during subscribe"
|
||||
<< message_type->str;
|
||||
RAY_LOG(FATAL) << "Fatal redis error during subscribe" << message_type->str;
|
||||
}
|
||||
|
||||
// NOTE(swang): We do not delete the callback after calling it since there
|
||||
@@ -72,8 +71,8 @@ void SubscribeRedisCallback(void *c, void *r, void *privdata) {
|
||||
} else if (reply->type == REDIS_REPLY_ERROR) {
|
||||
RAY_LOG(ERROR) << "Redis error " << reply->str;
|
||||
} else {
|
||||
RAY_LOG(FATAL) << "Fatal redis error of type " << reply->type
|
||||
<< " and with string " << reply->str;
|
||||
RAY_LOG(FATAL) << "Fatal redis error of type " << reply->type << " and with string "
|
||||
<< reply->str;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -144,8 +143,8 @@ Status RedisContext::Connect(const std::string &address, int port) {
|
||||
// Connect to subscribe context
|
||||
subscribe_context_ = redisAsyncConnect(address.c_str(), port);
|
||||
if (subscribe_context_ == nullptr || subscribe_context_->err) {
|
||||
RAY_LOG(FATAL) << "Could not establish subscribe connection to redis "
|
||||
<< address << ":" << port;
|
||||
RAY_LOG(FATAL) << "Could not establish subscribe connection to redis " << address
|
||||
<< ":" << port;
|
||||
}
|
||||
return Status::OK();
|
||||
}
|
||||
@@ -159,29 +158,24 @@ Status RedisContext::AttachToEventLoop(aeEventLoop *loop) {
|
||||
}
|
||||
}
|
||||
|
||||
Status RedisContext::RunAsync(const std::string &command,
|
||||
const UniqueID &id,
|
||||
uint8_t *data,
|
||||
int64_t length,
|
||||
const TablePubsub pubsub_channel,
|
||||
int64_t callback_index) {
|
||||
Status RedisContext::RunAsync(const std::string &command, const UniqueID &id,
|
||||
uint8_t *data, int64_t length,
|
||||
const TablePubsub pubsub_channel, int64_t callback_index) {
|
||||
if (length > 0) {
|
||||
std::string redis_command = command + " %d %b %b";
|
||||
int status = redisAsyncCommand(
|
||||
async_context_,
|
||||
reinterpret_cast<redisCallbackFn *>(&GlobalRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(),
|
||||
pubsub_channel, id.data(), id.size(), data, length);
|
||||
async_context_, reinterpret_cast<redisCallbackFn *>(&GlobalRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(), pubsub_channel,
|
||||
id.data(), id.size(), data, length);
|
||||
if (status == REDIS_ERR) {
|
||||
return Status::RedisError(std::string(async_context_->errstr));
|
||||
}
|
||||
} else {
|
||||
std::string redis_command = command + " %d %b";
|
||||
int status = redisAsyncCommand(
|
||||
async_context_,
|
||||
reinterpret_cast<redisCallbackFn *>(&GlobalRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(),
|
||||
pubsub_channel, id.data(), id.size());
|
||||
async_context_, reinterpret_cast<redisCallbackFn *>(&GlobalRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(), pubsub_channel,
|
||||
id.data(), id.size());
|
||||
if (status == REDIS_ERR) {
|
||||
return Status::RedisError(std::string(async_context_->errstr));
|
||||
}
|
||||
@@ -200,19 +194,16 @@ Status RedisContext::SubscribeAsync(const ClientID &client_id,
|
||||
// Subscribe to all messages.
|
||||
std::string redis_command = "SUBSCRIBE %d";
|
||||
status = redisAsyncCommand(
|
||||
subscribe_context_,
|
||||
reinterpret_cast<redisCallbackFn *>(&SubscribeRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(),
|
||||
pubsub_channel);
|
||||
subscribe_context_, reinterpret_cast<redisCallbackFn *>(&SubscribeRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(), pubsub_channel);
|
||||
} else {
|
||||
// Subscribe only to messages sent to this client.
|
||||
// TODO(swang): Nobody sends on this channel yet.
|
||||
std::string redis_command = "SUBSCRIBE %d:%b";
|
||||
status = redisAsyncCommand(
|
||||
subscribe_context_,
|
||||
reinterpret_cast<redisCallbackFn *>(&SubscribeRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(),
|
||||
pubsub_channel, client_id.data(), client_id.size());
|
||||
subscribe_context_, reinterpret_cast<redisCallbackFn *>(&SubscribeRedisCallback),
|
||||
reinterpret_cast<void *>(callback_index), redis_command.c_str(), pubsub_channel,
|
||||
client_id.data(), client_id.size());
|
||||
}
|
||||
|
||||
if (status == REDIS_ERR) {
|
||||
|
||||
@@ -50,14 +50,10 @@ class RedisContext {
|
||||
~RedisContext();
|
||||
Status Connect(const std::string &address, int port);
|
||||
Status AttachToEventLoop(aeEventLoop *loop);
|
||||
Status RunAsync(const std::string &command,
|
||||
const UniqueID &id,
|
||||
uint8_t *data,
|
||||
int64_t length,
|
||||
const TablePubsub pubsub_channel,
|
||||
Status RunAsync(const std::string &command, const UniqueID &id, uint8_t *data,
|
||||
int64_t length, const TablePubsub pubsub_channel,
|
||||
int64_t callback_index);
|
||||
Status SubscribeAsync(const ClientID &client_id,
|
||||
const TablePubsub pubsub_channel,
|
||||
Status SubscribeAsync(const ClientID &client_id, const TablePubsub pubsub_channel,
|
||||
int64_t callback_index);
|
||||
redisAsyncContext *async_context() { return async_context_; }
|
||||
|
||||
|
||||
+44
-12
@@ -2,15 +2,13 @@
|
||||
|
||||
#include "ray/gcs/client.h"
|
||||
|
||||
#include "task.h"
|
||||
#include "common_protocol.h"
|
||||
|
||||
namespace {
|
||||
|
||||
std::shared_ptr<TaskTableDataT> MakeTaskTableData(
|
||||
const TaskExecutionSpec &execution_spec,
|
||||
const DBClientID &local_scheduler_id,
|
||||
SchedulingState scheduling_state) {
|
||||
std::shared_ptr<TaskTableDataT> MakeTaskTableData(const TaskExecutionSpec &execution_spec,
|
||||
const ClientID &local_scheduler_id,
|
||||
SchedulingState scheduling_state) {
|
||||
auto data = std::make_shared<TaskTableDataT>();
|
||||
data->scheduling_state = scheduling_state;
|
||||
data->task_info =
|
||||
@@ -50,13 +48,10 @@ Status TaskTableAdd(AsyncGcsClient *gcs_client, Task *task) {
|
||||
|
||||
// TODO(pcm): This is a helper method that should go away once we get rid of
|
||||
// the Task* datastructure and replace it with TaskTableDataT.
|
||||
Status TaskTableTestAndUpdate(
|
||||
AsyncGcsClient *gcs_client,
|
||||
const TaskID &task_id,
|
||||
const DBClientID &local_scheduler_id,
|
||||
int test_state_bitmask,
|
||||
SchedulingState update_state,
|
||||
const TaskTable::TestAndUpdateCallback &callback) {
|
||||
Status TaskTableTestAndUpdate(AsyncGcsClient *gcs_client, const TaskID &task_id,
|
||||
const ClientID &local_scheduler_id, int test_state_bitmask,
|
||||
SchedulingState update_state,
|
||||
const TaskTable::TestAndUpdateCallback &callback) {
|
||||
auto data = std::make_shared<TaskTableTestAndUpdateT>();
|
||||
data->test_scheduler_id = local_scheduler_id.binary();
|
||||
data->test_state_bitmask = test_state_bitmask;
|
||||
@@ -65,6 +60,43 @@ Status TaskTableTestAndUpdate(
|
||||
data, callback);
|
||||
}
|
||||
|
||||
void ClientConnected(gcs::AsyncGcsClient *client, const ClientID &client_id,
|
||||
std::shared_ptr<ObjectTableDataT> data) {}
|
||||
|
||||
ClientTable::ClientTable(const std::shared_ptr<RedisContext> &context,
|
||||
AsyncGcsClient *client)
|
||||
: Table(context, client), client_id_(UniqueID::from_random()) {}
|
||||
|
||||
Status ClientTable::Connect(ClientID *client_id) {
|
||||
auto data = std::make_shared<ClientTableDataT>();
|
||||
data->client_id = client_id_.binary();
|
||||
// TODO(swang): Get the address and port from somewhere.
|
||||
data->node_manager_address = "";
|
||||
data->node_manager_port = 0;
|
||||
data->is_insertion = true;
|
||||
|
||||
// TODO(swang):
|
||||
// - Add ourselves to the client table (easier if this is synchronous).
|
||||
// - Subscribe to the client table.
|
||||
// - Once subscription is complete, read all client table entries once.
|
||||
|
||||
*client_id = client_id_;
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
Status ClientTable::Disconnect() {
|
||||
auto data = std::make_shared<ClientTableDataT>();
|
||||
data->client_id = client_id_.binary();
|
||||
// TODO(swang): Get the address and port from somewhere.
|
||||
data->node_manager_address = "";
|
||||
data->node_manager_port = 0;
|
||||
data->is_insertion = false;
|
||||
|
||||
// TODO(swang):
|
||||
// - Add ourselves to the client table (easier if this is synchronous).
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace gcs
|
||||
|
||||
} // namespace ray
|
||||
|
||||
+108
-59
@@ -42,18 +42,15 @@ class Table {
|
||||
};
|
||||
|
||||
Table(const std::shared_ptr<RedisContext> &context, AsyncGcsClient *client)
|
||||
: context_(context),
|
||||
client_(client),
|
||||
pubsub_channel_(TablePubsub_NO_PUBLISH){};
|
||||
: context_(context), client_(client), pubsub_channel_(TablePubsub_NO_PUBLISH){};
|
||||
|
||||
/// Add an entry to the table.
|
||||
///
|
||||
/// @param job_id The ID of the job (= driver).
|
||||
/// @param id The ID of the data that is added to the GCS.
|
||||
/// @param data Data that is added to the GCS.
|
||||
/// @param done Callback that is called once the data has been written to the
|
||||
/// GCS.
|
||||
/// @return Status
|
||||
/// \param job_id The ID of the job (= driver).
|
||||
/// \param id The ID of the data that is added to the GCS.
|
||||
/// \param data Data that is added to the GCS.
|
||||
/// \param done Callback that is called once the data has been written to the GCS.
|
||||
/// \return Status
|
||||
Status Add(const JobID &job_id,
|
||||
const ID &id,
|
||||
std::shared_ptr<DataT> data,
|
||||
@@ -65,18 +62,17 @@ class Table {
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
fbb.ForceDefaults(true);
|
||||
fbb.Finish(Data::Pack(fbb, data.get()));
|
||||
RAY_RETURN_NOT_OK(context_->RunAsync("RAY.TABLE_ADD", id,
|
||||
fbb.GetBufferPointer(), fbb.GetSize(),
|
||||
pubsub_channel_, callback_index));
|
||||
RAY_RETURN_NOT_OK(context_->RunAsync("RAY.TABLE_ADD", id, fbb.GetBufferPointer(),
|
||||
fbb.GetSize(), pubsub_channel_, callback_index));
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
/// Lookup an entry asynchronously.
|
||||
///
|
||||
/// @param job_id The ID of the job (= driver).
|
||||
/// @param id The ID of the data that is looked up in the GCS.
|
||||
/// @param lookup Callback that is called after lookup.
|
||||
/// @return Status
|
||||
/// \param job_id The ID of the job (= driver).
|
||||
/// \param id The ID of the data that is looked up in the GCS.
|
||||
/// \param lookup Callback that is called after lookup.
|
||||
/// \return Status
|
||||
Status Lookup(const JobID &job_id, const ID &id, const Callback &lookup) {
|
||||
auto d = std::shared_ptr<CallbackData>(
|
||||
new CallbackData({id, nullptr, lookup, this}));
|
||||
@@ -88,30 +84,27 @@ class Table {
|
||||
(d->callback)(d->client, d->id, result);
|
||||
});
|
||||
std::vector<uint8_t> nil;
|
||||
RAY_RETURN_NOT_OK(context_->RunAsync("RAY.TABLE_LOOKUP", id, nil.data(),
|
||||
nil.size(), pubsub_channel_,
|
||||
callback_index));
|
||||
RAY_RETURN_NOT_OK(context_->RunAsync("RAY.TABLE_LOOKUP", id, nil.data(), nil.size(),
|
||||
pubsub_channel_, callback_index));
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
/// Subscribe to updates of this table
|
||||
///
|
||||
/// @param job_id The ID of the job (= driver).
|
||||
/// @param client_id The type of update to listen to. If this is nil, then a
|
||||
/// \param job_id The ID of the job (= driver).
|
||||
/// \param client_id The type of update to listen to. If this is nil, then a
|
||||
/// message for each Add to the table will be received. Else, only
|
||||
/// messages for the given client will be received.
|
||||
/// @param subscribe Callback that is called on each received message.
|
||||
/// @param done Callback that is called when subscription is complete and we
|
||||
/// \param subscribe Callback that is called on each received message.
|
||||
/// \param done Callback that is called when subscription is complete and we
|
||||
/// are ready to receive messages..
|
||||
/// @return Status
|
||||
Status Subscribe(const JobID &job_id,
|
||||
const ClientID &client_id,
|
||||
const Callback &subscribe,
|
||||
const Callback &done) {
|
||||
/// \return Status
|
||||
Status Subscribe(const JobID &job_id, const ClientID &client_id,
|
||||
const Callback &subscribe, const Callback &done) {
|
||||
auto d = std::shared_ptr<CallbackData>(
|
||||
new CallbackData({client_id, nullptr, subscribe, this}));
|
||||
int64_t callback_index = RedisCallbackManager::instance().add(
|
||||
[done, d](const std::string &data) {
|
||||
int64_t callback_index =
|
||||
RedisCallbackManager::instance().add([done, d](const std::string &data) {
|
||||
if (data.empty()) {
|
||||
// No data is provided. This is the callback for the initial
|
||||
// subscription request.
|
||||
@@ -140,8 +133,7 @@ class Table {
|
||||
|
||||
class ObjectTable : public Table<ObjectID, ObjectTableData> {
|
||||
public:
|
||||
ObjectTable(const std::shared_ptr<RedisContext> &context,
|
||||
AsyncGcsClient *client)
|
||||
ObjectTable(const std::shared_ptr<RedisContext> &context, AsyncGcsClient *client)
|
||||
: Table(context, client) {
|
||||
pubsub_channel_ = TablePubsub_OBJECT;
|
||||
};
|
||||
@@ -151,12 +143,12 @@ class ObjectTable : public Table<ObjectID, ObjectTableData> {
|
||||
/// objects from the object table. The callback will be called once per
|
||||
/// notification received on this channel.
|
||||
///
|
||||
/// @param subscribe_all
|
||||
/// @param object_available_callback Callback to be called when new object
|
||||
/// \param subscribe_all
|
||||
/// \param object_available_callback Callback to be called when new object
|
||||
/// becomes available.
|
||||
/// @param done_callback Callback to be called when subscription is installed.
|
||||
/// \param done_callback Callback to be called when subscription is installed.
|
||||
/// This is only used for the tests.
|
||||
/// @return Status
|
||||
/// \return Status
|
||||
Status SubscribeToNotifications(const JobID &job_id,
|
||||
bool subscribe_all,
|
||||
const Callback &object_available,
|
||||
@@ -168,8 +160,8 @@ class ObjectTable : public Table<ObjectID, ObjectTableData> {
|
||||
/// notification channel, which was set up by the method
|
||||
/// ObjectTableSubscribeToNotifications.
|
||||
///
|
||||
/// @param object_ids The object IDs to receive notifications about.
|
||||
/// @return Status
|
||||
/// \param object_ids The object IDs to receive notifications about.
|
||||
/// \return Status
|
||||
Status RequestNotifications(const JobID &job_id,
|
||||
const std::vector<ObjectID> &object_ids);
|
||||
};
|
||||
@@ -183,8 +175,7 @@ using ActorTable = Table<ActorID, ActorTableData>;
|
||||
|
||||
class TaskTable : public Table<TaskID, TaskTableData> {
|
||||
public:
|
||||
TaskTable(const std::shared_ptr<RedisContext> &context,
|
||||
AsyncGcsClient *client)
|
||||
TaskTable(const std::shared_ptr<RedisContext> &context, AsyncGcsClient *client)
|
||||
: Table(context, client) {
|
||||
pubsub_channel_ = TablePubsub_TASK;
|
||||
};
|
||||
@@ -202,14 +193,14 @@ class TaskTable : public Table<TaskID, TaskTableData> {
|
||||
/// this function. This assumes that the task spec already exists in the task
|
||||
/// table entry.
|
||||
///
|
||||
/// @param task_id The task ID of the task entry to update.
|
||||
/// @param test_state_bitmask The bitmask to apply to the task entry's current
|
||||
/// \param task_id The task ID of the task entry to update.
|
||||
/// \param test_state_bitmask The bitmask to apply to the task entry's current
|
||||
/// scheduling state. The update happens if and only if the current
|
||||
/// scheduling state AND-ed with the bitmask is greater than 0.
|
||||
/// @param update_state The value to update the task entry's scheduling state
|
||||
/// \param update_state The value to update the task entry's scheduling state
|
||||
/// with, if the current state matches test_state_bitmask.
|
||||
/// @param callback Function to be called when database returns result.
|
||||
/// @return Status
|
||||
/// \param callback Function to be called when database returns result.
|
||||
/// \return Status
|
||||
Status TestAndUpdate(const JobID &job_id,
|
||||
const TaskID &id,
|
||||
std::shared_ptr<TaskTableTestAndUpdateT> data,
|
||||
@@ -236,23 +227,20 @@ class TaskTable : public Table<TaskID, TaskTableData> {
|
||||
/// Events include changes to the task's scheduling state or changes to the
|
||||
/// task's local scheduler ID.
|
||||
///
|
||||
/// @param local_scheduler_id The db_client_id of the local scheduler whose
|
||||
/// \param local_scheduler_id The db_client_id of the local scheduler whose
|
||||
/// events we want to listen to. If you want to subscribe to updates
|
||||
/// from
|
||||
/// all local schedulers, pass in NIL_ID.
|
||||
/// @param subscribe_callback Callback that will be called when the task table
|
||||
/// is
|
||||
/// \param subscribe_callback Callback that will be called when the task table is
|
||||
/// updated.
|
||||
/// @param state_filter Events we want to listen to. Can have values from the
|
||||
/// \param state_filter Events we want to listen to. Can have values from the
|
||||
/// enum "scheduling_state" in task.h.
|
||||
/// TODO(pcm): Make it possible to combine these using flags like
|
||||
/// TASK_STATUS_WAITING | TASK_STATUS_SCHEDULED.
|
||||
/// @param callback Function to be called when database returns result.
|
||||
/// @return Status
|
||||
Status SubscribeToTask(const JobID &job_id,
|
||||
const DBClientID &local_scheduler_id,
|
||||
int state_filter,
|
||||
const SubscribeToTaskCallback &callback,
|
||||
/// \param callback Function to be called when database returns result.
|
||||
/// \return Status
|
||||
Status SubscribeToTask(const JobID &job_id, const ClientID &local_scheduler_id,
|
||||
int state_filter, const SubscribeToTaskCallback &callback,
|
||||
const Callback &done);
|
||||
};
|
||||
|
||||
@@ -264,13 +252,74 @@ using ConfigTable = Table<ConfigID, ConfigTableData>;
|
||||
|
||||
Status TaskTableAdd(AsyncGcsClient *gcs_client, Task *task);
|
||||
|
||||
Status TaskTableTestAndUpdate(AsyncGcsClient *gcs_client,
|
||||
const TaskID &task_id,
|
||||
const DBClientID &local_scheduler_id,
|
||||
int test_state_bitmask,
|
||||
Status TaskTableTestAndUpdate(AsyncGcsClient *gcs_client, const TaskID &task_id,
|
||||
const ClientID &local_scheduler_id, int test_state_bitmask,
|
||||
SchedulingState update_state,
|
||||
const TaskTable::TestAndUpdateCallback &callback);
|
||||
|
||||
/// \class ClientInformation
|
||||
///
|
||||
/// Represents information in the client table about a particular client. Each
|
||||
/// client has an associated node manager.
|
||||
class ClientInformation {
|
||||
public:
|
||||
/// Create a client information object.
|
||||
///
|
||||
/// \param client_table_entry A serialized client table entry flatbuffer.
|
||||
ClientInformation(const ClientTableData &client_table_entry);
|
||||
|
||||
/// Get the client ID.
|
||||
///
|
||||
/// \return The ID of this client.
|
||||
const ClientID &GetClientId() const;
|
||||
|
||||
/// Get the IP address of the client's node manager.
|
||||
///
|
||||
/// \return The IP address of the client's node manager.
|
||||
const std::string GetIpAddress() const;
|
||||
|
||||
/// Get the port at which the client's node manager is listening for
|
||||
/// TCP connections.
|
||||
///
|
||||
/// \return The client's TCP port.
|
||||
int GetPort() const;
|
||||
|
||||
/// Get whether the client is alive.
|
||||
///
|
||||
/// \return Whether the client is alive.
|
||||
bool IsAlive() const;
|
||||
};
|
||||
|
||||
class ClientTable : private Table<ClientID, ClientTableData> {
|
||||
public:
|
||||
ClientTable(const std::shared_ptr<RedisContext> &context, AsyncGcsClient *client);
|
||||
|
||||
/// Connect as a client to the GCS. This registers us in the client table and
|
||||
/// begins subscription to client table notifications.
|
||||
///
|
||||
/// \param[out] client_id The assigned client ID will be written to this pointer.
|
||||
/// \return Status
|
||||
// TODO(swang): Call this from AsyncGcsClient::Connect?
|
||||
ray::Status Connect(ClientID *client_id);
|
||||
|
||||
/// Disconnect the client from the GCS. The client ID assigned during
|
||||
/// registration should never be reused after disconnecting.
|
||||
///
|
||||
/// \return Status
|
||||
ray::Status Disconnect();
|
||||
|
||||
/// Get a client's information from the cache.
|
||||
///
|
||||
/// \param client The client to get information about.
|
||||
const ClientInformation &GetClientInformation(const ClientID &client);
|
||||
|
||||
private:
|
||||
/// This client's ID.
|
||||
ClientID client_id_;
|
||||
/// A cache for information about all clients.
|
||||
std::unordered_map<ClientID, ClientInformation, UniqueIDHasher> client_cache_;
|
||||
};
|
||||
|
||||
} // namespace gcs
|
||||
|
||||
} // namespace ray
|
||||
|
||||
+4
-3
@@ -10,6 +10,9 @@
|
||||
#include "ray/constants.h"
|
||||
#include "ray/util/visibility.h"
|
||||
|
||||
// TODO(swang): Make task ID prefix of any object ID return values and puts so
|
||||
// that we can co-locate task and object entries in the GCS.
|
||||
|
||||
namespace ray {
|
||||
|
||||
class RAY_EXPORT UniqueID {
|
||||
@@ -55,10 +58,8 @@ typedef UniqueID ActorID;
|
||||
typedef UniqueID ActorHandleID;
|
||||
typedef UniqueID WorkerID;
|
||||
typedef UniqueID DriverID;
|
||||
// TODO(swang): Replace this with ClientID.
|
||||
typedef UniqueID DBClientID;
|
||||
typedef UniqueID ClientID;
|
||||
typedef UniqueID ConfigID;
|
||||
typedef UniqueID ClientID;
|
||||
|
||||
} // namespace ray
|
||||
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
include(${CMAKE_CURRENT_LIST_DIR}/../../common/cmake/Common.cmake)
|
||||
|
||||
set(OBJECT_MANAGER_FBS_SRC "${CMAKE_CURRENT_LIST_DIR}/format/object_manager.fbs")
|
||||
set(OUTPUT_DIR ${CMAKE_CURRENT_LIST_DIR}/format/)
|
||||
|
||||
set(OBJECT_MANAGER_FBS_OUTPUT_FILES
|
||||
"${OUTPUT_DIR}/object_manager_generated.h")
|
||||
|
||||
add_custom_command(
|
||||
OUTPUT ${OBJECT_MANAGER_FBS_OUTPUT_FILES}
|
||||
# The --gen-object-api flag generates a C++ class MessageT for each
|
||||
# flatbuffers message Message, which can be used to store deserialized
|
||||
# messages in data structures. This is currently used for ObjectInfo for
|
||||
# example.
|
||||
COMMAND ${FLATBUFFERS_COMPILER} -c -o ${OUTPUT_DIR} ${OBJECT_MANAGER_FBS_SRC} --cpp --gen-object-api --gen-mutable
|
||||
DEPENDS ${FBS_DEPENDS}
|
||||
COMMENT "Running flatc compiler on ${OBJECT_MANAGER_FBS_SRC}"
|
||||
VERBATIM)
|
||||
|
||||
add_custom_target(gen_object_manager_fbs DEPENDS ${OBJECT_MANAGER_FBS_OUTPUT_FILES})
|
||||
|
||||
ADD_RAY_TEST(object_manager_test STATIC_LINK_LIBS ray_static ${PLASMA_STATIC_LIB} ${ARROW_STATIC_LIB} gtest gtest_main pthread ${Boost_SYSTEM_LIBRARY})
|
||||
|
||||
add_library(object_manager object_manager.cc object_manager.h ${OBJECT_MANAGER_FBS_OUTPUT_FILES})
|
||||
target_link_libraries(object_manager common ray_static ${PLASMA_STATIC_LIB} ${ARROW_STATIC_LIB} ${Boost_SYSTEM_LIBRARY})
|
||||
|
||||
install(FILES
|
||||
object_manager
|
||||
DESTINATION "${CMAKE_INSTALL_INCLUDEDIR}/ray/object_manager")
|
||||
@@ -0,0 +1,30 @@
|
||||
// Object Manager protocol specification
|
||||
|
||||
enum OMMessageType:int {
|
||||
PullRequest = 1
|
||||
}
|
||||
|
||||
table PushRequest {
|
||||
|
||||
}
|
||||
|
||||
table PullRequest {
|
||||
// ID of the requesting client.
|
||||
client_id: string;
|
||||
// Requested ObjectID.
|
||||
object_id: string;
|
||||
}
|
||||
|
||||
table ClientConnectionInfo {
|
||||
// ID of the connecting client.
|
||||
client_id: string;
|
||||
// Whether this is a transfer connection.
|
||||
is_transfer: bool;
|
||||
}
|
||||
|
||||
table ObjectHeader {
|
||||
// The object ID being transferred.
|
||||
object_id: string;
|
||||
// The size of the object being transferred.
|
||||
object_size: ulong;
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
#include "object_directory.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
ObjectDirectory::ObjectDirectory(std::shared_ptr<GcsClient> gcs_client) {
|
||||
gcs_client_ = gcs_client;
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::ReportObjectAdded(const ObjectID &object_id,
|
||||
const ClientID &client_id) {
|
||||
return gcs_client_->object_table().Add(object_id, client_id, [] {});
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::ReportObjectRemoved(const ObjectID &object_id,
|
||||
const ClientID &client_id) {
|
||||
return gcs_client_->object_table().Remove(object_id, client_id, [] {});
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::GetInformation(const ClientID &client_id,
|
||||
const InfoSuccessCallback &success_cb,
|
||||
const InfoFailureCallback &fail_cb) {
|
||||
gcs_client_->client_table().GetClientInformation(
|
||||
client_id,
|
||||
[this, success_cb, client_id](ClientInformation client_info) {
|
||||
const auto &info =
|
||||
RemoteConnectionInfo(client_id, client_info.GetIp(), client_info.GetPort());
|
||||
success_cb(info);
|
||||
},
|
||||
fail_cb);
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::GetLocations(const ObjectID &object_id,
|
||||
const OnLocationsSuccess &success_cb,
|
||||
const OnLocationsFailure &fail_cb) {
|
||||
ray::Status status_code = ray::Status::OK();
|
||||
if (existing_requests_.count(object_id) == 0) {
|
||||
existing_requests_[object_id] = ODCallbacks({success_cb, fail_cb});
|
||||
status_code = ExecuteGetLocations(object_id);
|
||||
} else {
|
||||
// Do nothing. A request is in progress.
|
||||
}
|
||||
return status_code;
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::ExecuteGetLocations(const ObjectID &object_id) {
|
||||
// TODO(hme): Avoid callback hell.
|
||||
std::vector<RemoteConnectionInfo> remote_connections;
|
||||
ray::Status status = gcs_client_->object_table().GetObjectClientIDs(
|
||||
object_id,
|
||||
[this, object_id, &remote_connections](const std::vector<ClientID> &client_ids) {
|
||||
gcs_client_->client_table().GetClientInformationSet(
|
||||
client_ids,
|
||||
[this, object_id,
|
||||
&remote_connections](const std::vector<ClientInformation> &info_vec) {
|
||||
for (const auto &client_info : info_vec) {
|
||||
RemoteConnectionInfo info =
|
||||
RemoteConnectionInfo(client_info.GetClientId(), client_info.GetIp(),
|
||||
client_info.GetPort());
|
||||
remote_connections.push_back(info);
|
||||
}
|
||||
ray::Status cb_completion_status =
|
||||
GetLocationsComplete(Status::OK(), object_id, remote_connections);
|
||||
},
|
||||
[this, object_id, &remote_connections](const Status &status) {
|
||||
ray::Status cb_completion_status =
|
||||
GetLocationsComplete(status, object_id, remote_connections);
|
||||
});
|
||||
},
|
||||
[this, object_id, &remote_connections](const Status &status) {
|
||||
ray::Status cb_completion_status =
|
||||
GetLocationsComplete(status, object_id, remote_connections);
|
||||
});
|
||||
return status;
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::GetLocationsComplete(
|
||||
const ray::Status &status, const ObjectID &object_id,
|
||||
const std::vector<RemoteConnectionInfo> &remote_connections) {
|
||||
bool success = status.ok();
|
||||
// Only invoke a callback if the request was not cancelled.
|
||||
if (existing_requests_.count(object_id) > 0) {
|
||||
ODCallbacks cbs = existing_requests_[object_id];
|
||||
if (success) {
|
||||
cbs.success_cb(remote_connections, object_id);
|
||||
} else {
|
||||
cbs.fail_cb(status, object_id);
|
||||
}
|
||||
}
|
||||
existing_requests_.erase(object_id);
|
||||
return status;
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::Cancel(const ObjectID &object_id) {
|
||||
existing_requests_.erase(object_id);
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectDirectory::Terminate() { return ray::Status::OK(); };
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,134 @@
|
||||
#ifndef RAY_OBJECT_MANAGER_OBJECT_DIRECTORY_H
|
||||
#define RAY_OBJECT_MANAGER_OBJECT_DIRECTORY_H
|
||||
|
||||
#include <memory>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
#include <vector>
|
||||
|
||||
#include "ray/id.h"
|
||||
#include "ray/raylet/mock_gcs_client.h"
|
||||
#include "ray/status.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
struct RemoteConnectionInfo {
|
||||
RemoteConnectionInfo(const ClientID &id, const std::string &ip_address,
|
||||
uint16_t port_num)
|
||||
: client_id(id), ip(ip_address), port(port_num) {}
|
||||
ClientID client_id;
|
||||
std::string ip;
|
||||
uint16_t port;
|
||||
};
|
||||
|
||||
/// Connection information for remote object managers.
|
||||
class ObjectDirectoryInterface {
|
||||
public:
|
||||
ObjectDirectoryInterface() = default;
|
||||
virtual ~ObjectDirectoryInterface() = default;
|
||||
|
||||
/// Callbacks for GetInformation.
|
||||
using InfoSuccessCallback = std::function<void(const ray::RemoteConnectionInfo &info)>;
|
||||
using InfoFailureCallback = std::function<void(ray::Status status)>;
|
||||
|
||||
/// This is used to establish object manager client connections.
|
||||
///
|
||||
/// \param client_id The client for which information is required.
|
||||
/// \param success_cb A callback which handles the success of this method.
|
||||
/// \param fail_cb A callback which handles the failure of this method.
|
||||
/// \return Status of whether this asynchronous request succeeded.
|
||||
virtual ray::Status GetInformation(const ClientID &client_id,
|
||||
const InfoSuccessCallback &success_cb,
|
||||
const InfoFailureCallback &fail_cb) = 0;
|
||||
|
||||
// Callbacks for GetLocations.
|
||||
using OnLocationsSuccess = std::function<void(
|
||||
const std::vector<ray::RemoteConnectionInfo> &v, const ray::ObjectID &object_id)>;
|
||||
using OnLocationsFailure =
|
||||
std::function<void(ray::Status status, const ray::ObjectID &object_id)>;
|
||||
|
||||
/// Asynchronously obtain the locations of an object by ObjectID.
|
||||
/// This is used to handle object pulls.
|
||||
///
|
||||
/// \param object_id The required object's ObjectID.
|
||||
/// \param success_cb Invoked upon success with list of remote connection info.
|
||||
/// \param fail_cb Invoked upon failure with ray status and object id.
|
||||
/// \return Status of whether this asynchronous request succeeded.
|
||||
virtual ray::Status GetLocations(const ObjectID &object_id,
|
||||
const OnLocationsSuccess &success_cb,
|
||||
const OnLocationsFailure &fail_cb) = 0;
|
||||
|
||||
/// Cancels the invocation of the callback associated with callback_id.
|
||||
///
|
||||
/// \param object_id The object id invoked with GetLocations.
|
||||
/// \return Status of whether this method succeeded.
|
||||
virtual ray::Status Cancel(const ObjectID &object_id) = 0;
|
||||
|
||||
/// Report objects added to this node's store to the object directory.
|
||||
///
|
||||
/// \param object_id The object id that was put into the store.
|
||||
/// \param client_id The client id corresponding to this node.
|
||||
/// \return Status of whether this method succeeded.
|
||||
virtual ray::Status ReportObjectAdded(const ObjectID &object_id,
|
||||
const ClientID &client_id) = 0;
|
||||
|
||||
/// Report objects removed from this client's store to the object directory.
|
||||
///
|
||||
/// \param object_id The object id that was removed from the store.
|
||||
/// \param client_id The client id corresponding to this node.
|
||||
/// \return Status of whether this method succeeded.
|
||||
virtual ray::Status ReportObjectRemoved(const ObjectID &object_id,
|
||||
const ClientID &client_id) = 0;
|
||||
|
||||
/// Terminate this object.
|
||||
///
|
||||
/// \return Status of whether termination succeeded.
|
||||
virtual ray::Status Terminate() = 0;
|
||||
};
|
||||
|
||||
/// Ray ObjectDirectory declaration.
|
||||
class ObjectDirectory : public ObjectDirectoryInterface {
|
||||
public:
|
||||
ObjectDirectory() = default;
|
||||
~ObjectDirectory() override = default;
|
||||
|
||||
ray::Status GetInformation(const ClientID &client_id,
|
||||
const InfoSuccessCallback &success_cb,
|
||||
const InfoFailureCallback &fail_cb) override;
|
||||
ray::Status GetLocations(const ObjectID &object_id,
|
||||
const OnLocationsSuccess &success_cb,
|
||||
const OnLocationsFailure &fail_cb) override;
|
||||
ray::Status Cancel(const ObjectID &object_id) override;
|
||||
ray::Status Terminate() override;
|
||||
ray::Status ReportObjectAdded(const ObjectID &object_id,
|
||||
const ClientID &client_id) override;
|
||||
ray::Status ReportObjectRemoved(const ObjectID &object_id,
|
||||
const ClientID &client_id) override;
|
||||
/// Ray only (not part of the OD interface).
|
||||
ObjectDirectory(std::shared_ptr<GcsClient> gcs_client);
|
||||
|
||||
private:
|
||||
/// Callbacks associated with a call to GetLocations.
|
||||
// TODO(hme): I think these can be removed.
|
||||
struct ODCallbacks {
|
||||
OnLocationsSuccess success_cb;
|
||||
OnLocationsFailure fail_cb;
|
||||
};
|
||||
|
||||
/// Maintain map of in-flight GetLocation requests.
|
||||
std::unordered_map<ObjectID, ODCallbacks, UniqueIDHasher> existing_requests_;
|
||||
|
||||
/// Reference to the gcs client.
|
||||
std::shared_ptr<GcsClient> gcs_client_;
|
||||
|
||||
/// GetLocations registers a request for locations.
|
||||
/// This function actually carries out that request.
|
||||
ray::Status ExecuteGetLocations(const ObjectID &object_id);
|
||||
/// Invoked when call to ExecuteGetLocations completes.
|
||||
ray::Status GetLocationsComplete(const ray::Status &status, const ObjectID &object_id,
|
||||
const std::vector<RemoteConnectionInfo> &v);
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_OBJECT_MANAGER_OBJECT_DIRECTORY_H
|
||||
@@ -0,0 +1,448 @@
|
||||
#include "object_manager.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
ObjectManager::ObjectManager(boost::asio::io_service &io_service,
|
||||
ObjectManagerConfig config,
|
||||
std::shared_ptr<ray::GcsClient> gcs_client)
|
||||
: object_directory_(new ObjectDirectory(gcs_client)), work_(io_service_) {
|
||||
config_ = config;
|
||||
store_client_ = std::unique_ptr<ObjectStoreClient>(
|
||||
new ObjectStoreClient(io_service, config.store_socket_name));
|
||||
store_client_->SubscribeObjAdded(
|
||||
[this](const ObjectID &oid) { NotifyDirectoryObjectAdd(oid); });
|
||||
store_client_->SubscribeObjDeleted(
|
||||
[this](const ObjectID &oid) { NotifyDirectoryObjectDeleted(oid); });
|
||||
StartIOService();
|
||||
};
|
||||
|
||||
ObjectManager::ObjectManager(boost::asio::io_service &io_service,
|
||||
ObjectManagerConfig config,
|
||||
std::unique_ptr<ObjectDirectoryInterface> od)
|
||||
: object_directory_(std::move(od)), work_(io_service_) {
|
||||
config_ = config;
|
||||
store_client_ = std::unique_ptr<ObjectStoreClient>(
|
||||
new ObjectStoreClient(io_service, config.store_socket_name));
|
||||
store_client_->SubscribeObjAdded(
|
||||
[this](const ObjectID &oid) { NotifyDirectoryObjectAdd(oid); });
|
||||
store_client_->SubscribeObjDeleted(
|
||||
[this](const ObjectID &oid) { NotifyDirectoryObjectDeleted(oid); });
|
||||
StartIOService();
|
||||
};
|
||||
|
||||
void ObjectManager::StartIOService() {
|
||||
io_thread_ = std::thread(&ObjectManager::IOServiceLoop, this);
|
||||
// thread_group_.create_thread(boost::bind(&boost::asio::io_service::run,
|
||||
// &io_service_));
|
||||
}
|
||||
|
||||
void ObjectManager::IOServiceLoop() { io_service_.run(); }
|
||||
|
||||
void ObjectManager::StopIOService() {
|
||||
io_service_.stop();
|
||||
io_thread_.join();
|
||||
// thread_group_.join_all();
|
||||
}
|
||||
|
||||
void ObjectManager::SetClientID(const ClientID &client_id) { client_id_ = client_id; }
|
||||
|
||||
ClientID ObjectManager::GetClientID() { return client_id_; }
|
||||
|
||||
void ObjectManager::NotifyDirectoryObjectAdd(const ObjectID &object_id) {
|
||||
ray::Status status = object_directory_->ReportObjectAdded(object_id, client_id_);
|
||||
}
|
||||
|
||||
void ObjectManager::NotifyDirectoryObjectDeleted(const ObjectID &object_id) {
|
||||
ray::Status status = object_directory_->ReportObjectRemoved(object_id, client_id_);
|
||||
}
|
||||
|
||||
ray::Status ObjectManager::Terminate() {
|
||||
StopIOService();
|
||||
ray::Status status_code = object_directory_->Terminate();
|
||||
// TODO: evaluate store client termination status.
|
||||
store_client_->Terminate();
|
||||
return status_code;
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::SubscribeObjAdded(
|
||||
std::function<void(const ObjectID &)> callback) {
|
||||
store_client_->SubscribeObjAdded(callback);
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::SubscribeObjDeleted(
|
||||
std::function<void(const ObjectID &)> callback) {
|
||||
store_client_->SubscribeObjDeleted(callback);
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::Pull(const ObjectID &object_id) {
|
||||
// TODO(hme): Need to correct. Workaround to get all pull requests on the same thread.
|
||||
SchedulePull(object_id, 0);
|
||||
return Status::OK();
|
||||
};
|
||||
|
||||
void ObjectManager::SchedulePull(const ObjectID &object_id, int wait_ms) {
|
||||
pull_requests_[object_id] = Timer(new boost::asio::deadline_timer(
|
||||
io_service_, boost::posix_time::milliseconds(wait_ms)));
|
||||
pull_requests_[object_id]->async_wait(
|
||||
boost::bind(&ObjectManager::SchedulePullHandler, this, object_id));
|
||||
}
|
||||
|
||||
ray::Status ObjectManager::SchedulePullHandler(const ObjectID &object_id) {
|
||||
pull_requests_.erase(object_id);
|
||||
ray::Status status_code = object_directory_->GetLocations(
|
||||
object_id,
|
||||
[this](const std::vector<RemoteConnectionInfo> &vec, const ObjectID &object_id) {
|
||||
return GetLocationsSuccess(vec, object_id);
|
||||
},
|
||||
[this](ray::Status status, const ObjectID &object_id) {
|
||||
return GetLocationsFailed(status, object_id);
|
||||
});
|
||||
return status_code;
|
||||
}
|
||||
|
||||
void ObjectManager::GetLocationsSuccess(const std::vector<ray::RemoteConnectionInfo> &vec,
|
||||
const ray::ObjectID &object_id) {
|
||||
RemoteConnectionInfo info = vec.front();
|
||||
pull_requests_.erase(object_id);
|
||||
ray::Status status_code = Pull(object_id, info.client_id);
|
||||
};
|
||||
|
||||
void ObjectManager::GetLocationsFailed(ray::Status status, const ObjectID &object_id) {
|
||||
SchedulePull(object_id, config_.pull_timeout_ms);
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::Pull(const ObjectID &object_id, const ClientID &client_id) {
|
||||
Status status =
|
||||
GetMsgConnection(client_id, [this, object_id](SenderConnection::pointer client) {
|
||||
Status status = ExecutePull(object_id, client);
|
||||
});
|
||||
return status;
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::ExecutePull(const ObjectID &object_id,
|
||||
SenderConnection::pointer conn) {
|
||||
size_t message_type = OMMessageType_PullRequest;
|
||||
boost::system::error_code error_code;
|
||||
boost::asio::write(conn->GetSocket(),
|
||||
boost::asio::buffer(&message_type, sizeof(message_type)),
|
||||
error_code);
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
auto message = CreatePullRequest(fbb, fbb.CreateString(client_id_.binary()),
|
||||
fbb.CreateString(object_id.binary()));
|
||||
fbb.Finish(message);
|
||||
size_t length = fbb.GetSize();
|
||||
std::vector<boost::asio::const_buffer> buffer;
|
||||
buffer.push_back(boost::asio::buffer(&length, sizeof(length)));
|
||||
buffer.push_back(boost::asio::buffer(fbb.GetBufferPointer(), length));
|
||||
boost::asio::write(conn->GetSocket(), buffer);
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::Push(const ObjectID &object_id, const ClientID &client_id) {
|
||||
ray::Status status;
|
||||
status =
|
||||
GetTransferConnection(client_id, [this, object_id](SenderConnection::pointer conn) {
|
||||
ray::Status status = QueuePush(object_id, conn);
|
||||
});
|
||||
return status;
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::Cancel(const ObjectID &object_id) {
|
||||
// TODO(hme): Account for pull timers.
|
||||
ray::Status status = object_directory_->Cancel(object_id);
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::Wait(const std::vector<ObjectID> &object_ids,
|
||||
uint64_t timeout_ms, int num_ready_objects,
|
||||
const WaitCallback &callback) {
|
||||
// TODO: Implement wait.
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::GetMsgConnection(
|
||||
const ClientID &client_id, std::function<void(SenderConnection::pointer)> callback) {
|
||||
ray::Status status = Status::OK();
|
||||
if (message_send_connections_.count(client_id) > 0) {
|
||||
callback(message_send_connections_[client_id]);
|
||||
} else {
|
||||
status = object_directory_->GetInformation(
|
||||
client_id,
|
||||
[this, callback](RemoteConnectionInfo info) {
|
||||
Status status = CreateMsgConnection(info, callback);
|
||||
},
|
||||
[this](const Status &status) {
|
||||
// TODO: deal with failure.
|
||||
});
|
||||
}
|
||||
return status;
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::CreateMsgConnection(
|
||||
const RemoteConnectionInfo &info,
|
||||
std::function<void(SenderConnection::pointer)> callback) {
|
||||
message_send_connections_.emplace(
|
||||
info.client_id, SenderConnection::Create(io_service_, info.ip, info.port));
|
||||
// Prepare client connection info buffer.
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
bool is_transfer = false;
|
||||
auto message =
|
||||
CreateClientConnectionInfo(fbb, fbb.CreateString(client_id_.binary()), is_transfer);
|
||||
fbb.Finish(message);
|
||||
// Pack into asio buffer.
|
||||
size_t length = fbb.GetSize();
|
||||
std::vector<boost::asio::const_buffer> buffer;
|
||||
buffer.push_back(boost::asio::buffer(&length, sizeof(length)));
|
||||
buffer.push_back(boost::asio::buffer(fbb.GetBufferPointer(), length));
|
||||
// Send synchronously.
|
||||
SenderConnection::pointer conn = message_send_connections_[info.client_id];
|
||||
boost::system::error_code error;
|
||||
boost::asio::write(conn->GetSocket(), buffer);
|
||||
// The connection is ready, invoke callback with connection info.
|
||||
callback(message_send_connections_[info.client_id]);
|
||||
return Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::GetTransferConnection(
|
||||
const ClientID &client_id, std::function<void(SenderConnection::pointer)> callback) {
|
||||
ray::Status status = Status::OK();
|
||||
if (transfer_send_connections_.count(client_id) > 0) {
|
||||
callback(transfer_send_connections_[client_id]);
|
||||
} else {
|
||||
status = object_directory_->GetInformation(
|
||||
client_id,
|
||||
[this, callback](RemoteConnectionInfo info) {
|
||||
Status status = CreateTransferConnection(info, callback);
|
||||
},
|
||||
[this](const Status &status) {
|
||||
// TODO(hme): deal with failure.
|
||||
});
|
||||
}
|
||||
return status;
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::CreateTransferConnection(
|
||||
const RemoteConnectionInfo &info,
|
||||
std::function<void(SenderConnection::pointer)> callback) {
|
||||
transfer_send_connections_.emplace(
|
||||
info.client_id, SenderConnection::Create(io_service_, info.ip, info.port));
|
||||
// Prepare client connection info buffer.
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
bool is_transfer = true;
|
||||
auto message =
|
||||
CreateClientConnectionInfo(fbb, fbb.CreateString(client_id_.binary()), is_transfer);
|
||||
fbb.Finish(message);
|
||||
// Pack into asio buffer.
|
||||
size_t length = fbb.GetSize();
|
||||
std::vector<boost::asio::const_buffer> buffer;
|
||||
buffer.push_back(boost::asio::buffer(&length, sizeof(length)));
|
||||
buffer.push_back(boost::asio::buffer(fbb.GetBufferPointer(), length));
|
||||
// Send synchronously.
|
||||
SenderConnection::pointer conn = transfer_send_connections_[info.client_id];
|
||||
boost::system::error_code ec;
|
||||
boost::asio::write(conn->GetSocket(), buffer, ec);
|
||||
callback(transfer_send_connections_[info.client_id]);
|
||||
return Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::AcceptConnection(TCPClientConnection::pointer conn) {
|
||||
boost::system::error_code ec;
|
||||
// read header
|
||||
size_t length;
|
||||
std::vector<boost::asio::mutable_buffer> header;
|
||||
header.push_back(boost::asio::buffer(&length, sizeof(length)));
|
||||
boost::asio::read(conn->GetSocket(), header, ec);
|
||||
// read data
|
||||
std::vector<uint8_t> message;
|
||||
message.resize(length);
|
||||
boost::asio::read(conn->GetSocket(), boost::asio::buffer(message), ec);
|
||||
// Serialize
|
||||
auto info = flatbuffers::GetRoot<ClientConnectionInfo>(message.data());
|
||||
ClientID client_id = ObjectID::from_binary(info->client_id()->str());
|
||||
bool is_transfer = info->is_transfer();
|
||||
// TODO: trash connection if either fails.
|
||||
if (is_transfer) {
|
||||
transfer_receive_connections_[client_id] = conn;
|
||||
Status status = WaitPushReceive(conn);
|
||||
return status;
|
||||
} else {
|
||||
message_receive_connections_[client_id] = conn;
|
||||
Status status = WaitMessage(conn);
|
||||
return status;
|
||||
}
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::WaitPushReceive(TCPClientConnection::pointer conn) {
|
||||
boost::asio::async_read(
|
||||
conn->GetSocket(),
|
||||
boost::asio::buffer(&conn->message_length_, sizeof(conn->message_length_)),
|
||||
boost::bind(&ObjectManager::HandlePushReceive, this, conn,
|
||||
boost::asio::placeholders::error));
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
void ObjectManager::HandlePushReceive(TCPClientConnection::pointer conn,
|
||||
BoostEC length_ec) {
|
||||
std::vector<uint8_t> message;
|
||||
message.resize(conn->message_length_);
|
||||
boost::system::error_code ec;
|
||||
boost::asio::read(conn->GetSocket(), boost::asio::buffer(message), ec);
|
||||
// Serialize.
|
||||
auto object_header = flatbuffers::GetRoot<ObjectHeader>(message.data());
|
||||
ObjectID object_id = ObjectID::from_binary(object_header->object_id()->str());
|
||||
int64_t object_size = (int64_t)object_header->object_size();
|
||||
int64_t metadata_size = 0;
|
||||
// Try to create shared buffer.
|
||||
std::shared_ptr<Buffer> data;
|
||||
arrow::Status s = store_client_->GetClient().Create(
|
||||
object_id.to_plasma_id(), object_size, NULL, metadata_size, &data);
|
||||
if (s.ok()) {
|
||||
// Read object into store.
|
||||
uint8_t *mutable_data = data->mutable_data();
|
||||
boost::asio::read(conn->GetSocket(), boost::asio::buffer(mutable_data, object_size),
|
||||
ec);
|
||||
if (!ec.value()) {
|
||||
ARROW_CHECK_OK(store_client_->GetClient().Seal(object_id.to_plasma_id()));
|
||||
ARROW_CHECK_OK(store_client_->GetClient().Release(object_id.to_plasma_id()));
|
||||
} else {
|
||||
ARROW_CHECK_OK(store_client_->GetClient().Release(object_id.to_plasma_id()));
|
||||
ARROW_CHECK_OK(store_client_->GetClient().Abort(object_id.to_plasma_id()));
|
||||
RAY_LOG(ERROR) << "Receive Failed";
|
||||
}
|
||||
} else {
|
||||
RAY_LOG(ERROR) << "Buffer Create Failed: " << s.message();
|
||||
// Read object into empty buffer.
|
||||
uint8_t *mutable_data = (uint8_t *)malloc(object_size + metadata_size);
|
||||
boost::asio::read(conn->GetSocket(), boost::asio::buffer(mutable_data, object_size),
|
||||
ec);
|
||||
}
|
||||
// Wait for another push.
|
||||
ray::Status ray_status = WaitPushReceive(conn);
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::QueuePush(const ObjectID &object_id_const,
|
||||
SenderConnection::pointer conn) {
|
||||
ObjectID object_id = ObjectID(object_id_const);
|
||||
if (conn->ObjectIdQueued(object_id)) {
|
||||
// For now, return with status OK if the object is already in the send queue.
|
||||
return ray::Status::OK();
|
||||
}
|
||||
conn->QueueObjectId(object_id);
|
||||
if (num_transfers_ < max_transfers_) {
|
||||
return ExecutePushQueue(conn);
|
||||
}
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::ExecutePushQueue(SenderConnection::pointer conn) {
|
||||
ray::Status status = ray::Status::OK();
|
||||
while (num_transfers_ < max_transfers_) {
|
||||
if (conn->IsObjectIdQueueEmpty()) {
|
||||
return ray::Status::OK();
|
||||
}
|
||||
ObjectID object_id = conn->DequeueObjectId();
|
||||
// The threads that increment/decrement num_transfers_ are different.
|
||||
// It's important to increment num_transfers_ before executing the push.
|
||||
num_transfers_ += 1;
|
||||
status = ExecutePushHeaders(object_id, conn);
|
||||
}
|
||||
return status;
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::ExecutePushHeaders(const ObjectID &object_id_const,
|
||||
SenderConnection::pointer conn) {
|
||||
ObjectID object_id = ObjectID(object_id_const);
|
||||
// Allocate and append the request to the transfer queue.
|
||||
plasma::ObjectBuffer object_buffer;
|
||||
plasma::ObjectID plasma_id = object_id.to_plasma_id();
|
||||
ARROW_CHECK_OK(store_client_->GetClientOther().Get(&plasma_id, 1, 0, &object_buffer));
|
||||
if (object_buffer.data_size == -1) {
|
||||
RAY_LOG(ERROR) << "Failed to get object";
|
||||
// If the object wasn't locally available, exit immediately. If the object
|
||||
// later appears locally, the requesting plasma manager should request the
|
||||
// transfer again.
|
||||
return ray::Status::IOError(
|
||||
"Unable to transfer object to requesting plasma manager, object not local.");
|
||||
}
|
||||
RAY_CHECK(object_buffer.metadata->data() ==
|
||||
object_buffer.data->data() + object_buffer.data_size);
|
||||
SendRequest send_request;
|
||||
send_request.object_id = object_id;
|
||||
send_request.object_size = object_buffer.data_size;
|
||||
send_request.data = const_cast<uint8_t *>(object_buffer.data->data());
|
||||
conn->AddSendRequest(object_id, send_request);
|
||||
// Create buffer.
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
auto message = CreateObjectHeader(fbb, fbb.CreateString(object_id.binary()),
|
||||
send_request.object_size);
|
||||
fbb.Finish(message);
|
||||
// Pack into asio buffer.
|
||||
size_t length = fbb.GetSize();
|
||||
std::vector<boost::asio::const_buffer> buffer;
|
||||
buffer.push_back(boost::asio::buffer(&length, sizeof(length)));
|
||||
buffer.push_back(boost::asio::buffer(fbb.GetBufferPointer(), length));
|
||||
// Send asynchronously.
|
||||
boost::asio::async_write(conn->GetSocket(), buffer,
|
||||
boost::bind(&ObjectManager::ExecutePushObject, this, conn,
|
||||
object_id, boost::asio::placeholders::error));
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
void ObjectManager::ExecutePushObject(SenderConnection::pointer conn,
|
||||
const ObjectID &object_id,
|
||||
const boost::system::error_code &header_ec) {
|
||||
SendRequest &send_request = conn->GetSendRequest(object_id);
|
||||
boost::system::error_code ec;
|
||||
boost::asio::write(
|
||||
conn->GetSocket(),
|
||||
boost::asio::buffer(send_request.data, (size_t)send_request.object_size), ec);
|
||||
// Do this regardless of whether it failed or succeeded.
|
||||
ARROW_CHECK_OK(
|
||||
store_client_->GetClientOther().Release(send_request.object_id.to_plasma_id()));
|
||||
|
||||
ray::Status ray_status = ExecutePushCompleted(object_id, conn);
|
||||
}
|
||||
|
||||
ray::Status ObjectManager::ExecutePushCompleted(const ObjectID &object_id,
|
||||
SenderConnection::pointer conn) {
|
||||
conn->RemoveSendRequest(object_id);
|
||||
num_transfers_ -= 1;
|
||||
return ExecutePushQueue(conn);
|
||||
};
|
||||
|
||||
ray::Status ObjectManager::WaitMessage(TCPClientConnection::pointer conn) {
|
||||
boost::asio::async_read(
|
||||
conn->GetSocket(),
|
||||
boost::asio::buffer(&conn->message_type_, sizeof(conn->message_type_)),
|
||||
boost::bind(&ObjectManager::HandleMessage, this, conn,
|
||||
boost::asio::placeholders::error));
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
void ObjectManager::HandleMessage(TCPClientConnection::pointer conn, BoostEC msg_ec) {
|
||||
switch (conn->message_type_) {
|
||||
case OMMessageType_PullRequest:
|
||||
ReceivePullRequest(conn);
|
||||
}
|
||||
}
|
||||
|
||||
void ObjectManager::ReceivePullRequest(TCPClientConnection::pointer conn) {
|
||||
boost::asio::read(
|
||||
conn->GetSocket(),
|
||||
boost::asio::buffer(&conn->message_length_, sizeof(conn->message_length_)));
|
||||
std::vector<uint8_t> message;
|
||||
message.resize(conn->message_length_);
|
||||
boost::system::error_code error_code;
|
||||
boost::asio::read(conn->GetSocket(), boost::asio::buffer(message), error_code);
|
||||
// Serialize.
|
||||
auto pull_request = flatbuffers::GetRoot<PullRequest>(message.data());
|
||||
ObjectID object_id = ObjectID::from_binary(pull_request->object_id()->str());
|
||||
ClientID client_id = ClientID::from_binary(pull_request->client_id()->str());
|
||||
// Push object to requesting client.
|
||||
ray::Status push_status = Push(object_id, client_id);
|
||||
ray::Status wait_status = WaitMessage(conn);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,268 @@
|
||||
#ifndef RAY_OBJECT_MANAGER_OBJECT_MANAGER_H
|
||||
#define RAY_OBJECT_MANAGER_OBJECT_MANAGER_H
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstdint>
|
||||
#include <deque>
|
||||
#include <map>
|
||||
#include <memory>
|
||||
#include <thread>
|
||||
|
||||
#include <boost/asio.hpp>
|
||||
#include <boost/asio/error.hpp>
|
||||
#include <boost/bind.hpp>
|
||||
|
||||
#include "plasma/client.h"
|
||||
#include "plasma/events.h"
|
||||
#include "plasma/plasma.h"
|
||||
|
||||
#include "format/object_manager_generated.h"
|
||||
#include "object_directory.h"
|
||||
#include "object_manager_client_connection.h"
|
||||
#include "object_store_client.h"
|
||||
#include "ray/id.h"
|
||||
#include "ray/status.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
struct ObjectManagerConfig {
|
||||
// The time in milliseconds to wait before retrying a pull
|
||||
// that failed due to client id lookup.
|
||||
int pull_timeout_ms = 100;
|
||||
// TODO(hme): Implement num retries (to avoid infinite retries).
|
||||
std::string store_socket_name;
|
||||
};
|
||||
|
||||
// TODO(hme): Comment everything doxygen-style.
|
||||
// TODO(hme): Implement connection cleanup.
|
||||
// TODO(hme): Add success/failure callbacks for push and pull.
|
||||
// TODO(hme): Use boost thread pool.
|
||||
// TODO(hme): Add incoming connections to io_service tied to thread pool.
|
||||
class ObjectManager {
|
||||
public:
|
||||
/// Implicitly instantiates Ray implementation of ObjectDirectory.
|
||||
///
|
||||
/// \param io_service The asio io_service tied to the object manager.
|
||||
/// \param config ObjectManager configuration.
|
||||
/// \param gcs_client A client connection to the Ray GCS.
|
||||
explicit ObjectManager(boost::asio::io_service &io_service, ObjectManagerConfig config,
|
||||
std::shared_ptr<ray::GcsClient> gcs_client);
|
||||
|
||||
/// Takes user-defined ObjectDirectoryInterface implementation.
|
||||
/// When this constructor is used, the ObjectManager assumes ownership of
|
||||
/// the given ObjectDirectory instance.
|
||||
///
|
||||
/// \param io_service The asio io_service tied to the object manager.
|
||||
/// \param config ObjectManager configuration.
|
||||
/// \param od An object implementing the object directory interface.
|
||||
explicit ObjectManager(boost::asio::io_service &io_service, ObjectManagerConfig config,
|
||||
std::unique_ptr<ObjectDirectoryInterface> od);
|
||||
|
||||
/// \param client_id Set the client id associated with this node.
|
||||
void SetClientID(const ClientID &client_id);
|
||||
|
||||
/// \return Get the client id associated with this node.
|
||||
ClientID GetClientID();
|
||||
|
||||
/// Subscribe to notifications of objects added to local store.
|
||||
/// Upon subscribing, the callback will be invoked for all objects that
|
||||
///
|
||||
/// already exist in the local store.
|
||||
/// \param callback The callback to invoke when objects are added to the local store.
|
||||
/// \return Status of whether adding the subscription succeeded.
|
||||
ray::Status SubscribeObjAdded(std::function<void(const ray::ObjectID &)> callback);
|
||||
|
||||
/// Subscribe to notifications of objects deleted from local store.
|
||||
///
|
||||
/// \param callback The callback to invoke when objects are removed from the local
|
||||
/// store.
|
||||
/// \return Status of whether adding the subscription succeeded.
|
||||
ray::Status SubscribeObjDeleted(std::function<void(const ray::ObjectID &)> callback);
|
||||
|
||||
/// Push an object to to the node manager on the node corresponding to client id.
|
||||
///
|
||||
/// \param object_id The object's object id.
|
||||
/// \param client_id The remote node's client id.
|
||||
/// \return Status of whether the push request successfully initiated.
|
||||
ray::Status Push(const ObjectID &object_id, const ClientID &client_id);
|
||||
|
||||
/// Pull an object from ClientID. Returns UniqueID asociated with
|
||||
/// an invocation of this method.
|
||||
///
|
||||
/// \param object_id The object's object id.
|
||||
/// \return Status of whether the pull request successfully initiated.
|
||||
ray::Status Pull(const ObjectID &object_id);
|
||||
|
||||
/// Discover ClientID via ObjectDirectory, then pull object
|
||||
/// from ClientID associated with ObjectID.
|
||||
///
|
||||
/// \param object_id The object's object id.
|
||||
/// \param client_id The remote node's client id.
|
||||
/// \return Status of whether the pull request successfully initiated.
|
||||
ray::Status Pull(const ObjectID &object_id, const ClientID &client_id);
|
||||
|
||||
/// Add a connection to a remote object manager.
|
||||
/// This is invoked by an external server.
|
||||
///
|
||||
/// \param conn The connection.
|
||||
/// \return Status of whether the connection was successfully established.
|
||||
ray::Status AcceptConnection(TCPClientConnection::pointer conn);
|
||||
|
||||
/// Cancels all requests (Push/Pull) associated with the given ObjectID.
|
||||
///
|
||||
/// \param object_id The ObjectID.
|
||||
/// \return Status of whether requests were successfully cancelled.
|
||||
ray::Status Cancel(const ObjectID &object_id);
|
||||
|
||||
// Callback definition for wait.
|
||||
using WaitCallback = std::function<void(const ray::Status, uint64_t,
|
||||
const std::vector<ray::ObjectID> &)>;
|
||||
/// Wait for timeout_ms before invoking the provided callback.
|
||||
/// If num_ready_objects is satisfied before the timeout, then
|
||||
/// invoke the callback.
|
||||
///
|
||||
/// \param object_ids The object ids to wait on.
|
||||
/// \param timeout_ms The time in milliseconds to wait before invoking the callback.
|
||||
/// \param num_ready_objects The minimum number of objects required before
|
||||
/// invoking the callback.
|
||||
/// \param callback Invoked when either timeout_ms is satisfied OR num_ready_objects
|
||||
/// is satisfied.
|
||||
/// \return Status of whether the wait successfully initiated.
|
||||
ray::Status Wait(const std::vector<ObjectID> &object_ids, uint64_t timeout_ms,
|
||||
int num_ready_objects, const WaitCallback &callback);
|
||||
|
||||
/// \return Whether this object was successfully terminated.
|
||||
ray::Status Terminate();
|
||||
|
||||
private:
|
||||
using BoostEC = const boost::system::error_code &;
|
||||
|
||||
ClientID client_id_;
|
||||
ObjectManagerConfig config_;
|
||||
std::unique_ptr<ObjectDirectoryInterface> object_directory_;
|
||||
std::unique_ptr<ObjectStoreClient> store_client_;
|
||||
|
||||
/// An io service for creating connections to other object managers.
|
||||
boost::asio::io_service io_service_;
|
||||
|
||||
/// Used to create "work" for an io service, so when it's run, it doesn't exit.
|
||||
boost::asio::io_service::work work_;
|
||||
|
||||
/// Single thread for executing asynchronous handlers.
|
||||
/// This runs the (currently only) io_service, which handles all outgoing requests
|
||||
/// and object transfers (push).
|
||||
std::thread io_thread_;
|
||||
|
||||
/// Relatively simple way to add thread pooling.
|
||||
/// boost::thread_group thread_group_;
|
||||
|
||||
/// Timeout for failed pull requests.
|
||||
using Timer = std::shared_ptr<boost::asio::deadline_timer>;
|
||||
std::unordered_map<ObjectID, Timer, UniqueIDHasher> pull_requests_;
|
||||
|
||||
// TODO (hme): This needs to account for receives as well.
|
||||
/// This number is incremented whenever a push is started.
|
||||
int num_transfers_ = 0;
|
||||
// TODO (hme): Allow for concurrent sends.
|
||||
/// This is the maximum number of pushes allowed.
|
||||
/// We can only increase this number if we increase the number of
|
||||
/// plasma client connections.
|
||||
int max_transfers_ = 1;
|
||||
|
||||
/// Note that (currently) receives take place on the main thread,
|
||||
/// and sends take place on a dedicated thread.
|
||||
std::unordered_map<ray::ClientID, SenderConnection::pointer, ray::UniqueIDHasher>
|
||||
message_send_connections_;
|
||||
std::unordered_map<ray::ClientID, SenderConnection::pointer, ray::UniqueIDHasher>
|
||||
transfer_send_connections_;
|
||||
|
||||
std::unordered_map<ray::ClientID, TCPClientConnection::pointer, ray::UniqueIDHasher>
|
||||
message_receive_connections_;
|
||||
std::unordered_map<ray::ClientID, TCPClientConnection::pointer, ray::UniqueIDHasher>
|
||||
transfer_receive_connections_;
|
||||
|
||||
/// Handle starting, running, and stopping asio io_service.
|
||||
void StartIOService();
|
||||
void IOServiceLoop();
|
||||
void StopIOService();
|
||||
|
||||
/// Wait wait_ms milliseconds before triggering a pull request for object_id.
|
||||
/// This is invoked when a pull fails. Only point of failure currently considered
|
||||
/// is GetLocationsFailed.
|
||||
void SchedulePull(const ObjectID &object_id, int wait_ms);
|
||||
|
||||
/// The handler for SchedulePull. Invokes a pull and removes the deadline timer
|
||||
/// that was added to schedule the pull.
|
||||
ray::Status SchedulePullHandler(const ObjectID &object_id);
|
||||
|
||||
/// Synchronously send a pull request.
|
||||
/// Invoked once a connection to a remote manager that contains the required ObjectID
|
||||
/// is established.
|
||||
ray::Status ExecutePull(const ObjectID &object_id, SenderConnection::pointer conn);
|
||||
|
||||
/// Invoked once a connection to the remote manager to which the ObjectID
|
||||
/// is to be sent is established.
|
||||
ray::Status QueuePush(const ObjectID &object_id, SenderConnection::pointer client);
|
||||
/// Starts as many queued pushes as possible without exceeding max_transfers_
|
||||
/// concurrent transfers.
|
||||
ray::Status ExecutePushQueue(SenderConnection::pointer client);
|
||||
/// Initiate a push. This method asynchronously sends the object id and object size
|
||||
/// to the remote object manager.
|
||||
ray::Status ExecutePushHeaders(const ObjectID &object_id,
|
||||
SenderConnection::pointer client);
|
||||
/// Called by the handler for ExecutePushMeta.
|
||||
/// This method initiates the actual object transfer.
|
||||
void ExecutePushObject(SenderConnection::pointer conn, const ObjectID &object_id,
|
||||
const boost::system::error_code &header_ec);
|
||||
/// Invoked when a push is completed. This method will decrement num_transfers_
|
||||
/// and invoke ExecutePushQueue.
|
||||
ray::Status ExecutePushCompleted(const ObjectID &object_id,
|
||||
SenderConnection::pointer client);
|
||||
|
||||
/// Private callback implementation for success on get location. Called inside OD.
|
||||
void GetLocationsSuccess(const std::vector<RemoteConnectionInfo> &vec,
|
||||
const ObjectID &object_id);
|
||||
|
||||
/// Private callback implementation for failure on get location. Called inside OD.
|
||||
void GetLocationsFailed(ray::Status status, const ObjectID &object_id);
|
||||
|
||||
/// Asynchronously obtain a connection to client_id.
|
||||
/// If a connection to client_id already exists, the callback is invoked immediately.
|
||||
ray::Status GetMsgConnection(const ClientID &client_id,
|
||||
std::function<void(SenderConnection::pointer)> callback);
|
||||
/// Asynchronously create a connection to client_id.
|
||||
ray::Status CreateMsgConnection(
|
||||
const RemoteConnectionInfo &info,
|
||||
std::function<void(SenderConnection::pointer)> callback);
|
||||
/// Asynchronously create a connection to client_id.
|
||||
ray::Status GetTransferConnection(
|
||||
const ClientID &client_id, std::function<void(SenderConnection::pointer)> callback);
|
||||
/// Asynchronously obtain a connection to client_id.
|
||||
/// If a connection to client_id already exists, the callback is invoked immediately.
|
||||
ray::Status CreateTransferConnection(
|
||||
const RemoteConnectionInfo &info,
|
||||
std::function<void(SenderConnection::pointer)> callback);
|
||||
|
||||
/// A socket connection doing an asynchronous read on a transfer connection that was
|
||||
/// added by AcceptConnection.
|
||||
ray::Status WaitPushReceive(TCPClientConnection::pointer conn);
|
||||
/// Invoked when a remote object manager pushes an object to this object manager.
|
||||
void HandlePushReceive(TCPClientConnection::pointer conn, BoostEC length_ec);
|
||||
|
||||
/// A socket connection doing an asynchronous read on a message connection that was
|
||||
/// added by AcceptConnection.
|
||||
ray::Status WaitMessage(TCPClientConnection::pointer conn);
|
||||
/// Handle messages.
|
||||
void HandleMessage(TCPClientConnection::pointer conn, BoostEC msg_ec);
|
||||
/// Process the receive pull request message.
|
||||
void ReceivePullRequest(TCPClientConnection::pointer conn);
|
||||
|
||||
/// Register object add with directory.
|
||||
void NotifyDirectoryObjectAdd(const ObjectID &object_id);
|
||||
/// Register object remove with directory.
|
||||
void NotifyDirectoryObjectDeleted(const ObjectID &object_id);
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_OBJECT_MANAGER_OBJECT_MANAGER_H
|
||||
@@ -0,0 +1,59 @@
|
||||
#include "object_manager_client_connection.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
SenderConnection::pointer SenderConnection::Create(boost::asio::io_service &io_service,
|
||||
const std::string &ip, uint16_t port) {
|
||||
return pointer(new SenderConnection(io_service, ip, port));
|
||||
};
|
||||
|
||||
SenderConnection::SenderConnection(boost::asio::io_service &io_service,
|
||||
const std::string &ip, uint16_t port)
|
||||
: socket_(io_service), send_queue_() {
|
||||
boost::asio::ip::address ip_address = boost::asio::ip::address::from_string(ip);
|
||||
boost::asio::ip::tcp::endpoint endpoint(ip_address, port);
|
||||
socket_.connect(endpoint);
|
||||
};
|
||||
|
||||
boost::asio::ip::tcp::socket &SenderConnection::GetSocket() { return socket_; };
|
||||
|
||||
bool SenderConnection::IsObjectIdQueueEmpty() { return send_queue_.empty(); }
|
||||
|
||||
bool SenderConnection::ObjectIdQueued(const ObjectID &object_id) {
|
||||
return std::find(send_queue_.begin(), send_queue_.end(), object_id) !=
|
||||
send_queue_.end();
|
||||
}
|
||||
|
||||
void SenderConnection::QueueObjectId(const ObjectID &object_id) {
|
||||
send_queue_.push_back(ObjectID(object_id));
|
||||
}
|
||||
|
||||
ObjectID SenderConnection::DequeueObjectId() {
|
||||
ObjectID object_id = send_queue_.front();
|
||||
send_queue_.pop_front();
|
||||
return object_id;
|
||||
}
|
||||
|
||||
void SenderConnection::AddSendRequest(const ObjectID &object_id,
|
||||
SendRequest &send_request) {
|
||||
send_requests_.emplace(object_id, send_request);
|
||||
}
|
||||
|
||||
void SenderConnection::RemoveSendRequest(const ObjectID &object_id) {
|
||||
send_requests_.erase(object_id);
|
||||
}
|
||||
|
||||
SendRequest &SenderConnection::GetSendRequest(const ObjectID &object_id) {
|
||||
return send_requests_[object_id];
|
||||
};
|
||||
|
||||
TCPClientConnection::TCPClientConnection(boost::asio::io_service &io_service)
|
||||
: socket_(io_service) {}
|
||||
|
||||
TCPClientConnection::pointer TCPClientConnection::Create(
|
||||
boost::asio::io_service &io_service) {
|
||||
return TCPClientConnection::pointer(new TCPClientConnection(io_service));
|
||||
}
|
||||
|
||||
boost::asio::ip::tcp::socket &TCPClientConnection::GetSocket() { return socket_; }
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,71 @@
|
||||
#ifndef RAY_OBJECT_MANAGER_OBJECT_MANAGER_CLIENT_CONNECTION_H
|
||||
#define RAY_OBJECT_MANAGER_OBJECT_MANAGER_CLIENT_CONNECTION_H
|
||||
|
||||
#include <deque>
|
||||
#include <memory>
|
||||
#include <unordered_map>
|
||||
|
||||
#include <boost/asio.hpp>
|
||||
#include <boost/asio/error.hpp>
|
||||
#include <boost/enable_shared_from_this.hpp>
|
||||
|
||||
#include "ray/id.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
struct SendRequest {
|
||||
ObjectID object_id;
|
||||
ClientID client_id;
|
||||
int64_t object_size;
|
||||
uint8_t *data;
|
||||
};
|
||||
|
||||
// TODO(hme): Document public API after integration with common connection.
|
||||
class SenderConnection : public boost::enable_shared_from_this<SenderConnection> {
|
||||
public:
|
||||
typedef boost::shared_ptr<SenderConnection> pointer;
|
||||
typedef std::unordered_map<ray::ObjectID, SendRequest, UniqueIDHasher> SendRequestsType;
|
||||
typedef std::deque<ray::ObjectID> SendQueueType;
|
||||
|
||||
static pointer Create(boost::asio::io_service &io_service, const std::string &ip,
|
||||
uint16_t port);
|
||||
|
||||
explicit SenderConnection(boost::asio::io_service &io_service, const std::string &ip,
|
||||
uint16_t port);
|
||||
|
||||
boost::asio::ip::tcp::socket &GetSocket();
|
||||
|
||||
bool IsObjectIdQueueEmpty();
|
||||
bool ObjectIdQueued(const ObjectID &object_id);
|
||||
void QueueObjectId(const ObjectID &object_id);
|
||||
ObjectID DequeueObjectId();
|
||||
|
||||
void AddSendRequest(const ObjectID &object_id, SendRequest &send_request);
|
||||
void RemoveSendRequest(const ObjectID &object_id);
|
||||
SendRequest &GetSendRequest(const ObjectID &object_id);
|
||||
|
||||
private:
|
||||
boost::asio::ip::tcp::socket socket_;
|
||||
SendQueueType send_queue_;
|
||||
SendRequestsType send_requests_;
|
||||
};
|
||||
|
||||
// TODO(hme): Document public API after integration with common connection.
|
||||
class TCPClientConnection : public boost::enable_shared_from_this<TCPClientConnection> {
|
||||
public:
|
||||
typedef boost::shared_ptr<TCPClientConnection> pointer;
|
||||
static pointer Create(boost::asio::io_service &io_service);
|
||||
boost::asio::ip::tcp::socket &GetSocket();
|
||||
|
||||
TCPClientConnection(boost::asio::io_service &io_service);
|
||||
|
||||
int64_t message_type_;
|
||||
uint64_t message_length_;
|
||||
|
||||
private:
|
||||
boost::asio::ip::tcp::socket socket_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_OBJECT_MANAGER_OBJECT_MANAGER_CLIENT_CONNECTION_H
|
||||
@@ -0,0 +1 @@
|
||||
// TODO(hme): Move all messaging code here.
|
||||
@@ -0,0 +1,6 @@
|
||||
#ifndef RAY_OBJECT_MANAGER_OBJECT_MANAGER_PROTOCOL_H
|
||||
#define RAY_OBJECT_MANAGER_OBJECT_MANAGER_PROTOCOL_H
|
||||
|
||||
// TODO(hme): Move all messaging code here.
|
||||
|
||||
#endif // RAY_OBJECT_MANAGER_OBJECT_MANAGER_PROTOCOL_H
|
||||
@@ -0,0 +1,137 @@
|
||||
#include <iostream>
|
||||
#include <thread>
|
||||
|
||||
#include "gtest/gtest.h"
|
||||
#include "plasma/client.h"
|
||||
#include "plasma/events.h"
|
||||
#include "plasma/plasma.h"
|
||||
#include "plasma/protocol.h"
|
||||
|
||||
#include "ray/status.h"
|
||||
|
||||
#include "ray/object_manager/object_manager.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
std::string test_executable; // NOLINT
|
||||
|
||||
class TestObjectManager : public ::testing::Test {
|
||||
public:
|
||||
TestObjectManager() { RAY_LOG(DEBUG) << "TestObjectManager: started."; }
|
||||
|
||||
void SetUp() {
|
||||
// start store
|
||||
std::string om_dir = test_executable.substr(0, test_executable.find_last_of("/"));
|
||||
std::string plasma_dir = om_dir + "./../plasma";
|
||||
std::string plasma_command =
|
||||
plasma_dir +
|
||||
"/plasma_store -m 1000000000 -s /tmp/store 1> /dev/null 2> /dev/null &";
|
||||
int s = system(plasma_command.c_str());
|
||||
ASSERT_TRUE(!s);
|
||||
|
||||
// Start mock global control store.
|
||||
mock_gcs_client_ = std::shared_ptr<GcsClient>(new GcsClient());
|
||||
// mock_gcs_client_->Register();
|
||||
|
||||
// Start node server.
|
||||
|
||||
// Start object manager 1.
|
||||
ObjectManagerConfig config;
|
||||
config.store_socket_name = "/tmp/store";
|
||||
object_manager_1_ = std::unique_ptr<ObjectManager>(
|
||||
new ObjectManager(io_service_, config, mock_gcs_client_));
|
||||
|
||||
// Start object manager 2.
|
||||
// ObjectManagerConfig config2;
|
||||
// config2.store_socket_name = "/tmp/store";
|
||||
// std::shared_ptr<ObjectDirectory> od2 = std::shared_ptr<ObjectDirectory>(new
|
||||
// ObjectDirectory());
|
||||
// od2->InitGcs(mock_gcs_client_);
|
||||
// object_manager_2_ = std::unique_ptr<ObjectManager>(new ObjectManager(io_service,
|
||||
// config2, od2));
|
||||
|
||||
// Initiate client connection.
|
||||
ARROW_CHECK_OK(client_.Connect("/tmp/store", "", PLASMA_DEFAULT_RELEASE_DELAY));
|
||||
|
||||
this->StartLoop();
|
||||
}
|
||||
|
||||
void TearDown() {
|
||||
this->StopLoop();
|
||||
arrow::Status arrow_status = client_.Disconnect();
|
||||
ASSERT_TRUE(arrow_status.ok());
|
||||
ray::Status ray_status = object_manager_1_->Terminate();
|
||||
ASSERT_TRUE(ray_status.ok());
|
||||
// object_manager_2_->Terminate();
|
||||
int s = system("killall plasma_store &");
|
||||
ASSERT_TRUE(!s);
|
||||
}
|
||||
|
||||
void Loop() { io_service_.run(); };
|
||||
|
||||
void StartLoop() { process_thread_ = std::thread(&TestObjectManager::Loop, this); };
|
||||
|
||||
void StopLoop() {
|
||||
io_service_.stop();
|
||||
process_thread_.join();
|
||||
}
|
||||
|
||||
protected:
|
||||
std::thread process_thread_;
|
||||
plasma::PlasmaClient client_;
|
||||
plasma::PlasmaClient client2_;
|
||||
boost::asio::io_service io_service_;
|
||||
|
||||
std::shared_ptr<GcsClient> mock_gcs_client_;
|
||||
std::unique_ptr<ObjectManager> object_manager_1_;
|
||||
std::unique_ptr<ObjectManager> object_manager_2_;
|
||||
};
|
||||
|
||||
// TODO: get rid of dead code?
|
||||
// TEST_F(TestObjectManager, TestPush) {
|
||||
// // test object push between two object managers.
|
||||
// ASSERT_TRUE(true);
|
||||
// sleep(1);
|
||||
//}
|
||||
|
||||
// TEST_F(TestObjectManager, TestPull) {
|
||||
// ObjectID object_id = ObjectID().from_random();
|
||||
// ClientID dbc_id = ClientID().from_random();
|
||||
// RAY_LOG(INFO) << "ObjectID: " << object_id.hex().c_str();
|
||||
// RAY_LOG(INFO) << "ClientID: " << dbc_id.hex().c_str();
|
||||
// om->Pull(object_id, dbc_id);
|
||||
// om->Pull(object_id);
|
||||
// ASSERT_TRUE(true);
|
||||
// sleep(1);
|
||||
//}
|
||||
|
||||
void ObjectAdded(const ObjectID &object_id) {
|
||||
RAY_LOG(INFO) << "ObjectID Added: " << object_id.hex().c_str();
|
||||
}
|
||||
|
||||
TEST_F(TestObjectManager, TestNotifications) {
|
||||
ray::Status status = object_manager_1_->SubscribeObjAdded(ObjectAdded);
|
||||
ASSERT_TRUE(status.ok());
|
||||
// put object
|
||||
for (int i = 0; i < 10; ++i) {
|
||||
ObjectID object_id = ObjectID::from_random();
|
||||
RAY_LOG(INFO) << "ObjectID Created: " << object_id.hex().c_str();
|
||||
int64_t data_size = 100;
|
||||
uint8_t metadata[] = {5};
|
||||
int64_t metadata_size = sizeof(metadata);
|
||||
std::shared_ptr<Buffer> data;
|
||||
ARROW_CHECK_OK(client_.Create(object_id.to_plasma_id(), data_size, metadata,
|
||||
metadata_size, &data));
|
||||
ARROW_CHECK_OK(client_.Seal(object_id.to_plasma_id()));
|
||||
}
|
||||
// TODO(hme): Can we do this without sleeping?
|
||||
sleep(1);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
|
||||
int main(int argc, char **argv) {
|
||||
::testing::InitGoogleTest(&argc, argv);
|
||||
ray::test_executable = std::string(argv[0]);
|
||||
return RUN_ALL_TESTS();
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
#include <future>
|
||||
#include <iostream>
|
||||
|
||||
#include <boost/asio.hpp>
|
||||
#include <boost/bind.hpp>
|
||||
#include <boost/function.hpp>
|
||||
|
||||
#include "common.h"
|
||||
#include "common_protocol.h"
|
||||
#include "ray/object_manager/object_store_client.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
// TODO(hme): Dedicate this class to notifications.
|
||||
// TODO(hme): Create object store client pool for object manager.
|
||||
ObjectStoreClient::ObjectStoreClient(boost::asio::io_service &io_service,
|
||||
std::string &store_socket_name)
|
||||
: client_one_(), client_two_(), socket_(io_service) {
|
||||
ARROW_CHECK_OK(
|
||||
client_two_.Connect(store_socket_name.c_str(), "", PLASMA_DEFAULT_RELEASE_DELAY));
|
||||
ARROW_CHECK_OK(
|
||||
client_one_.Connect(store_socket_name.c_str(), "", PLASMA_DEFAULT_RELEASE_DELAY));
|
||||
|
||||
// Connect to two clients, but subscribe to only one.
|
||||
ARROW_CHECK_OK(client_one_.Subscribe(&c_socket_));
|
||||
boost::system::error_code ec;
|
||||
socket_.assign(boost::asio::local::stream_protocol(), c_socket_, ec);
|
||||
assert(!ec.value());
|
||||
NotificationWait();
|
||||
};
|
||||
|
||||
void ObjectStoreClient::Terminate() {
|
||||
ARROW_CHECK_OK(client_two_.Disconnect());
|
||||
ARROW_CHECK_OK(client_one_.Disconnect());
|
||||
}
|
||||
|
||||
void ObjectStoreClient::NotificationWait() {
|
||||
boost::asio::async_read(socket_, boost::asio::buffer(&length_, sizeof(length_)),
|
||||
boost::bind(&ObjectStoreClient::ProcessStoreLength, this,
|
||||
boost::asio::placeholders::error));
|
||||
}
|
||||
|
||||
void ObjectStoreClient::ProcessStoreLength(const boost::system::error_code &error) {
|
||||
notification_.resize(length_);
|
||||
boost::asio::async_read(socket_, boost::asio::buffer(notification_),
|
||||
boost::bind(&ObjectStoreClient::ProcessStoreNotification, this,
|
||||
boost::asio::placeholders::error));
|
||||
}
|
||||
|
||||
void ObjectStoreClient::ProcessStoreNotification(const boost::system::error_code &error) {
|
||||
if (error) {
|
||||
throw std::runtime_error("ObjectStore may have died.");
|
||||
}
|
||||
|
||||
auto object_info = flatbuffers::GetRoot<ObjectInfo>(notification_.data());
|
||||
ObjectID object_id = from_flatbuf(*object_info->object_id());
|
||||
if (object_info->is_deletion()) {
|
||||
ProcessStoreRemove(object_id);
|
||||
} else {
|
||||
ProcessStoreAdd(object_id);
|
||||
// why all these params?
|
||||
// ProcessStoreAdd(
|
||||
// object_id, object_info->data_size(),
|
||||
// object_info->metadata_size(),
|
||||
// (unsigned char *) object_info->digest()->data());
|
||||
}
|
||||
NotificationWait();
|
||||
}
|
||||
|
||||
void ObjectStoreClient::ProcessStoreAdd(const ObjectID &object_id) {
|
||||
for (auto handler : add_handlers_) {
|
||||
handler(object_id);
|
||||
}
|
||||
};
|
||||
|
||||
void ObjectStoreClient::ProcessStoreRemove(const ObjectID &object_id) {
|
||||
for (auto handler : rem_handlers_) {
|
||||
handler(object_id);
|
||||
}
|
||||
};
|
||||
|
||||
void ObjectStoreClient::SubscribeObjAdded(
|
||||
std::function<void(const ObjectID &)> callback) {
|
||||
add_handlers_.push_back(callback);
|
||||
};
|
||||
|
||||
void ObjectStoreClient::SubscribeObjDeleted(
|
||||
std::function<void(const ObjectID &)> callback) {
|
||||
rem_handlers_.push_back(callback);
|
||||
};
|
||||
|
||||
plasma::PlasmaClient &ObjectStoreClient::GetClient() { return client_one_; };
|
||||
|
||||
plasma::PlasmaClient &ObjectStoreClient::GetClientOther() { return client_two_; };
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,67 @@
|
||||
#ifndef RAY_OBJECT_MANAGER_OBJECT_STORE_CLIENT_H
|
||||
#define RAY_OBJECT_MANAGER_OBJECT_STORE_CLIENT_H
|
||||
|
||||
#include <list>
|
||||
#include <memory>
|
||||
#include <vector>
|
||||
|
||||
#include <boost/asio.hpp>
|
||||
#include <boost/asio/error.hpp>
|
||||
#include <boost/bind.hpp>
|
||||
|
||||
#include "plasma/client.h"
|
||||
#include "plasma/events.h"
|
||||
#include "plasma/plasma.h"
|
||||
|
||||
#include "object_directory.h"
|
||||
#include "ray/id.h"
|
||||
#include "ray/status.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
// TODO(hme): document public API after refactor.
|
||||
class ObjectStoreClient {
|
||||
public:
|
||||
// Encapsulates communication with the object store.
|
||||
ObjectStoreClient(boost::asio::io_service &io_service, std::string &store_socket_name);
|
||||
|
||||
// Subscribe to notifications of objects added to local store.
|
||||
// Upon subscribing, the callback will be invoked for all objects that
|
||||
// already exist in the local store.
|
||||
void SubscribeObjAdded(std::function<void(const ray::ObjectID &)> callback);
|
||||
|
||||
// Subscribe to notifications of objects deleted from local store.
|
||||
void SubscribeObjDeleted(std::function<void(const ray::ObjectID &)> callback);
|
||||
|
||||
// TODO(hme): There should be as many client connections as there are threads.
|
||||
// Two client connections are made to enable concurrent communication with the store.
|
||||
plasma::PlasmaClient &GetClient();
|
||||
plasma::PlasmaClient &GetClientOther();
|
||||
|
||||
// Terminate this object.
|
||||
void Terminate();
|
||||
|
||||
private:
|
||||
std::vector<std::function<void(const ray::ObjectID &)>> add_handlers_;
|
||||
std::vector<std::function<void(const ray::ObjectID &)>> rem_handlers_;
|
||||
|
||||
plasma::PlasmaClient client_one_;
|
||||
plasma::PlasmaClient client_two_;
|
||||
int c_socket_;
|
||||
int64_t length_;
|
||||
std::vector<uint8_t> notification_;
|
||||
boost::asio::local::stream_protocol::socket socket_;
|
||||
|
||||
// Async loop for handling object store notifications.
|
||||
void NotificationWait();
|
||||
void ProcessStoreLength(const boost::system::error_code &error);
|
||||
void ProcessStoreNotification(const boost::system::error_code &error);
|
||||
|
||||
// Support for rebroadcasting object add/rem events.
|
||||
void ProcessStoreAdd(const ObjectID &object_id);
|
||||
void ProcessStoreRemove(const ObjectID &object_id);
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_OBJECT_MANAGER_OBJECT_STORE_CLIENT_H
|
||||
@@ -0,0 +1,36 @@
|
||||
include(${CMAKE_CURRENT_LIST_DIR}/../../common/cmake/Common.cmake)
|
||||
|
||||
set(NODE_MANAGER_FBS_SRC "${CMAKE_CURRENT_LIST_DIR}/format/node_manager.fbs")
|
||||
set(OUTPUT_DIR ${CMAKE_CURRENT_LIST_DIR}/format/)
|
||||
|
||||
set(NODE_MANAGER_FBS_OUTPUT_FILES
|
||||
"${OUTPUT_DIR}/node_manager_generated.h")
|
||||
|
||||
add_custom_command(
|
||||
OUTPUT ${NODE_MANAGER_FBS_OUTPUT_FILES}
|
||||
# The --gen-object-api flag generates a C++ class MessageT for each
|
||||
# flatbuffers message Message, which can be used to store deserialized
|
||||
# messages in data structures. This is currently used for ObjectInfo for
|
||||
# example.
|
||||
COMMAND ${FLATBUFFERS_COMPILER} -c -o ${OUTPUT_DIR} ${NODE_MANAGER_FBS_SRC} --cpp --gen-object-api --gen-mutable
|
||||
DEPENDS ${FBS_DEPENDS}
|
||||
COMMENT "Running flatc compiler on ${NODE_MANAGER_FBS_SRC}"
|
||||
VERBATIM)
|
||||
|
||||
add_custom_target(gen_node_manager_fbs DEPENDS ${NODE_MANAGER_FBS_OUTPUT_FILES})
|
||||
|
||||
ADD_RAY_TEST(raylet_test STATIC_LINK_LIBS ray_static ${PLASMA_STATIC_LIB} ${ARROW_STATIC_LIB} gtest gtest_main pthread ${Boost_SYSTEM_LIBRARY})
|
||||
|
||||
ADD_RAY_TEST(worker_pool_test STATIC_LINK_LIBS ray_static ${PLASMA_STATIC_LIB} ${ARROW_STATIC_LIB} gtest gtest_main gmock_main pthread ${Boost_SYSTEM_LIBRARY})
|
||||
|
||||
add_library(rayletlib raylet.cc ${NODE_MANAGER_FBS_OUTPUT_FILES})
|
||||
target_link_libraries(rayletlib ray_static ${Boost_SYSTEM_LIBRARY})
|
||||
|
||||
add_executable(raylet main.cc)
|
||||
target_link_libraries(raylet rayletlib ${Boost_SYSTEM_LIBRARY} pthread)
|
||||
add_executable(raylet_demo remote_dependencies_demo.cc)
|
||||
target_link_libraries(raylet_demo rayletlib ${Boost_SYSTEM_LIBRARY} pthread)
|
||||
|
||||
install(FILES
|
||||
raylet
|
||||
DESTINATION "${CMAKE_INSTALL_INCLUDEDIR}/ray/raylet")
|
||||
@@ -0,0 +1,11 @@
|
||||
#include "actor.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
ActorInformation::ActorInformation() : id_(UniqueID::nil()) {}
|
||||
|
||||
ActorInformation::~ActorInformation() {}
|
||||
|
||||
const ActorID &ActorInformation::GetActorId() const { return this->id_; }
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,26 @@
|
||||
#ifndef RAY_RAYLET_ACTOR_H
|
||||
#define RAY_RAYLET_ACTOR_H
|
||||
|
||||
#include "ray/id.h"
|
||||
|
||||
namespace ray {
|
||||
class ActorInformation {
|
||||
public:
|
||||
/// \brief ActorInformation constructor.
|
||||
ActorInformation();
|
||||
|
||||
/// \brief ActorInformation destructor.
|
||||
~ActorInformation();
|
||||
|
||||
/// \brief Return the id of this actor.
|
||||
/// \return actor id.
|
||||
const ActorID &GetActorId() const;
|
||||
|
||||
private:
|
||||
/// Unique identifier for this actor.
|
||||
ActorID id_;
|
||||
}; // class ActorInformation
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_ACTOR_H
|
||||
@@ -0,0 +1,145 @@
|
||||
// Local scheduler protocol specification
|
||||
|
||||
enum MessageType:int {
|
||||
// Task is submitted to the local scheduler. This is sent from a worker to a
|
||||
// local scheduler.
|
||||
SubmitTask = 1,
|
||||
// Notify the local scheduler that a task has finished. This is sent from a
|
||||
// worker to a local scheduler.
|
||||
TaskDone,
|
||||
// Log a message to the event table. This is sent from a worker to a local
|
||||
// scheduler.
|
||||
EventLogMessage,
|
||||
// Send an initial connection message to the local scheduler. This is sent
|
||||
// from a worker or driver to a local scheduler.
|
||||
RegisterClientRequest,
|
||||
// Send a reply confirming the successful registration of a worker or driver.
|
||||
// This is sent from the local scheduler to a worker or driver.
|
||||
RegisterClientReply,
|
||||
// Notify the local scheduler that this client is disconnecting gracefully.
|
||||
// This is sent from a worker to a local scheduler.
|
||||
DisconnectClient,
|
||||
// Get a new task from the local scheduler. This is sent from a worker to a
|
||||
// local scheduler.
|
||||
GetTask,
|
||||
// Tell a worker to execute a task. This is sent from a local scheduler to a
|
||||
// worker.
|
||||
ExecuteTask,
|
||||
// Reconstruct a possibly lost object. This is sent from a worker to a local
|
||||
// scheduler.
|
||||
ReconstructObject,
|
||||
// For a worker that was blocked on some object(s), tell the local scheduler
|
||||
// that the worker is now unblocked. This is sent from a worker to a local
|
||||
// scheduler.
|
||||
NotifyUnblocked,
|
||||
// Add a result table entry for an object put.
|
||||
PutObject,
|
||||
// A request to get the task frontier for an actor, called by the actor when
|
||||
// saving a checkpoint.
|
||||
GetActorFrontierRequest,
|
||||
// The ActorFrontier response to a GetActorFrontierRequest. The local
|
||||
// scheduler returns the actor's per-handle task counts and execution
|
||||
// dependencies, which can later be used as the argument to SetActorFrontier
|
||||
// when resuming from the checkpoint.
|
||||
GetActorFrontierReply,
|
||||
// A request to set the task frontier for an actor, called when resuming from
|
||||
// a checkpoint. The local scheduler will update the actor's per-handle task
|
||||
// counts and execution dependencies, discard any tasks that already executed
|
||||
// before the checkpoint, and make any tasks on the frontier runnable by
|
||||
// making their execution dependencies available.
|
||||
SetActorFrontier
|
||||
}
|
||||
|
||||
table SubmitTaskRequest {
|
||||
execution_dependencies: [string];
|
||||
task_spec: string;
|
||||
}
|
||||
|
||||
// This message is sent from the local scheduler to a worker.
|
||||
table GetTaskReply {
|
||||
// A string of bytes representing the task specification.
|
||||
task_spec: string;
|
||||
// The IDs of the GPUs that the worker is allowed to use for this task.
|
||||
gpu_ids: [int];
|
||||
}
|
||||
|
||||
table EventLogMessage {
|
||||
key: string;
|
||||
value: string;
|
||||
timestamp: double;
|
||||
}
|
||||
|
||||
// This struct is used to register a new worker with the local scheduler.
|
||||
// It is shipped as part of local_scheduler_connect.
|
||||
table RegisterClientRequest {
|
||||
// True if the client is a worker and false if the client is a driver.
|
||||
is_worker: bool;
|
||||
// The ID of the worker or driver.
|
||||
client_id: string;
|
||||
// The ID of the actor. This is NIL_ACTOR_ID if the worker is not an actor.
|
||||
actor_id: string;
|
||||
// The process ID of this worker.
|
||||
worker_pid: long;
|
||||
// The number of GPUs required by this actor.
|
||||
num_gpus: long;
|
||||
}
|
||||
|
||||
table RegisterClientReply {
|
||||
// The IDs of the GPUs that are reserved for this worker.
|
||||
gpu_ids: [int];
|
||||
}
|
||||
|
||||
table DisconnectClient {
|
||||
}
|
||||
|
||||
table ReconstructObject {
|
||||
// Object ID of the object that needs to be reconstructed.
|
||||
object_id: string;
|
||||
}
|
||||
|
||||
table PutObject {
|
||||
// Task ID of the task that performed the put.
|
||||
task_id: string;
|
||||
// Object ID of the object that is being put.
|
||||
object_id: string;
|
||||
}
|
||||
|
||||
// The ActorFrontier is used to represent the current frontier of tasks that
|
||||
// the local scheduler has marked as runnable for a particular actor. It is
|
||||
// used to save the point in an actor's lifetime at which a checkpoint was
|
||||
// taken, so that the same frontier of tasks can be made runnable again if the
|
||||
// actor is resumed from that checkpoint.
|
||||
table ActorFrontier {
|
||||
// Actor ID of the actor whose frontier is described.
|
||||
actor_id: string;
|
||||
// A list of handle IDs, representing the callers of the actor that have
|
||||
// submitted a runnable task to the local scheduler. A nil ID represents the
|
||||
// creator of the actor.
|
||||
handle_ids: [string];
|
||||
// A list representing the number of tasks executed so far, per handle. Each
|
||||
// count in task_counters corresponds to the handle at the same in index in
|
||||
// handle_ids.
|
||||
task_counters: [long];
|
||||
// A list representing the execution dependency for the next runnable task,
|
||||
// per handle. Each execution dependency in frontier_dependencies corresponds
|
||||
// to the handle at the same in index in handle_ids.
|
||||
frontier_dependencies: [string];
|
||||
}
|
||||
|
||||
table GetActorFrontierRequest {
|
||||
actor_id: string;
|
||||
}
|
||||
|
||||
table RegisterNodeManagerRequest {
|
||||
// GCS ClientID of the connecting node manager.
|
||||
client_id: string;
|
||||
}
|
||||
|
||||
table ForwardTaskRequest {
|
||||
// The task to be forwarded.
|
||||
// TODO(swang): Replace with a Task flatbuffer type.
|
||||
task: string;
|
||||
// The uncommitted lineage of the forwarded task, according to the sending
|
||||
// node manager.
|
||||
uncommitted_lineage: [string];
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
#include "lineage_cache.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
LineageCache::LineageCache() {}
|
||||
|
||||
ray::Status LineageCache::AddTask(const Task &task) {
|
||||
throw std::runtime_error("method not implemented");
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
ray::Status LineageCache::AddTask(const Task &task, const Lineage &uncommitted_lineage) {
|
||||
throw std::runtime_error("method not implemented");
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
ray::Status LineageCache::AddObjectLocation(const ObjectID &object_id) {
|
||||
throw std::runtime_error("method not implemented");
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
Lineage &LineageCache::GetUncommittedLineage(const ObjectID &object_id) {
|
||||
throw std::runtime_error("method not implemented");
|
||||
}
|
||||
|
||||
Status LineageCache::Flush() {
|
||||
throw std::runtime_error("method not implemented");
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,83 @@
|
||||
#ifndef RAY_RAYLET_LINEAGE_CACHE_H
|
||||
#define RAY_RAYLET_LINEAGE_CACHE_H
|
||||
|
||||
// clang-format off
|
||||
#include "ray/raylet/task.h"
|
||||
#include "ray/id.h"
|
||||
#include "ray/status.h"
|
||||
// clang-format on
|
||||
|
||||
namespace ray {
|
||||
|
||||
// TODO(swang): Define this class.
|
||||
class Lineage {};
|
||||
|
||||
class LineageCacheEntry {
|
||||
private:
|
||||
// TODO(swang): This should be an enum of the state of the entry - goes from
|
||||
// completely local, to dirty, to in flight, to committed.
|
||||
bool dirty_;
|
||||
};
|
||||
|
||||
class LineageCacheTaskEntry : public LineageCacheEntry {};
|
||||
class LineageCacheObjectEntry : public LineageCacheEntry {};
|
||||
|
||||
/// \class LineageCache
|
||||
///
|
||||
/// A cache of the object and task tables. This consists of all tasks that this
|
||||
/// node owns, as well as their task lineage, that have not yet been added
|
||||
/// durably to the GCS.
|
||||
class LineageCache {
|
||||
public:
|
||||
/// Create a lineage cache policy.
|
||||
/// TODO(swang): Pass in the policy (interface?) and a GCS client.
|
||||
LineageCache();
|
||||
|
||||
/// Add a task and its object outputs asynchronously to the GCS. This
|
||||
/// overwrites the task’s mutable fields in the execution specification.
|
||||
///
|
||||
/// \param task The task to add.
|
||||
/// \return Status.
|
||||
ray::Status AddTask(const Task &task);
|
||||
|
||||
/// Add a task and its uncommitted lineage asynchronously to the GCS. The
|
||||
/// mutable fields for the given task will be overwritten, but not for the
|
||||
/// tasks in the uncommitted lineage.
|
||||
///
|
||||
/// \param task The task to add.
|
||||
/// \param uncommitted_lineage The task's uncommitted lineage. These are the
|
||||
/// tasks that the given task is data-dependent on, but that have not
|
||||
/// been made durable in the GCS, as far as we know.
|
||||
/// \return Status.
|
||||
ray::Status AddTask(const Task &task, const Lineage &uncommitted_lineage);
|
||||
|
||||
/// Add this node as an object location, to be asynchronously committed to
|
||||
/// the GCS.
|
||||
///
|
||||
/// \param object_id The object to add a location for.
|
||||
/// \return Status.
|
||||
ray::Status AddObjectLocation(const ObjectID &object_id);
|
||||
|
||||
/// Get the uncommitted lineage of an object. These are the tasks that the
|
||||
/// given object is data-dependent on, but that have not been made durable in
|
||||
/// the GCS, as far as we know.
|
||||
///
|
||||
/// \param object_id The object to get the uncommitted lineage for.
|
||||
/// \return The uncommitted lineage of the object.
|
||||
Lineage &GetUncommittedLineage(const ObjectID &object_id);
|
||||
|
||||
/// Asynchronously write any tasks and object locations that have been added
|
||||
/// since the last flush to the GCS. When each write is acknowledged, its
|
||||
/// entry will be marked as committed.
|
||||
///
|
||||
/// \return Status.
|
||||
Status Flush();
|
||||
|
||||
private:
|
||||
std::unordered_map<TaskID, LineageCacheTaskEntry, UniqueIDHasher> task_table_;
|
||||
std::unordered_map<ObjectID, LineageCacheObjectEntry, UniqueIDHasher> object_table_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_LINEAGE_CACHE_H
|
||||
@@ -0,0 +1,38 @@
|
||||
#include <iostream>
|
||||
|
||||
#include "ray/raylet/raylet.h"
|
||||
#include "ray/status.h"
|
||||
|
||||
#ifndef RAYLET_TEST
|
||||
int main(int argc, char *argv[]) {
|
||||
RAY_CHECK(argc == 2);
|
||||
|
||||
// start store
|
||||
std::string executable_str = std::string(argv[0]);
|
||||
std::string exec_dir = executable_str.substr(0, executable_str.find_last_of("/"));
|
||||
std::string plasma_dir = exec_dir + "./../plasma";
|
||||
std::string plasma_command =
|
||||
plasma_dir +
|
||||
"/plasma_store -m 1000000000 -s /tmp/store 1> /dev/null 2> /dev/null &";
|
||||
RAY_LOG(INFO) << plasma_command;
|
||||
int s = system(plasma_command.c_str());
|
||||
RAY_CHECK(s == 0);
|
||||
|
||||
// configure
|
||||
std::unordered_map<std::string, double> static_resource_conf;
|
||||
static_resource_conf = {{"CPU", 1}, {"GPU", 1}};
|
||||
ray::ResourceSet resource_config(std::move(static_resource_conf));
|
||||
ray::ObjectManagerConfig om_config;
|
||||
om_config.store_socket_name = "/tmp/store";
|
||||
|
||||
// initialize mock gcs & object directory
|
||||
std::shared_ptr<ray::GcsClient> mock_gcs_client =
|
||||
std::shared_ptr<ray::GcsClient>(new ray::GcsClient());
|
||||
|
||||
// Initialize the node manager.
|
||||
boost::asio::io_service io_service;
|
||||
ray::Raylet server(io_service, std::string(argv[1]), resource_config, om_config,
|
||||
mock_gcs_client);
|
||||
io_service.run();
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,121 @@
|
||||
#include <functional>
|
||||
#include <iostream>
|
||||
|
||||
#include "ray/raylet/mock_gcs_client.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
ray::Status ObjectTable::GetObjectClientIDs(const ray::ObjectID &object_id,
|
||||
const ClientIDsCallback &success,
|
||||
const FailCallback &fail) {
|
||||
RAY_LOG(DEBUG) << "GetObjectClientIDs " << object_id.hex();
|
||||
if (client_lookup.count(object_id) > 0) {
|
||||
if (!client_lookup[object_id].empty()) {
|
||||
std::vector<ClientID> v;
|
||||
for (auto client_id : client_lookup[object_id]) {
|
||||
v.push_back(client_id);
|
||||
}
|
||||
success(std::move(v));
|
||||
return Status::OK();
|
||||
} else {
|
||||
fail(Status::KeyError("ObjectID has no clients."));
|
||||
return Status::OK();
|
||||
}
|
||||
} else {
|
||||
fail(Status::KeyError("ObjectID doesn't exist."));
|
||||
return Status::OK();
|
||||
}
|
||||
}
|
||||
|
||||
ray::Status ObjectTable::Add(const ObjectID &object_id, const ClientID &client_id,
|
||||
const DoneCallback &done_callback) {
|
||||
if (client_lookup.count(object_id) == 0) {
|
||||
RAY_LOG(DEBUG) << "Add ObjectID set " << object_id.hex();
|
||||
client_lookup[object_id] = std::unordered_set<ClientID, UniqueIDHasher>();
|
||||
} else if (client_lookup[object_id].count(client_id) != 0) {
|
||||
return ray::Status::KeyError("ClientID already exists.");
|
||||
}
|
||||
RAY_LOG(DEBUG) << "Insert ClientID " << client_id.hex();
|
||||
client_lookup[object_id].insert(client_id);
|
||||
done_callback();
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
ray::Status ObjectTable::Remove(const ObjectID &object_id, const ClientID &client_id,
|
||||
const DoneCallback &done_callback) {
|
||||
if (client_lookup.count(object_id) == 0) {
|
||||
return ray::Status::KeyError("ObjectID doesn't exist.");
|
||||
} else if (client_lookup[object_id].count(client_id) == 0) {
|
||||
return ray::Status::KeyError("ClientID doesn't exist.");
|
||||
}
|
||||
client_lookup[object_id].erase(client_id);
|
||||
done_callback();
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
ray::Status ClientTable::GetClientIds(ClientIDsCallback callback) {
|
||||
std::vector<ClientID> keys;
|
||||
keys.reserve(info_lookup.size());
|
||||
for (auto kv : info_lookup) {
|
||||
keys.push_back(kv.first);
|
||||
}
|
||||
callback(keys);
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
void ClientTable::GetClientInformationSet(const std::vector<ClientID> &client_ids,
|
||||
ManyInfoCallback callback,
|
||||
FailCallback failcb) {
|
||||
std::vector<ClientInformation> info_vec;
|
||||
for (const auto &client_id : client_ids) {
|
||||
if (info_lookup.count(client_id) != 0) {
|
||||
info_vec.push_back(info_lookup.at(client_id));
|
||||
}
|
||||
}
|
||||
if (info_vec.empty()) {
|
||||
failcb(Status::KeyError("ClientID not found."));
|
||||
} else {
|
||||
callback(info_vec);
|
||||
}
|
||||
}
|
||||
|
||||
void ClientTable::GetClientInformation(const ClientID &client_id,
|
||||
SingleInfoCallback callback, FailCallback failcb) {
|
||||
if (info_lookup.count(client_id) == 0) {
|
||||
failcb(ray::Status::KeyError("CleintID not found."));
|
||||
} else {
|
||||
callback(info_lookup.at(client_id));
|
||||
}
|
||||
}
|
||||
|
||||
ray::Status ClientTable::Add(const ClientID &client_id, const std::string &ip,
|
||||
uint16_t port, DoneCallback done_callback) {
|
||||
if (info_lookup.count(client_id) != 0) {
|
||||
return ray::Status::KeyError("ClientID already exists.");
|
||||
}
|
||||
info_lookup.emplace(client_id, ClientInformation(client_id, ip, port));
|
||||
done_callback();
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
ray::Status ClientTable::Remove(const ClientID &client_id, DoneCallback done_callback) {
|
||||
if (info_lookup.count(client_id) == 0) {
|
||||
return ray::Status::KeyError("ClientID doesn't exist.");
|
||||
}
|
||||
info_lookup.erase(client_id);
|
||||
done_callback();
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
ClientID GcsClient::Register(const std::string &ip, uint16_t port) {
|
||||
ClientID client_id = ClientID().from_random();
|
||||
// TODO: handle client registration failure.
|
||||
ray::Status status = client_table().Add(std::move(client_id), ip, port, []() {});
|
||||
return client_id;
|
||||
}
|
||||
|
||||
ObjectTable &GcsClient::object_table() { return *object_table_; }
|
||||
|
||||
ClientTable &GcsClient::client_table() { return *client_table_; }
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,94 @@
|
||||
#ifndef RAY_RAYLET_MOCK_GCS_CLIENT_H
|
||||
#define RAY_RAYLET_MOCK_GCS_CLIENT_H
|
||||
|
||||
#include <cstdint>
|
||||
#include <functional>
|
||||
#include <map>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
#include <vector>
|
||||
|
||||
#include <boost/asio.hpp>
|
||||
#include <boost/asio/error.hpp>
|
||||
#include <boost/bind.hpp>
|
||||
|
||||
#include "ray/id.h"
|
||||
#include "ray/status.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
class ObjectTable {
|
||||
public:
|
||||
using DoneCallback = std::function<void()>;
|
||||
using ClientIDsCallback = std::function<void(const std::vector<ray::ClientID> &)>;
|
||||
using FailCallback = std::function<void(const ray::Status &)>;
|
||||
ray::Status Add(const ObjectID &object_id, const ClientID &client_id,
|
||||
const DoneCallback &done);
|
||||
ray::Status Remove(const ObjectID &object_id, const ClientID &client_id,
|
||||
const DoneCallback &done);
|
||||
ray::Status GetObjectClientIDs(const ObjectID &object_id, const ClientIDsCallback &,
|
||||
const FailCallback &);
|
||||
|
||||
private:
|
||||
std::vector<ClientID> empty_set_;
|
||||
std::unordered_map<ObjectID, std::unordered_set<ClientID, UniqueIDHasher>,
|
||||
UniqueIDHasher>
|
||||
client_lookup;
|
||||
};
|
||||
|
||||
class ClientInformation {
|
||||
public:
|
||||
ClientInformation(const ClientID &client_id, const std::string &ip_address,
|
||||
uint16_t port)
|
||||
: client_id_(client_id), ip_address_(ip_address), port_(port) {}
|
||||
const ClientID &GetClientId() const { return client_id_; }
|
||||
const std::string &GetIp() const { return ip_address_; }
|
||||
const uint16_t &GetPort() const { return port_; }
|
||||
|
||||
private:
|
||||
ClientID client_id_;
|
||||
std::string ip_address_;
|
||||
uint16_t port_;
|
||||
};
|
||||
|
||||
class ClientTable {
|
||||
public:
|
||||
typedef std::unordered_map<ClientID, ClientInformation, UniqueIDHasher> info_type;
|
||||
|
||||
using ClientIDsCallback = std::function<void(std::vector<ray::ClientID>)>;
|
||||
using SingleInfoCallback = std::function<void(ClientInformation info)>;
|
||||
using ManyInfoCallback = std::function<void(std::vector<ClientInformation> info_vec)>;
|
||||
using DoneCallback = std::function<void()>;
|
||||
using FailCallback = std::function<void(ray::Status)>;
|
||||
|
||||
ray::Status GetClientIds(ClientIDsCallback cb);
|
||||
void GetClientInformationSet(const std::vector<ClientID> &client_ids,
|
||||
ManyInfoCallback cb, FailCallback failcb);
|
||||
void GetClientInformation(const ClientID &client_id, SingleInfoCallback callback,
|
||||
FailCallback failcb);
|
||||
ray::Status Add(const ClientID &client_id, const std::string &ip, uint16_t port,
|
||||
DoneCallback cb);
|
||||
ray::Status Remove(const ClientID &client_id, DoneCallback done);
|
||||
|
||||
private:
|
||||
info_type info_lookup;
|
||||
};
|
||||
|
||||
class GcsClient {
|
||||
public:
|
||||
GcsClient() {
|
||||
this->object_table_.reset(new ObjectTable());
|
||||
this->client_table_.reset(new ClientTable());
|
||||
}
|
||||
// Register the ip and port of the connecting client.
|
||||
ClientID Register(const std::string &ip, uint16_t port);
|
||||
ObjectTable &object_table();
|
||||
ClientTable &client_table();
|
||||
|
||||
private:
|
||||
std::unique_ptr<ObjectTable> object_table_;
|
||||
std::unique_ptr<ClientTable> client_table_;
|
||||
};
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_MOCK_GCS_CLIENT_H
|
||||
@@ -0,0 +1,174 @@
|
||||
#include "ray/raylet/node_manager.h"
|
||||
|
||||
#include "common_protocol.h"
|
||||
#include "ray/raylet/format/node_manager_generated.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
NodeManager::NodeManager(const std::string &socket_name,
|
||||
const ResourceSet &resource_config,
|
||||
ObjectManager &object_manager)
|
||||
: local_resources_(resource_config),
|
||||
worker_pool_(WorkerPool(0)),
|
||||
local_queues_(SchedulingQueue()),
|
||||
scheduling_policy_(local_queues_),
|
||||
reconstruction_policy_([this](const TaskID &task_id) { ResubmitTask(task_id); }),
|
||||
task_dependency_manager_(
|
||||
object_manager,
|
||||
// reconstruction_policy_,
|
||||
[this](const TaskID &task_id) { HandleWaitingTaskReady(task_id); }) {
|
||||
//// TODO(atumanov): need to add the self-knowledge of ClientID, using nill().
|
||||
// cluster_resource_map_[ClientID::nil()] = local_resources_;
|
||||
}
|
||||
|
||||
void NodeManager::ProcessNewClient(std::shared_ptr<LocalClientConnection> client) {
|
||||
// The new client is a worker, so begin listening for messages.
|
||||
client->ProcessMessages();
|
||||
}
|
||||
|
||||
void NodeManager::ProcessClientMessage(std::shared_ptr<LocalClientConnection> client,
|
||||
int64_t message_type,
|
||||
const uint8_t *message_data) {
|
||||
RAY_LOG(DEBUG) << "Message of type " << message_type;
|
||||
|
||||
switch (message_type) {
|
||||
case MessageType_RegisterClientRequest: {
|
||||
auto message = flatbuffers::GetRoot<RegisterClientRequest>(message_data);
|
||||
if (message->is_worker()) {
|
||||
// Create a new worker from the registration request.
|
||||
std::shared_ptr<Worker> worker(new Worker(message->worker_pid(), client));
|
||||
// Register the new worker.
|
||||
worker_pool_.RegisterWorker(std::move(worker));
|
||||
}
|
||||
|
||||
// Build the reply to the worker's registration request. TODO(swang): This
|
||||
// is legacy code and should be removed once actor creation tasks are
|
||||
// implemented.
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
auto reply = CreateRegisterClientReply(fbb, fbb.CreateVector(std::vector<int>()));
|
||||
fbb.Finish(reply);
|
||||
// Reply to the worker's registration request, then listen for more
|
||||
// messages.
|
||||
client->WriteMessage(MessageType_RegisterClientReply, fbb.GetSize(),
|
||||
fbb.GetBufferPointer());
|
||||
} break;
|
||||
case MessageType_GetTask: {
|
||||
const std::shared_ptr<Worker> worker = worker_pool_.GetRegisteredWorker(client);
|
||||
RAY_CHECK(worker);
|
||||
// If the worker was assigned a task, mark it as finished.
|
||||
if (!worker->GetAssignedTaskId().is_nil()) {
|
||||
FinishTask(worker->GetAssignedTaskId());
|
||||
}
|
||||
// Return the worker to the idle pool.
|
||||
worker_pool_.PushWorker(worker);
|
||||
auto scheduled_tasks = local_queues_.GetScheduledTasks();
|
||||
if (!scheduled_tasks.empty()) {
|
||||
const TaskID &scheduled_task_id =
|
||||
scheduled_tasks.front().GetTaskSpecification().TaskId();
|
||||
auto scheduled_tasks = local_queues_.RemoveTasks({scheduled_task_id});
|
||||
AssignTask(scheduled_tasks.front());
|
||||
}
|
||||
} break;
|
||||
case MessageType_DisconnectClient: {
|
||||
// Remove the dead worker from the pool and stop listening for messages.
|
||||
const std::shared_ptr<Worker> worker = worker_pool_.GetRegisteredWorker(client);
|
||||
if (worker) {
|
||||
worker_pool_.DisconnectWorker(worker);
|
||||
}
|
||||
} break;
|
||||
case MessageType_SubmitTask: {
|
||||
// Read the task submitted by the client.
|
||||
auto message = flatbuffers::GetRoot<SubmitTaskRequest>(message_data);
|
||||
TaskExecutionSpecification task_execution_spec(
|
||||
from_flatbuf(*message->execution_dependencies()));
|
||||
TaskSpecification task_spec(*message->task_spec());
|
||||
Task task(task_execution_spec, task_spec);
|
||||
// Submit the task to the local scheduler.
|
||||
SubmitTask(task);
|
||||
// Listen for more messages.
|
||||
client->ProcessMessages();
|
||||
} break;
|
||||
default:
|
||||
RAY_LOG(FATAL) << "Received unexpected message type " << message_type;
|
||||
}
|
||||
}
|
||||
|
||||
void NodeManager::HandleWaitingTaskReady(const TaskID &task_id) {
|
||||
auto ready_tasks = local_queues_.RemoveTasks({task_id});
|
||||
local_queues_.QueueReadyTasks(std::vector<Task>(ready_tasks));
|
||||
// Schedule the newly ready tasks if possible.
|
||||
ScheduleTasks();
|
||||
}
|
||||
|
||||
void NodeManager::ScheduleTasks() {
|
||||
// Ask policy for scheduling decision.
|
||||
// TODO(alexey): Give the policy all cluster resources instead of just the
|
||||
// local one.
|
||||
std::unordered_map<ClientID, SchedulingResources, UniqueIDHasher> cluster_resource_map;
|
||||
cluster_resource_map[ClientID::nil()] = local_resources_;
|
||||
const auto &policy_decision = scheduling_policy_.Schedule(cluster_resource_map);
|
||||
// Extract decision for this local scheduler.
|
||||
// TODO(alexey): Check for this node's own client ID, not for nil.
|
||||
std::unordered_set<TaskID, UniqueIDHasher> task_ids;
|
||||
for (auto &task_schedule : policy_decision) {
|
||||
if (task_schedule.second.is_nil()) {
|
||||
task_ids.insert(task_schedule.first);
|
||||
}
|
||||
}
|
||||
|
||||
// Assign the tasks to workers.
|
||||
std::vector<Task> tasks = local_queues_.RemoveTasks(task_ids);
|
||||
for (auto &task : tasks) {
|
||||
AssignTask(task);
|
||||
}
|
||||
}
|
||||
|
||||
void NodeManager::SubmitTask(const Task &task) {
|
||||
if (task_dependency_manager_.TaskReady(task)) {
|
||||
local_queues_.QueueReadyTasks(std::vector<Task>({task}));
|
||||
ScheduleTasks();
|
||||
} else {
|
||||
local_queues_.QueueWaitingTasks(std::vector<Task>({task}));
|
||||
task_dependency_manager_.SubscribeTaskReady(task);
|
||||
}
|
||||
}
|
||||
|
||||
void NodeManager::AssignTask(const Task &task) {
|
||||
if (worker_pool_.PoolSize() == 0) {
|
||||
// Start a new worker.
|
||||
worker_pool_.StartWorker();
|
||||
// Queue this task for future assignment. The task will be assigned to a
|
||||
// worker once one becomes available.
|
||||
local_queues_.QueueScheduledTasks(std::vector<Task>({task}));
|
||||
// TODO(swang): Acquire resources here or when a worker becomes available?
|
||||
return;
|
||||
}
|
||||
|
||||
std::shared_ptr<Worker> worker = worker_pool_.PopWorker();
|
||||
RAY_LOG(DEBUG) << "Assigning task to worker with pid " << worker->Pid();
|
||||
|
||||
// TODO(swang): Acquire resources for the task.
|
||||
// local_resources_.Acquire(task.GetTaskSpecification().GetRequiredResources());
|
||||
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
const TaskSpecification &spec = task.GetTaskSpecification();
|
||||
auto message = CreateGetTaskReply(fbb, spec.ToFlatbuffer(fbb),
|
||||
fbb.CreateVector(std::vector<int>()));
|
||||
fbb.Finish(message);
|
||||
worker->Connection()->WriteMessage(MessageType_ExecuteTask, fbb.GetSize(),
|
||||
fbb.GetBufferPointer());
|
||||
worker->AssignTaskId(spec.TaskId());
|
||||
local_queues_.QueueRunningTasks(std::vector<Task>({task}));
|
||||
}
|
||||
|
||||
void NodeManager::FinishTask(const TaskID &task_id) {
|
||||
RAY_LOG(DEBUG) << "Finished task " << task_id.hex();
|
||||
local_queues_.RemoveTasks({task_id});
|
||||
// TODO(swang): Release resources that were held for the task.
|
||||
}
|
||||
|
||||
void NodeManager::ResubmitTask(const TaskID &task_id) {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,75 @@
|
||||
#ifndef RAY_RAYLET_NODE_MANAGER_H
|
||||
#define RAY_RAYLET_NODE_MANAGER_H
|
||||
|
||||
// clang-format off
|
||||
#include "ray/common/client_connection.h"
|
||||
#include "ray/raylet/scheduling_policy.h"
|
||||
#include "ray/raylet/scheduling_queue.h"
|
||||
#include "ray/raylet/scheduling_resources.h"
|
||||
#include "ray/object_manager/object_manager.h"
|
||||
#include "ray/raylet/reconstruction_policy.h"
|
||||
#include "ray/raylet/task_dependency_manager.h"
|
||||
#include "ray/raylet/worker_pool.h"
|
||||
// clang-format on
|
||||
|
||||
namespace ray {
|
||||
|
||||
class NodeManager : public ClientManager<boost::asio::local::stream_protocol> {
|
||||
public:
|
||||
/// Create a node manager.
|
||||
///
|
||||
/// \param socket_name The pathname of the Unix domain socket to listen at
|
||||
/// for local connections.
|
||||
/// \param resource_config The initial set of node resources.
|
||||
/// \param object_manager A reference to the local object manager.
|
||||
NodeManager(const std::string &socket_name, const ResourceSet &resource_config,
|
||||
ObjectManager &object_manager);
|
||||
|
||||
/// Process a new client connection.
|
||||
void ProcessNewClient(std::shared_ptr<LocalClientConnection> client);
|
||||
|
||||
/// Process a message from a client. This method is responsible for
|
||||
/// explicitly listening for more messages from the client if the client is
|
||||
/// still alive.
|
||||
///
|
||||
/// \param client The client that sent the message.
|
||||
/// \param message_type The message type (e.g., a flatbuffer enum).
|
||||
/// \param message A pointer to the message data.
|
||||
void ProcessClientMessage(std::shared_ptr<LocalClientConnection> client,
|
||||
int64_t message_type, const uint8_t *message);
|
||||
|
||||
private:
|
||||
/// Submit a task to this node.
|
||||
void SubmitTask(const Task &task);
|
||||
/// Assign a task.
|
||||
void AssignTask(const Task &task);
|
||||
/// Finish a task.
|
||||
void FinishTask(const TaskID &task_id);
|
||||
/// Schedule tasks.
|
||||
void ScheduleTasks();
|
||||
/// Handle a task whose local dependencies were missing and are now
|
||||
/// available.
|
||||
void HandleWaitingTaskReady(const TaskID &task_id);
|
||||
/// Resubmit a task whose return value needs to be reconstructed.
|
||||
void ResubmitTask(const TaskID &task_id);
|
||||
|
||||
/// The resources local to this node.
|
||||
SchedulingResources local_resources_;
|
||||
// TODO(atumanov): Add resource information from other nodes.
|
||||
// std::unordered_map<ClientID, SchedulingResources&, UniqueIDHasher>
|
||||
// cluster_resource_map_;
|
||||
/// A pool of workers.
|
||||
WorkerPool worker_pool_;
|
||||
/// A set of queues to maintain tasks.
|
||||
SchedulingQueue local_queues_;
|
||||
/// The scheduling policy in effect for this local scheduler.
|
||||
SchedulingPolicy scheduling_policy_;
|
||||
/// The reconstruction policy for deciding when to re-execute a task.
|
||||
ReconstructionPolicy reconstruction_policy_;
|
||||
/// A manager to make waiting tasks's missing object dependencies available.
|
||||
TaskDependencyManager task_dependency_manager_;
|
||||
};
|
||||
|
||||
} // end namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_NODE_MANAGER_H
|
||||
@@ -0,0 +1,73 @@
|
||||
#include "raylet.h"
|
||||
|
||||
#include <boost/bind.hpp>
|
||||
#include <iostream>
|
||||
|
||||
#include "ray/status.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
Raylet::Raylet(boost::asio::io_service &io_service, const std::string &socket_name,
|
||||
const ResourceSet &resource_config,
|
||||
const ObjectManagerConfig &object_manager_config,
|
||||
std::shared_ptr<ray::GcsClient> gcs_client)
|
||||
: acceptor_(io_service, boost::asio::local::stream_protocol::endpoint(socket_name)),
|
||||
socket_(io_service),
|
||||
tcp_acceptor_(io_service,
|
||||
boost::asio::ip::tcp::endpoint(boost::asio::ip::tcp::v4(), 0)),
|
||||
tcp_socket_(io_service),
|
||||
object_manager_(io_service, object_manager_config, gcs_client),
|
||||
node_manager_(socket_name, resource_config, object_manager_),
|
||||
gcs_client_(gcs_client) {
|
||||
ClientID client_id = RegisterGcs();
|
||||
object_manager_.SetClientID(client_id);
|
||||
// Start listening for clients.
|
||||
DoAccept();
|
||||
DoAcceptTcp();
|
||||
}
|
||||
|
||||
Raylet::~Raylet() { RAY_CHECK_OK(object_manager_.Terminate()); }
|
||||
|
||||
ClientID Raylet::RegisterGcs() {
|
||||
boost::asio::ip::tcp::endpoint endpoint = tcp_acceptor_.local_endpoint();
|
||||
std::string ip = endpoint.address().to_string();
|
||||
uint16_t port = endpoint.port();
|
||||
ClientID client_id = gcs_client_->Register(ip, port);
|
||||
return client_id;
|
||||
}
|
||||
|
||||
void Raylet::DoAcceptTcp() {
|
||||
TCPClientConnection::pointer new_connection =
|
||||
TCPClientConnection::Create(acceptor_.get_io_service());
|
||||
tcp_acceptor_.async_accept(new_connection->GetSocket(),
|
||||
boost::bind(&Raylet::HandleAcceptTcp, this, new_connection,
|
||||
boost::asio::placeholders::error));
|
||||
}
|
||||
|
||||
void Raylet::HandleAcceptTcp(TCPClientConnection::pointer new_connection,
|
||||
const boost::system::error_code &error) {
|
||||
if (!error) {
|
||||
// Pass it off to object manager for now.
|
||||
ray::Status status = object_manager_.AcceptConnection(std::move(new_connection));
|
||||
}
|
||||
DoAcceptTcp();
|
||||
}
|
||||
|
||||
void Raylet::DoAccept() {
|
||||
acceptor_.async_accept(socket_, boost::bind(&Raylet::HandleAccept, this,
|
||||
boost::asio::placeholders::error));
|
||||
}
|
||||
|
||||
void Raylet::HandleAccept(const boost::system::error_code &error) {
|
||||
if (!error) {
|
||||
// Accept a new local client and dispatch it to the node manager.
|
||||
auto new_connection =
|
||||
LocalClientConnection::Create(node_manager_, std::move(socket_));
|
||||
}
|
||||
// We're ready to accept another client.
|
||||
DoAccept();
|
||||
}
|
||||
|
||||
ObjectManager &Raylet::GetObjectManager() { return object_manager_; }
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,76 @@
|
||||
#ifndef RAY_RAYLET_RAYLET_H
|
||||
#define RAY_RAYLET_RAYLET_H
|
||||
|
||||
#include <list>
|
||||
|
||||
#include <boost/asio.hpp>
|
||||
#include <boost/asio/error.hpp>
|
||||
|
||||
// clang-format off
|
||||
#include "ray/raylet/node_manager.h"
|
||||
#include "ray/object_manager/object_manager.h"
|
||||
#include "ray/raylet/scheduling_resources.h"
|
||||
// clang-format on
|
||||
|
||||
namespace ray {
|
||||
|
||||
class Task;
|
||||
class NodeManager;
|
||||
|
||||
// TODO(swang): Rename class and source files to Raylet.
|
||||
class Raylet {
|
||||
public:
|
||||
/// Create a node manager server and listen for new clients.
|
||||
///
|
||||
/// \param io_service The event loop to run the server on.
|
||||
/// \param socket_name The Unix domain socket to listen on for local clients.
|
||||
/// \param resource_config The initial set of resources to start the local
|
||||
/// scheduler with.
|
||||
/// \param object_manager_config Configuration to initialize the object
|
||||
/// manager.
|
||||
/// \param gcs_client A client connection to the GCS.
|
||||
Raylet(boost::asio::io_service &io_service, const std::string &socket_name,
|
||||
const ResourceSet &resource_config,
|
||||
const ObjectManagerConfig &object_manager_config,
|
||||
std::shared_ptr<ray::GcsClient> gcs_client);
|
||||
|
||||
/// Destroy the NodeServer.
|
||||
~Raylet();
|
||||
|
||||
// TODO(melih): Get rid of this method.
|
||||
ObjectManager &GetObjectManager();
|
||||
|
||||
private:
|
||||
/// Register GCS client.
|
||||
ClientID RegisterGcs();
|
||||
/// Accept a client connection.
|
||||
void DoAccept();
|
||||
/// Handle an accepted client connection.
|
||||
void HandleAccept(const boost::system::error_code &error);
|
||||
/// Accept a tcp client connection.
|
||||
void DoAcceptTcp();
|
||||
/// Handle an accepted tcp client connection.
|
||||
void HandleAcceptTcp(TCPClientConnection::pointer new_connection,
|
||||
const boost::system::error_code &error);
|
||||
|
||||
/// An acceptor for new clients.
|
||||
boost::asio::local::stream_protocol::acceptor acceptor_;
|
||||
/// The socket to listen on for new clients.
|
||||
boost::asio::local::stream_protocol::socket socket_;
|
||||
/// An acceptor for new tcp clients.
|
||||
boost::asio::ip::tcp::acceptor tcp_acceptor_;
|
||||
/// The socket to listen on for new tcp clients.
|
||||
boost::asio::ip::tcp::socket tcp_socket_;
|
||||
|
||||
// TODO(swang): Lineage cache.
|
||||
/// Manages client requests for object transfers and availability.
|
||||
ObjectManager object_manager_;
|
||||
/// Manages client requests for task submission and execution.
|
||||
NodeManager node_manager_;
|
||||
/// A client connection to the GCS.
|
||||
std::shared_ptr<ray::GcsClient> gcs_client_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_RAYLET_H
|
||||
@@ -0,0 +1,275 @@
|
||||
#include <iostream>
|
||||
#include <thread>
|
||||
|
||||
#include "gtest/gtest.h"
|
||||
|
||||
#include "ray/raylet/raylet.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
std::string test_executable; // NOLINT
|
||||
|
||||
class TestRaylet : public ::testing::Test {
|
||||
public:
|
||||
TestRaylet() { RAY_LOG(INFO) << "TestRaylet: started."; }
|
||||
|
||||
std::string StartStore(const std::string &id) {
|
||||
std::string store_id = "/tmp/store";
|
||||
store_id = store_id + id;
|
||||
std::string test_dir = test_executable.substr(0, test_executable.find_last_of("/"));
|
||||
std::string plasma_dir = test_dir + "./../plasma";
|
||||
std::string plasma_command = plasma_dir + "/plasma_store -m 1000000000 -s " +
|
||||
store_id + " 1> /dev/null 2> /dev/null &";
|
||||
RAY_LOG(INFO) << plasma_command;
|
||||
int ec = system(plasma_command.c_str());
|
||||
if (ec != 0) {
|
||||
throw std::runtime_error("failed to start plasma store.");
|
||||
};
|
||||
return store_id;
|
||||
}
|
||||
|
||||
void SetUp() {
|
||||
// start store
|
||||
std::string store_sock_1 = StartStore("1");
|
||||
std::string store_sock_2 = StartStore("2");
|
||||
|
||||
// configure
|
||||
std::unordered_map<std::string, double> static_resource_config;
|
||||
static_resource_config = {{"num_cpus", 1}, {"num_gpus", 1}};
|
||||
ray::ResourceSet resource_config(std::move(static_resource_config));
|
||||
|
||||
// start mock gcs
|
||||
mock_gcs_client = std::shared_ptr<GcsClient>(new GcsClient());
|
||||
|
||||
// start first server
|
||||
ray::ObjectManagerConfig om_config_1;
|
||||
om_config_1.store_socket_name = store_sock_1;
|
||||
server1.reset(new Raylet(io_service, std::string("hello1"), resource_config,
|
||||
om_config_1, mock_gcs_client));
|
||||
|
||||
// start second server
|
||||
ray::ObjectManagerConfig om_config_2;
|
||||
om_config_2.store_socket_name = store_sock_2;
|
||||
server2.reset(new Raylet(io_service, std::string("hello2"), resource_config,
|
||||
om_config_2, mock_gcs_client));
|
||||
|
||||
// connect to stores.
|
||||
ARROW_CHECK_OK(client1.Connect(store_sock_1, "", PLASMA_DEFAULT_RELEASE_DELAY));
|
||||
ARROW_CHECK_OK(client2.Connect(store_sock_2, "", PLASMA_DEFAULT_RELEASE_DELAY));
|
||||
this->StartLoop();
|
||||
}
|
||||
|
||||
void TearDown() {
|
||||
this->StopLoop();
|
||||
arrow::Status client1_status = client1.Disconnect();
|
||||
arrow::Status client2_status = client2.Disconnect();
|
||||
ASSERT_TRUE(client1_status.ok() && client2_status.ok());
|
||||
|
||||
this->server1.reset();
|
||||
this->server2.reset();
|
||||
|
||||
int s = system("killall plasma_store &");
|
||||
ASSERT_TRUE(!s);
|
||||
|
||||
std::string cmd_str = test_executable.substr(0, test_executable.find_last_of("/"));
|
||||
s = system(("rm " + cmd_str + "/hello1").c_str());
|
||||
ASSERT_TRUE(!s);
|
||||
s = system(("rm " + cmd_str + "/hello2").c_str());
|
||||
ASSERT_TRUE(!s);
|
||||
}
|
||||
|
||||
void Loop() { io_service.run(); };
|
||||
|
||||
void StartLoop() { p = std::thread(&TestRaylet::Loop, this); };
|
||||
|
||||
void StopLoop() {
|
||||
io_service.stop();
|
||||
p.join();
|
||||
}
|
||||
|
||||
ObjectID WriteDataToClient(plasma::PlasmaClient &client, int64_t data_size) {
|
||||
ObjectID object_id = ObjectID::from_random();
|
||||
RAY_LOG(DEBUG) << "ObjectID Created: " << object_id.hex().c_str();
|
||||
uint8_t metadata[] = {5};
|
||||
int64_t metadata_size = sizeof(metadata);
|
||||
std::shared_ptr<Buffer> data;
|
||||
ARROW_CHECK_OK(client.Create(object_id.to_plasma_id(), data_size, metadata,
|
||||
metadata_size, &data));
|
||||
ARROW_CHECK_OK(client.Seal(object_id.to_plasma_id()));
|
||||
return object_id;
|
||||
}
|
||||
|
||||
void object_added_handler_1(const ObjectID &object_id) {
|
||||
RAY_LOG(INFO) << "Store 1 added: " << object_id.hex();
|
||||
v1.push_back(object_id);
|
||||
};
|
||||
|
||||
void object_added_handler_2(const ObjectID &object_id) {
|
||||
RAY_LOG(INFO) << "Store 2 added: " << object_id.hex();
|
||||
v2.push_back(object_id);
|
||||
};
|
||||
|
||||
protected:
|
||||
std::thread p;
|
||||
boost::asio::io_service io_service;
|
||||
std::shared_ptr<ray::GcsClient> mock_gcs_client;
|
||||
std::unique_ptr<ray::Raylet> server1;
|
||||
std::unique_ptr<ray::Raylet> server2;
|
||||
|
||||
plasma::PlasmaClient client1;
|
||||
plasma::PlasmaClient client2;
|
||||
std::vector<ObjectID> v1;
|
||||
std::vector<ObjectID> v2;
|
||||
};
|
||||
|
||||
TEST_F(TestRaylet, TestRayletCommands) {
|
||||
ray::Status status = ray::Status::OK();
|
||||
// TODO(atumanov): assert status is OK everywhere it's returned.
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "All connected clients:"
|
||||
<< "\n";
|
||||
status = mock_gcs_client->client_table().GetClientIds(
|
||||
[this](const std::vector<ClientID> &client_ids) {
|
||||
mock_gcs_client->client_table().GetClientInformationSet(
|
||||
client_ids,
|
||||
[this](const std::vector<ClientInformation> &info_vec) {
|
||||
for (const auto &info : info_vec) {
|
||||
RAY_LOG(INFO) << "ClientID=" << info.GetClientId().hex();
|
||||
RAY_LOG(INFO) << "ClientIp=" << info.GetIp();
|
||||
RAY_LOG(INFO) << "ClientPort=" << info.GetPort();
|
||||
}
|
||||
},
|
||||
[](Status status) {});
|
||||
});
|
||||
|
||||
sleep(1);
|
||||
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "Server client ids:"
|
||||
<< "\n";
|
||||
|
||||
status = server1->GetObjectManager().SubscribeObjAdded(
|
||||
[this](const ObjectID &object_id) { object_added_handler_1(object_id); });
|
||||
ASSERT_TRUE(status.ok());
|
||||
|
||||
status = server2->GetObjectManager().SubscribeObjAdded(
|
||||
[this](const ObjectID &object_id) { object_added_handler_2(object_id); });
|
||||
ASSERT_TRUE(status.ok());
|
||||
|
||||
ClientID client_id_1 = server1->GetObjectManager().GetClientID();
|
||||
ClientID client_id_2 = server2->GetObjectManager().GetClientID();
|
||||
RAY_LOG(INFO) << "Server 1: " << client_id_1.hex();
|
||||
RAY_LOG(INFO) << "Server 2: " << client_id_2.hex();
|
||||
|
||||
sleep(1);
|
||||
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "Test bidirectional pull"
|
||||
<< "\n";
|
||||
for (int i = -1; ++i < 100;) {
|
||||
ObjectID oid1 = WriteDataToClient(client1, 100);
|
||||
ObjectID oid2 = WriteDataToClient(client2, 100);
|
||||
status = server1->GetObjectManager().Pull(oid2);
|
||||
status = server2->GetObjectManager().Pull(oid1);
|
||||
}
|
||||
sleep(1);
|
||||
RAY_LOG(INFO) << v1.size() << " " << v2.size();
|
||||
ASSERT_TRUE(v1.size() == v2.size());
|
||||
for (int i = -1; ++i < (int)v1.size();) {
|
||||
ASSERT_TRUE(std::find(v1.begin(), v1.end(), v2[i]) != v1.end());
|
||||
}
|
||||
v1.clear();
|
||||
v2.clear();
|
||||
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "Test pull 1 from 2"
|
||||
<< "\n";
|
||||
for (int i = -1; ++i < 3;) {
|
||||
ObjectID oid2 = WriteDataToClient(client2, 100);
|
||||
status = server1->GetObjectManager().Pull(oid2);
|
||||
}
|
||||
sleep(1);
|
||||
RAY_LOG(INFO) << v1.size() << " " << v2.size();
|
||||
ASSERT_TRUE(v1.size() == v2.size());
|
||||
for (int i = -1; ++i < (int)v1.size();) {
|
||||
ASSERT_TRUE(std::find(v1.begin(), v1.end(), v2[i]) != v1.end());
|
||||
}
|
||||
v1.clear();
|
||||
v2.clear();
|
||||
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "Test pull 2 from 1"
|
||||
<< "\n";
|
||||
for (int i = -1; ++i < 3;) {
|
||||
ObjectID oid1 = WriteDataToClient(client1, 100);
|
||||
status = server2->GetObjectManager().Pull(oid1);
|
||||
}
|
||||
sleep(1);
|
||||
RAY_LOG(INFO) << v1.size() << " " << v2.size();
|
||||
ASSERT_TRUE(v1.size() == v2.size());
|
||||
for (int i = -1; ++i < (int)v1.size();) {
|
||||
ASSERT_TRUE(std::find(v1.begin(), v1.end(), v2[i]) != v1.end());
|
||||
}
|
||||
v1.clear();
|
||||
v2.clear();
|
||||
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "Test push 1 to 2"
|
||||
<< "\n";
|
||||
for (int i = -1; ++i < 3;) {
|
||||
ObjectID oid1 = WriteDataToClient(client1, 100);
|
||||
status = server1->GetObjectManager().Push(oid1, client_id_2);
|
||||
}
|
||||
sleep(1);
|
||||
RAY_LOG(INFO) << v1.size() << " " << v2.size();
|
||||
ASSERT_TRUE(v1.size() == v2.size());
|
||||
for (int i = -1; ++i < (int)v1.size();) {
|
||||
ASSERT_TRUE(std::find(v1.begin(), v1.end(), v2[i]) != v1.end());
|
||||
}
|
||||
v1.clear();
|
||||
v2.clear();
|
||||
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "Test push 2 to 1"
|
||||
<< "\n";
|
||||
for (int i = -1; ++i < 3;) {
|
||||
ObjectID oid2 = WriteDataToClient(client2, 100);
|
||||
status = server2->GetObjectManager().Push(oid2, client_id_1);
|
||||
}
|
||||
sleep(1);
|
||||
RAY_LOG(INFO) << v1.size() << " " << v2.size();
|
||||
ASSERT_TRUE(v1.size() == v2.size());
|
||||
for (int i = -1; ++i < (int)v1.size();) {
|
||||
ASSERT_TRUE(std::find(v1.begin(), v1.end(), v2[i]) != v1.end());
|
||||
}
|
||||
v1.clear();
|
||||
v2.clear();
|
||||
|
||||
RAY_LOG(INFO) << "\n"
|
||||
<< "Test bidirectional push"
|
||||
<< "\n";
|
||||
for (int i = -1; ++i < 3;) {
|
||||
ObjectID oid1 = WriteDataToClient(client1, 100);
|
||||
ObjectID oid2 = WriteDataToClient(client2, 100);
|
||||
status = server1->GetObjectManager().Push(oid1, client_id_2);
|
||||
status = server2->GetObjectManager().Push(oid2, client_id_1);
|
||||
}
|
||||
sleep(1);
|
||||
RAY_LOG(INFO) << v1.size() << " " << v2.size();
|
||||
ASSERT_TRUE(v1.size() == v2.size());
|
||||
for (int i = -1; ++i < (int)v1.size();) {
|
||||
ASSERT_TRUE(std::find(v1.begin(), v1.end(), v2[i]) != v1.end());
|
||||
}
|
||||
v1.clear();
|
||||
v2.clear();
|
||||
|
||||
ASSERT_TRUE(true);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
|
||||
int main(int argc, char **argv) {
|
||||
::testing::InitGoogleTest(&argc, argv);
|
||||
ray::test_executable = std::string(argv[0]);
|
||||
return RUN_ALL_TESTS();
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
#include "reconstruction_policy.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
void ReconstructionPolicy::CheckObjectReconstruction(const ObjectID &object) {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
|
||||
} // end namespace ray
|
||||
@@ -0,0 +1,34 @@
|
||||
#ifndef RAY_RAYLET_RECONSTRUCTION_POLICY_H
|
||||
#define RAY_RAYLET_RECONSTRUCTION_POLICY_H
|
||||
|
||||
#include <functional>
|
||||
|
||||
#include "ray/id.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
// TODO(swang): Use std::function instead of boost.
|
||||
|
||||
class ReconstructionPolicy {
|
||||
public:
|
||||
/// Create the reconstruction policy.
|
||||
///
|
||||
/// \param reconstruction_handler The handler to call if a task needs to be
|
||||
/// re-executed.
|
||||
// TODO(swang): This requires at minimum references to the Raylet's lineage
|
||||
// cache and GCS client.
|
||||
ReconstructionPolicy(std::function<void(const TaskID &)> reconstruction_handler) {}
|
||||
|
||||
/// Check whether an object requires reconstruction. If this object requires
|
||||
/// reconstruction, the registered task reconstruction handler will be called
|
||||
/// for each task that needs to be re-executed.
|
||||
///
|
||||
/// \param object_id The object to check for reconstruction.
|
||||
void CheckObjectReconstruction(const ObjectID &object_id);
|
||||
|
||||
private:
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_RECONSTRUCTION_POLICY_H
|
||||
@@ -0,0 +1,42 @@
|
||||
#include <iostream>
|
||||
|
||||
#include "ray/raylet/raylet.h"
|
||||
|
||||
/// A demo that starts two Raylets, with one object store each. The two Raylets
|
||||
/// share a mock GCS client for communication between the two (e.g., for
|
||||
/// ObjectManager::Push).
|
||||
int main(int argc, char *argv[]) {
|
||||
RAY_CHECK(argc == 3);
|
||||
std::string store1 = "/tmp/store1";
|
||||
std::string store2 = "/tmp/store2";
|
||||
// start store
|
||||
std::string plasma_dir = "../../plasma";
|
||||
std::string plasma_command1 = plasma_dir + "/plasma_store -m 1000000000 -s ";
|
||||
std::string plasma_command2 = " 1> /dev/null 2> /dev/null &";
|
||||
RAY_LOG(INFO) << plasma_command1 << store1 << plasma_command2;
|
||||
RAY_LOG(INFO) << plasma_command1 << store2 << plasma_command2;
|
||||
int s;
|
||||
s = system((plasma_command1 + store1 + plasma_command2).c_str());
|
||||
RAY_CHECK(s == 0);
|
||||
s = system((plasma_command1 + store2 + plasma_command2).c_str());
|
||||
|
||||
// configure
|
||||
std::unordered_map<std::string, double> static_resource_conf;
|
||||
static_resource_conf = {{"CPU", 1}, {"GPU", 1}};
|
||||
ray::ResourceSet resource_config(std::move(static_resource_conf));
|
||||
ray::ObjectManagerConfig om_config;
|
||||
|
||||
// initialize mock gcs & object directory
|
||||
std::shared_ptr<ray::GcsClient> mock_gcs_client =
|
||||
std::shared_ptr<ray::GcsClient>(new ray::GcsClient());
|
||||
|
||||
// Initialize the node manager.
|
||||
boost::asio::io_service io_service;
|
||||
om_config.store_socket_name = store1;
|
||||
ray::Raylet server1(io_service, std::string(argv[1]), resource_config, om_config,
|
||||
mock_gcs_client);
|
||||
om_config.store_socket_name = store2;
|
||||
ray::Raylet server2(io_service, std::string(argv[2]), resource_config, om_config,
|
||||
mock_gcs_client);
|
||||
io_service.run();
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
#include "scheduling_policy.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
SchedulingPolicy::SchedulingPolicy(const SchedulingQueue &scheduling_queue)
|
||||
: scheduling_queue_(scheduling_queue) {}
|
||||
|
||||
std::unordered_map<TaskID, ClientID, UniqueIDHasher> SchedulingPolicy::Schedule(
|
||||
const std::unordered_map<ClientID, SchedulingResources, UniqueIDHasher>
|
||||
&cluster_resources) {
|
||||
static ClientID local_node_id = ClientID::nil();
|
||||
std::unordered_map<TaskID, ClientID, UniqueIDHasher> decision;
|
||||
// TODO(atumanov): consider all cluster resources.
|
||||
SchedulingResources resource_supply = cluster_resources.at(local_node_id);
|
||||
const auto &resource_supply_set = resource_supply.GetAvailableResources();
|
||||
|
||||
// Iterate over running tasks, get their resource demand and try to schedule.
|
||||
for (const auto &t : scheduling_queue_.GetReadyTasks()) {
|
||||
// Get task's resource demand
|
||||
const auto &resource_demand = t.GetTaskSpecification().GetRequiredResources();
|
||||
bool task_feasible = resource_demand.IsSubset(resource_supply_set);
|
||||
if (task_feasible) {
|
||||
const TaskID &task_id = t.GetTaskSpecification().TaskId();
|
||||
decision[task_id] = local_node_id;
|
||||
}
|
||||
}
|
||||
return decision;
|
||||
}
|
||||
|
||||
SchedulingPolicy::~SchedulingPolicy() {}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,42 @@
|
||||
#ifndef RAY_RAYLET_SCHEDULING_POLICY_H
|
||||
#define RAY_RAYLET_SCHEDULING_POLICY_H
|
||||
|
||||
#include <unordered_map>
|
||||
|
||||
#include "ray/raylet/scheduling_queue.h"
|
||||
#include "ray/raylet/scheduling_resources.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// \class SchedulingPolicy
|
||||
/// \brief Implements a scheduling policy for the node manager.
|
||||
class SchedulingPolicy {
|
||||
public:
|
||||
/// \brief SchedulingPolicy constructor.
|
||||
///
|
||||
/// \param scheduling_queue: reference to a scheduler queues object for access to
|
||||
/// tasks.
|
||||
/// \return None.
|
||||
SchedulingPolicy(const SchedulingQueue &scheduling_queue);
|
||||
|
||||
/// Perform a scheduling operation, given a set of cluster resources and
|
||||
/// producing a mapping of tasks to node managers.
|
||||
///
|
||||
/// \param cluster_resources: a set of cluster resources representing
|
||||
/// configured and current resource capacity on each node.
|
||||
/// \return Scheduling decision, mapping tasks to node managers for placement.
|
||||
std::unordered_map<TaskID, ClientID, UniqueIDHasher> Schedule(
|
||||
const std::unordered_map<ClientID, SchedulingResources, UniqueIDHasher>
|
||||
&cluster_resources);
|
||||
|
||||
/// \brief SchedulingPolicy destructor.
|
||||
virtual ~SchedulingPolicy();
|
||||
|
||||
private:
|
||||
/// An immutable reference to the scheduling task queues.
|
||||
const SchedulingQueue &scheduling_queue_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_SCHEDULING_POLICY_H
|
||||
@@ -0,0 +1,91 @@
|
||||
#include "scheduling_queue.h"
|
||||
|
||||
#include "ray/status.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
const std::list<Task> &SchedulingQueue::GetWaitingTasks() const {
|
||||
return this->waiting_tasks_;
|
||||
}
|
||||
|
||||
const std::list<Task> &SchedulingQueue::GetReadyTasks() const {
|
||||
return this->ready_tasks_;
|
||||
}
|
||||
|
||||
const std::list<Task> &SchedulingQueue::GetScheduledTasks() const {
|
||||
return this->scheduled_tasks_;
|
||||
}
|
||||
|
||||
const std::list<Task> &SchedulingQueue::GetRunningTasks() const {
|
||||
return this->running_tasks_;
|
||||
}
|
||||
|
||||
const std::list<Task> &SchedulingQueue::GetReadyMethods() const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
|
||||
// Helper function to remove tasks in the given set of task_ids from a
|
||||
// queue, and append them to the given vector removed_tasks.
|
||||
void removeTasksFromQueue(std::list<Task> &queue,
|
||||
std::unordered_set<TaskID, UniqueIDHasher> &task_ids,
|
||||
std::vector<Task> &removed_tasks) {
|
||||
for (auto it = queue.begin(); it != queue.end();) {
|
||||
auto task_id = task_ids.find(it->GetTaskSpecification().TaskId());
|
||||
if (task_id != task_ids.end()) {
|
||||
task_ids.erase(task_id);
|
||||
removed_tasks.push_back(std::move(*it));
|
||||
it = queue.erase(it);
|
||||
} else {
|
||||
it++;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Helper function to queue the given tasks to the given queue.
|
||||
void queueTasks(std::list<Task> &queue, const std::vector<Task> &tasks) {
|
||||
for (auto &task : tasks) {
|
||||
queue.push_back(task);
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<Task> SchedulingQueue::RemoveTasks(
|
||||
std::unordered_set<TaskID, UniqueIDHasher> task_ids) {
|
||||
// List of removed tasks to be returned.
|
||||
std::vector<Task> removed_tasks;
|
||||
|
||||
// Try to find the tasks to remove from the waiting tasks.
|
||||
removeTasksFromQueue(waiting_tasks_, task_ids, removed_tasks);
|
||||
removeTasksFromQueue(ready_tasks_, task_ids, removed_tasks);
|
||||
removeTasksFromQueue(scheduled_tasks_, task_ids, removed_tasks);
|
||||
removeTasksFromQueue(running_tasks_, task_ids, removed_tasks);
|
||||
// TODO(swang): Remove from running methods.
|
||||
|
||||
RAY_CHECK(task_ids.size() == 0);
|
||||
return removed_tasks;
|
||||
}
|
||||
|
||||
void SchedulingQueue::QueueWaitingTasks(const std::vector<Task> &tasks) {
|
||||
queueTasks(waiting_tasks_, tasks);
|
||||
}
|
||||
|
||||
void SchedulingQueue::QueueReadyTasks(const std::vector<Task> &tasks) {
|
||||
queueTasks(ready_tasks_, tasks);
|
||||
}
|
||||
|
||||
void SchedulingQueue::QueueScheduledTasks(const std::vector<Task> &tasks) {
|
||||
queueTasks(scheduled_tasks_, tasks);
|
||||
}
|
||||
|
||||
void SchedulingQueue::QueueRunningTasks(const std::vector<Task> &tasks) {
|
||||
queueTasks(running_tasks_, tasks);
|
||||
}
|
||||
|
||||
// RegisterActor is responsible for recording provided actor_information
|
||||
// in the actor registry.
|
||||
bool SchedulingQueue::RegisterActor(ActorID actor_id,
|
||||
const ActorInformation &actor_information) {
|
||||
actor_registry_[actor_id] = actor_information;
|
||||
return true;
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,108 @@
|
||||
#ifndef RAY_RAYLET_SCHEDULING_QUEUE_H
|
||||
#define RAY_RAYLET_SCHEDULING_QUEUE_H
|
||||
|
||||
#include <list>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
#include <vector>
|
||||
|
||||
#include "ray/raylet/actor.h"
|
||||
#include "ray/raylet/task.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// \class SchedulingQueue
|
||||
///
|
||||
/// Encapsulates task queues. Each queue represents a scheduling state for a
|
||||
/// task. The scheduling state is one of (1) waiting: for object dependencies
|
||||
/// to become available, (2) ready: object dependencies are available and the
|
||||
/// task is ready to be scheduled, (3) scheduled: the task has been scheduled
|
||||
/// but is waiting for a worker, or (4) running: the task has been scheduled
|
||||
/// and is running on a worker.
|
||||
class SchedulingQueue {
|
||||
public:
|
||||
/// Create a scheduling queue.
|
||||
SchedulingQueue() {}
|
||||
|
||||
/// SchedulingQueue destructor.
|
||||
virtual ~SchedulingQueue() {}
|
||||
|
||||
/// Get the queue of tasks in the waiting state.
|
||||
///
|
||||
/// \return A const reference to the queue of tasks that are waiting for
|
||||
/// object dependencies to become available.
|
||||
const std::list<Task> &GetWaitingTasks() const;
|
||||
|
||||
/// Get the queue of tasks in the ready state.
|
||||
///
|
||||
/// \return A const reference to the queue of tasks that have all
|
||||
/// dependencies local and that are waiting to be scheduled.
|
||||
const std::list<Task> &GetReadyTasks() const;
|
||||
|
||||
/// Get the queue of actor methods in the ready state.
|
||||
///
|
||||
/// \return A const reference to the queue of actor methods that have all
|
||||
/// dependencies local and that are waiting to be scheduled.
|
||||
const std::list<Task> &GetReadyMethods() const;
|
||||
|
||||
/// Get the queue of tasks in the scheduled state.
|
||||
///
|
||||
/// \return A const reference to the queue of tasks that have been scheduled
|
||||
/// to execute but that are waiting for a worker.
|
||||
const std::list<Task> &GetScheduledTasks() const;
|
||||
|
||||
/// Get the queue of tasks in the running state.
|
||||
///
|
||||
/// \return A const reference to the queue of tasks that are currently
|
||||
/// executing on a worker.
|
||||
const std::list<Task> &GetRunningTasks() const;
|
||||
|
||||
/// Remove tasks from the task queue.
|
||||
///
|
||||
/// \param tasks The set of task IDs to remove from the queue. The
|
||||
/// corresponding tasks must be contained in the queue.
|
||||
/// \return A vector of the tasks that were removed.
|
||||
std::vector<Task> RemoveTasks(std::unordered_set<TaskID, UniqueIDHasher> tasks);
|
||||
|
||||
/// Queue tasks in the waiting state.
|
||||
///
|
||||
/// \param tasks The tasks to queue.
|
||||
void QueueWaitingTasks(const std::vector<Task> &tasks);
|
||||
|
||||
/// Queue tasks in the ready state.
|
||||
///
|
||||
/// \param tasks The tasks to queue.
|
||||
void QueueReadyTasks(const std::vector<Task> &tasks);
|
||||
|
||||
/// Queue tasks in the scheduled state.
|
||||
///
|
||||
/// \param tasks The tasks to queue.
|
||||
void QueueScheduledTasks(const std::vector<Task> &tasks);
|
||||
|
||||
/// Queue tasks in the running state.
|
||||
///
|
||||
/// \param tasks The tasks to queue.
|
||||
void QueueRunningTasks(const std::vector<Task> &tasks);
|
||||
|
||||
/// Register an actor.
|
||||
///
|
||||
/// \param actor_id The ID of the actor to register.
|
||||
/// \param actor_information Information about the actor.
|
||||
bool RegisterActor(ActorID actor_id, const ActorInformation &actor_information);
|
||||
|
||||
private:
|
||||
/// Tasks that are waiting for an object dependency to appear locally.
|
||||
std::list<Task> waiting_tasks_;
|
||||
/// Tasks whose object dependencies are locally available, but that are
|
||||
/// waiting to be scheduled.
|
||||
std::list<Task> ready_tasks_;
|
||||
/// Tasks that have been scheduled to run, but that are waiting for a worker.
|
||||
std::list<Task> scheduled_tasks_;
|
||||
/// Tasks that are running on a worker.
|
||||
std::list<Task> running_tasks_;
|
||||
/// The registry of known actors.
|
||||
std::unordered_map<ActorID, ActorInformation, UniqueIDHasher> actor_registry_;
|
||||
};
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_SCHEDULING_QUEUE_H
|
||||
@@ -0,0 +1,106 @@
|
||||
#include "scheduling_resources.h"
|
||||
|
||||
#include <cmath>
|
||||
|
||||
namespace ray {
|
||||
|
||||
ResourceSet::ResourceSet() {}
|
||||
|
||||
ResourceSet::ResourceSet(const std::unordered_map<std::string, double> &resource_map)
|
||||
: resource_capacity_(resource_map) {}
|
||||
|
||||
ResourceSet::~ResourceSet() {}
|
||||
|
||||
bool ResourceSet::operator==(const ResourceSet &rhs) const {
|
||||
return (this->IsSubset(rhs) && rhs.IsSubset(*this));
|
||||
}
|
||||
|
||||
bool ResourceSet::IsSubset(const ResourceSet &other) const {
|
||||
// Check to make sure all keys of this are in other.
|
||||
for (const auto &resource_pair : resource_capacity_) {
|
||||
const auto &resource_name = resource_pair.first;
|
||||
const double lhs_quantity = resource_pair.second;
|
||||
double rhs_quantity = 0;
|
||||
if (!other.GetResource(resource_name, &rhs_quantity)) {
|
||||
// Resource not found in rhs, therefore lhs is not a subset of rhs.
|
||||
return false;
|
||||
}
|
||||
if (lhs_quantity > rhs_quantity) {
|
||||
// Resource found in rhs, but lhs capacity exceeds rhs capacity.
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
/// Test whether this ResourceSet is a superset of the other ResourceSet
|
||||
bool ResourceSet::IsSuperset(const ResourceSet &other) const {
|
||||
return other.IsSubset(*this);
|
||||
}
|
||||
/// Test whether this ResourceSet is precisely equal to the other ResourceSet.
|
||||
bool ResourceSet::IsEqual(const ResourceSet &rhs) const {
|
||||
return (this->IsSubset(rhs) && rhs.IsSubset(*this));
|
||||
}
|
||||
|
||||
bool ResourceSet::AddResource(const std::string &resource_name, double capacity) {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
bool ResourceSet::RemoveResource(const std::string &resource_name) {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
bool ResourceSet::SubtractResources(const ResourceSet &other) {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
bool ResourceSet::AddResources(const ResourceSet &other) {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
|
||||
bool ResourceSet::GetResource(const std::string &resource_name, double *value) const {
|
||||
if (!value) {
|
||||
return false;
|
||||
}
|
||||
if (this->resource_capacity_.count(resource_name) == 0) {
|
||||
*value = std::nan("");
|
||||
return false;
|
||||
}
|
||||
*value = this->resource_capacity_.at(resource_name);
|
||||
return true;
|
||||
}
|
||||
|
||||
/// SchedulingResources class implementation
|
||||
|
||||
SchedulingResources::SchedulingResources()
|
||||
: resources_total_(ResourceSet()), resources_available_(ResourceSet()) {}
|
||||
|
||||
SchedulingResources::SchedulingResources(const ResourceSet &total)
|
||||
: resources_total_(total), resources_available_(total) {}
|
||||
|
||||
SchedulingResources::~SchedulingResources() {}
|
||||
|
||||
ResourceAvailabilityStatus SchedulingResources::CheckResourcesSatisfied(
|
||||
ResourceSet &resources) const {
|
||||
if (!resources.IsSubset(this->resources_total_)) {
|
||||
return kInfeasible;
|
||||
}
|
||||
// Resource demand specified is feasible. Check if it's available.
|
||||
if (!resources.IsSubset(this->resources_available_)) {
|
||||
return kResourcesUnavailable;
|
||||
}
|
||||
return kFeasible;
|
||||
}
|
||||
|
||||
const ResourceSet &SchedulingResources::GetAvailableResources() const {
|
||||
return this->resources_available_;
|
||||
}
|
||||
|
||||
// Return specified resources back to SchedulingResources.
|
||||
bool SchedulingResources::Release(const ResourceSet &resources) {
|
||||
return this->resources_available_.AddResources(resources);
|
||||
}
|
||||
|
||||
// Take specified resources from SchedulingResources.
|
||||
bool SchedulingResources::Acquire(const ResourceSet &resources) {
|
||||
return this->resources_available_.SubtractResources(resources);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,150 @@
|
||||
#ifndef RAY_RAYLET_SCHEDULING_RESOURCES_H
|
||||
#define RAY_RAYLET_SCHEDULING_RESOURCES_H
|
||||
|
||||
#include <string>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// Resource availability status reports whether the resource requirement is
|
||||
/// (1) infeasible, (2) feasible but currently unavailable, or (3) available.
|
||||
typedef enum {
|
||||
kInfeasible, ///< Cannot ever satisfy resource requirements.
|
||||
kResourcesUnavailable, ///< Feasible, but not currently available.
|
||||
kFeasible ///< Feasible and currently available.
|
||||
} ResourceAvailabilityStatus;
|
||||
|
||||
/// \class ResourceSet
|
||||
/// \brief Encapsulates and operates on a set of resources, including CPUs,
|
||||
/// GPUs, and custom labels.
|
||||
class ResourceSet {
|
||||
public:
|
||||
/// \brief empty ResourceSet constructor.
|
||||
ResourceSet();
|
||||
|
||||
/// \brief Constructs ResourceSet from the specified resource map.
|
||||
ResourceSet(const std::unordered_map<std::string, double> &resource_map);
|
||||
|
||||
/// \brief Empty ResourceSet destructor.
|
||||
~ResourceSet();
|
||||
|
||||
/// \brief Test equality with the other specified ResourceSet object.
|
||||
///
|
||||
/// \param rhs: Right-hand side object for equality comparison.
|
||||
/// \return True if objects are equal, False otherwise.
|
||||
bool operator==(const ResourceSet &rhs) const;
|
||||
|
||||
/// \brief Test equality with the other specified ResourceSet object.
|
||||
///
|
||||
/// \param other: Right-hand side object for equality comparison.
|
||||
/// \return True if objects are equal, False otherwise.
|
||||
bool IsEqual(const ResourceSet &other) const;
|
||||
|
||||
/// \brief Test whether this ResourceSet is a subset of the other ResourceSet.
|
||||
///
|
||||
/// \param other: The resource set we check being a subset of.
|
||||
/// \return True if the current resource set is the subset of other. False
|
||||
/// otherwise.
|
||||
bool IsSubset(const ResourceSet &other) const;
|
||||
|
||||
/// \brief Test if this ResourceSet is a superset of the other ResourceSet.
|
||||
///
|
||||
/// \param other: The resource set we check being a superset of.
|
||||
/// \return True if the current resource set is the superset of other.
|
||||
/// False otherwise.
|
||||
bool IsSuperset(const ResourceSet &other) const;
|
||||
|
||||
/// \brief Add a new resource to the resource set.
|
||||
///
|
||||
/// \param resource_name: name/label of the resource to add.
|
||||
/// \param capacity: numeric capacity value for the resource to add.
|
||||
/// \return True, if the resource was successfully added. False otherwise.
|
||||
bool AddResource(const std::string &resource_name, double capacity);
|
||||
|
||||
/// \brief Remove the specified resource from the resource set.
|
||||
///
|
||||
/// \param resource_name: name/label of the resource to remove.
|
||||
/// \return True, if the resource was successfully removed. False otherwise.
|
||||
bool RemoveResource(const std::string &resource_name);
|
||||
|
||||
/// \brief Add a set of resources to the current set of resources.
|
||||
///
|
||||
/// \param other: The other resource set to add.
|
||||
/// \return True if the resource set was added successfully. False otherwise.
|
||||
bool AddResources(const ResourceSet &other);
|
||||
|
||||
/// \brief Subtract a set of resources from the current set of resources.
|
||||
///
|
||||
/// \param other: The resource set to subtract from the current resource set.
|
||||
/// \return True if the resource set was subtracted successfully.
|
||||
/// False otherwise.
|
||||
bool SubtractResources(const ResourceSet &other);
|
||||
|
||||
/// Return the capacity value associated with the specified resource.
|
||||
///
|
||||
/// \param resource_name: Resource name for which capacity is requested.
|
||||
/// \param[out] value: Resource capacity value.
|
||||
/// \return True if the resource capacity value was successfully retrieved.
|
||||
/// False otherwise.
|
||||
bool GetResource(const std::string &resource_name, double *value) const;
|
||||
|
||||
private:
|
||||
/// Resource capacity map.
|
||||
std::unordered_map<std::string, double> resource_capacity_;
|
||||
};
|
||||
|
||||
/// \class SchedulingResources
|
||||
/// SchedulingResources class encapsulates the state of all local resources and
|
||||
/// manages accounting of those resources. Resources include configured resource
|
||||
/// bundle capacity, and GPU allocation map.
|
||||
class SchedulingResources {
|
||||
public:
|
||||
/// SchedulingResources constructor: sets configured and available resources
|
||||
/// to an empty set.
|
||||
SchedulingResources();
|
||||
|
||||
/// SchedulingResources constructor: sets available and configured capacity
|
||||
/// to the resource set specified.
|
||||
///
|
||||
/// \param total: The amount of total configured capacity.
|
||||
SchedulingResources(const ResourceSet &total);
|
||||
|
||||
/// \brief SchedulingResources destructor.
|
||||
~SchedulingResources();
|
||||
|
||||
/// \brief Check if the specified resource request can be satisfied.
|
||||
///
|
||||
/// \param set: The set of resources representing the resource request.
|
||||
/// \return Availability status that specifies if the requested resource set
|
||||
/// is feasible, infeasible, or feasible but unavailable.
|
||||
ResourceAvailabilityStatus CheckResourcesSatisfied(ResourceSet &set) const;
|
||||
|
||||
/// \brief Request the set and capacity of resources currently available.
|
||||
///
|
||||
/// \return Immutable set of resources with currently available capacity.
|
||||
const ResourceSet &GetAvailableResources() const;
|
||||
|
||||
/// \brief Release the amount of resources specified.
|
||||
///
|
||||
/// \param resources: the amount of resources to be released.
|
||||
/// \return True if resources were successfully released. False otherwise.
|
||||
bool Release(const ResourceSet &resources);
|
||||
|
||||
/// \brief Acquire the amount of resources specified.
|
||||
///
|
||||
/// \param resources: the amount of resources to be acquired.
|
||||
/// \return True if resources were successfully acquired. False otherwise.
|
||||
bool Acquire(const ResourceSet &resources);
|
||||
|
||||
private:
|
||||
/// Static resource configuration (e.g., static_resources).
|
||||
ResourceSet resources_total_;
|
||||
/// Dynamic resource capacity (e.g., dynamic_resources).
|
||||
ResourceSet resources_available_;
|
||||
/// gpu_map - replace with ResourceMap (for generality).
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_SCHEDULING_RESOURCES_H
|
||||
@@ -0,0 +1,50 @@
|
||||
#include "task.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
const TaskExecutionSpecification &Task::GetTaskExecutionSpec() const {
|
||||
return task_execution_spec_;
|
||||
}
|
||||
|
||||
const TaskSpecification &Task::GetTaskSpecification() const { return task_spec_; }
|
||||
|
||||
const std::vector<ObjectID> Task::GetDependencies() const {
|
||||
std::vector<ObjectID> dependencies;
|
||||
for (int i = 0; i < task_spec_.NumArgs(); ++i) {
|
||||
int count = task_spec_.ArgIdCount(i);
|
||||
for (int j = 0; j < count; j++) {
|
||||
dependencies.push_back(task_spec_.ArgId(i, j));
|
||||
}
|
||||
}
|
||||
// TODO(atumanov): why not just return a const reference to ExecutionDependencies() and
|
||||
// avoid a copy.
|
||||
auto execution_dependencies = task_execution_spec_.ExecutionDependencies();
|
||||
dependencies.insert(dependencies.end(), execution_dependencies.begin(),
|
||||
execution_dependencies.end());
|
||||
return dependencies;
|
||||
}
|
||||
|
||||
bool Task::DependsOn(const ObjectID &object_id) const {
|
||||
// Iterate through the task arguments to see if it contains object_id.
|
||||
int64_t num_args = task_spec_.NumArgs();
|
||||
for (int i = 0; i < num_args; ++i) {
|
||||
int count = task_spec_.ArgIdCount(i);
|
||||
for (int j = 0; j < count; j++) {
|
||||
ObjectID arg_id = task_spec_.ArgId(i, j);
|
||||
if (arg_id == object_id) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
// Iterate through the execution dependencies to see if it contains object_id.
|
||||
for (const auto &dependency_id : task_execution_spec_.ExecutionDependencies()) {
|
||||
if (dependency_id == object_id) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
// The requested object ID was not a task argument or an execution dependency.
|
||||
// This task is not dependent on it.
|
||||
return false;
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,69 @@
|
||||
#ifndef RAY_RAYLET_TASK_H
|
||||
#define RAY_RAYLET_TASK_H
|
||||
|
||||
#include <inttypes.h>
|
||||
|
||||
#include "ray/raylet/task_execution_spec.h"
|
||||
#include "ray/raylet/task_spec.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// \class Task
|
||||
///
|
||||
/// A Task represents a Ray task and a specification of its execution (e.g.,
|
||||
/// resource demands). The task's specification contains both immutable fields,
|
||||
/// determined at submission time, and mutable fields, determined at execution
|
||||
/// time.
|
||||
class Task {
|
||||
public:
|
||||
/// Create a task.
|
||||
///
|
||||
/// \param execution_spec The execution specification for the task. These are
|
||||
/// the mutable fields in the task specification that may change at task
|
||||
/// execution time.
|
||||
/// \param task_spec The immutable specification for the task. These fields
|
||||
/// are determined at task submission time.
|
||||
Task(const TaskExecutionSpecification &execution_spec,
|
||||
const TaskSpecification &task_spec)
|
||||
: task_execution_spec_(execution_spec), task_spec_(task_spec) {}
|
||||
|
||||
/// Destroy the task.
|
||||
virtual ~Task() {}
|
||||
|
||||
/// Get the execution specification for the task.
|
||||
///
|
||||
/// \return The mutable specification for the task.
|
||||
const TaskExecutionSpecification &GetTaskExecutionSpec() const;
|
||||
|
||||
/// Get the immutable specification for the task.
|
||||
///
|
||||
/// \return The immutable specification for the task.
|
||||
const TaskSpecification &GetTaskSpecification() const;
|
||||
|
||||
/// Get the task's object dependencies. This comprises the immutable task
|
||||
/// arguments and the mutable execution dependencies.
|
||||
///
|
||||
/// \return The object dependencies.
|
||||
/// TODO(atumanov): consider returning a constant reference.
|
||||
const std::vector<ObjectID> GetDependencies() const;
|
||||
|
||||
/// Compute whether the task is dependent on an object ID.
|
||||
///
|
||||
/// \param object_id The object ID that the task may be dependent on.
|
||||
/// \return Returns true if the task is dependent on the given object ID and
|
||||
/// false otherwise.
|
||||
bool DependsOn(const ObjectID &object_id) const;
|
||||
|
||||
private:
|
||||
/// Task execution specification, consisting of all dynamic/mutable
|
||||
/// information about this task determined at execution time..
|
||||
TaskExecutionSpecification task_execution_spec_;
|
||||
/// Task specification object, consisting of immutable information about this
|
||||
/// task determined at submission time. Includes resource demand, object
|
||||
/// dependencies, etc.
|
||||
TaskSpecification task_spec_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_TASK_H
|
||||
@@ -0,0 +1,107 @@
|
||||
#include "task_dependency_manager.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
TaskDependencyManager::TaskDependencyManager(
|
||||
ObjectManager &object_manager,
|
||||
// ReconstructionPolicy &reconstruction_policy,
|
||||
std::function<void(const TaskID &)> handler)
|
||||
: object_manager_(object_manager),
|
||||
// reconstruction_policy_(reconstruction_policy),
|
||||
task_ready_callback_(handler) {
|
||||
// TODO(swang): Check return status.
|
||||
ray::Status status = object_manager_.SubscribeObjAdded(
|
||||
[this](const ObjectID &object_id) { handleObjectReady(object_id); });
|
||||
// TODO(swang): Subscribe to object removed notifications.
|
||||
}
|
||||
|
||||
bool TaskDependencyManager::argumentsReady(const std::vector<ObjectID> arguments) const {
|
||||
for (auto &argument : arguments) {
|
||||
// Check if any argument is missing.
|
||||
if (local_objects_.count(argument) == 0) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
// All arguments are ready.
|
||||
return true;
|
||||
}
|
||||
|
||||
void TaskDependencyManager::handleObjectReady(const ray::ObjectID &object_id) {
|
||||
RAY_LOG(DEBUG) << "object ready " << object_id.hex();
|
||||
// Add the object to the table of locally available objects.
|
||||
RAY_CHECK(local_objects_.count(object_id) == 0);
|
||||
local_objects_.insert(object_id);
|
||||
|
||||
// Handle any tasks that were dependent on the newly available object.
|
||||
std::vector<TaskID> ready_task_ids;
|
||||
auto dependent_tasks = remote_object_dependencies_.find(object_id);
|
||||
if (dependent_tasks != remote_object_dependencies_.end()) {
|
||||
for (auto &dependent_task_id : dependent_tasks->second) {
|
||||
// If the dependent task now has all of its arguments ready, it's ready
|
||||
// to run.
|
||||
if (argumentsReady(task_dependencies_[dependent_task_id])) {
|
||||
ready_task_ids.push_back(dependent_task_id);
|
||||
}
|
||||
}
|
||||
remote_object_dependencies_.erase(dependent_tasks);
|
||||
}
|
||||
// Process callbacks for all of the tasks dependent on the object that are
|
||||
// now ready to run.
|
||||
for (auto &ready_task_id : ready_task_ids) {
|
||||
UnsubscribeTaskReady(ready_task_id);
|
||||
task_ready_callback_(ready_task_id);
|
||||
}
|
||||
}
|
||||
|
||||
bool TaskDependencyManager::TaskReady(const Task &task) const {
|
||||
const std::vector<ObjectID> arguments = task.GetDependencies();
|
||||
return argumentsReady(arguments);
|
||||
}
|
||||
|
||||
void TaskDependencyManager::SubscribeTaskReady(const Task &task) {
|
||||
TaskID task_id = task.GetTaskSpecification().TaskId();
|
||||
const std::vector<ObjectID> arguments = task.GetDependencies();
|
||||
// Add the task's arguments to the table of subscribed tasks.
|
||||
task_dependencies_[task_id] = arguments;
|
||||
// Add the task's remote arguments to the table of remote objects.
|
||||
int num_missing_arguments = 0;
|
||||
for (auto &argument : arguments) {
|
||||
if (local_objects_.count(argument) == 0) {
|
||||
remote_object_dependencies_[argument].push_back(task_id);
|
||||
num_missing_arguments++;
|
||||
// TODO(swang): Check return status.
|
||||
// TODO(swang): Handle Pull failure (if object manager does not retry).
|
||||
// TODO(atumanov): pull return status should be propagated back to the caller.
|
||||
ray::Status status = object_manager_.Pull(argument);
|
||||
}
|
||||
}
|
||||
// Check that the task has some missing arguments.
|
||||
RAY_CHECK(num_missing_arguments > 0);
|
||||
}
|
||||
|
||||
void TaskDependencyManager::UnsubscribeTaskReady(const TaskID &task_id) {
|
||||
const std::vector<ObjectID> arguments = task_dependencies_[task_id];
|
||||
// Remove the task from the table of subscribed tasks.
|
||||
task_dependencies_.erase(task_id);
|
||||
// Remove the task from the table of remote objects to dependent tasks.
|
||||
for (auto &argument : arguments) {
|
||||
if (local_objects_.count(argument) == 1) {
|
||||
continue;
|
||||
}
|
||||
// The argument is not local. Remove the task from the list of tasks that
|
||||
// are dependent on this object.
|
||||
std::vector<TaskID> &dependent_tasks = remote_object_dependencies_[task_id];
|
||||
for (auto it = dependent_tasks.begin(); it != dependent_tasks.end(); it++) {
|
||||
if (*it == task_id) {
|
||||
it = dependent_tasks.erase(it);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void TaskDependencyManager::MarkDependencyReady(const ObjectID &object) {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,82 @@
|
||||
#ifndef RAY_RAYLET_TASK_DEPENDENCY_MANAGER_H
|
||||
#define RAY_RAYLET_TASK_DEPENDENCY_MANAGER_H
|
||||
|
||||
// clang-format off
|
||||
#include "ray/id.h"
|
||||
#include "ray/raylet/task.h"
|
||||
#include "ray/object_manager/object_manager.h"
|
||||
#include "ray/raylet/reconstruction_policy.h"
|
||||
// clang-format on
|
||||
|
||||
namespace ray {
|
||||
|
||||
class ReconstructionPolicy;
|
||||
|
||||
/// \class TaskDependencyManager
|
||||
///
|
||||
/// Responsible for managing task dependencies. Tasks that have object
|
||||
/// dependencies that are missing locally should contact the manager to try to
|
||||
/// make their dependencies available.
|
||||
class TaskDependencyManager {
|
||||
public:
|
||||
/// Create a task dependency manager.
|
||||
///
|
||||
/// \param object_manager A reference to the object manager so that the task
|
||||
/// dependency manager can issue requests to transfer objects.
|
||||
/// \param handler The handler to call for subscribed tasks whose
|
||||
/// dependencies have become available locally.
|
||||
TaskDependencyManager(ObjectManager &object_manager,
|
||||
// ReconstructionPolicy &reconstruction_policy,
|
||||
std::function<void(const TaskID &)> handler);
|
||||
|
||||
/// Check whether a task's object dependencies are locally available.
|
||||
///
|
||||
/// \param task The task whose object dependencies will be checked.
|
||||
/// \return Whether the task's object dependencies are ready.
|
||||
bool TaskReady(const Task &task) const;
|
||||
|
||||
/// Subscribe to a task that has missing dependencies. The manager will
|
||||
/// attempt to make any missing dependencies available locally by transfer or
|
||||
/// by reconstruction. The registered handler will be called when the task's
|
||||
/// dependencies become locally available.
|
||||
///
|
||||
/// \param task The task with missing dependencies.
|
||||
void SubscribeTaskReady(const Task &task);
|
||||
|
||||
/// Stop waiting for a task's dependencies to become available.
|
||||
///
|
||||
/// \param task_id The task ID of the task with missing dependencies.
|
||||
void UnsubscribeTaskReady(const TaskID &task_id);
|
||||
|
||||
/// Mark an object as locally available. This is used for objects that do not
|
||||
/// have a stored value (e.g., actor execution dependencies).
|
||||
///
|
||||
/// \param object_id The object ID of the object to mark as locally
|
||||
/// available.
|
||||
void MarkDependencyReady(const ObjectID &object_id);
|
||||
|
||||
private:
|
||||
/// Check whether the given list of objects are ready.
|
||||
bool argumentsReady(const std::vector<ObjectID> arguments) const;
|
||||
/// Handle an object added to the object store.
|
||||
void handleObjectReady(const ray::ObjectID &object_id);
|
||||
/// A reference to the object manager so that we can issue Pull requests of
|
||||
/// missing objects.
|
||||
ObjectManager &object_manager_;
|
||||
/// A mapping from task ID of each subscribed task to its list of
|
||||
/// dependencies.
|
||||
std::unordered_map<ray::TaskID, std::vector<ray::ObjectID>, UniqueIDHasher>
|
||||
task_dependencies_;
|
||||
// A mapping from object ID of each object that is not locally available to
|
||||
// the list of subscribed tasks that are dependent on it.
|
||||
std::unordered_map<ray::ObjectID, std::vector<ray::TaskID>, UniqueIDHasher>
|
||||
remote_object_dependencies_;
|
||||
// The set of locally available objects.
|
||||
std::unordered_set<ray::ObjectID, UniqueIDHasher> local_objects_;
|
||||
// The callback to call when a subscribed task becomes ready.
|
||||
std::function<void(const TaskID &)> task_ready_callback_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_TASK_DEPENDENCY_MANAGER_H
|
||||
@@ -0,0 +1,36 @@
|
||||
#include "ray/raylet/task_execution_spec.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
TaskExecutionSpecification::TaskExecutionSpecification(
|
||||
const std::vector<ObjectID> &&execution_dependencies)
|
||||
: execution_dependencies_(std::move(execution_dependencies)),
|
||||
last_timestamp_(0),
|
||||
spillback_count_(0) {}
|
||||
|
||||
TaskExecutionSpecification::TaskExecutionSpecification(
|
||||
const std::vector<ObjectID> &&execution_dependencies, int spillback_count)
|
||||
: execution_dependencies_(std::move(execution_dependencies)),
|
||||
last_timestamp_(0),
|
||||
spillback_count_(spillback_count) {}
|
||||
|
||||
const std::vector<ObjectID> &TaskExecutionSpecification::ExecutionDependencies() const {
|
||||
return execution_dependencies_;
|
||||
}
|
||||
|
||||
void TaskExecutionSpecification::SetExecutionDependencies(
|
||||
const std::vector<ObjectID> &dependencies) {
|
||||
execution_dependencies_ = dependencies;
|
||||
}
|
||||
|
||||
int TaskExecutionSpecification::SpillbackCount() const { return spillback_count_; }
|
||||
|
||||
void TaskExecutionSpecification::IncrementSpillbackCount() { ++spillback_count_; }
|
||||
|
||||
int64_t TaskExecutionSpecification::LastTimeStamp() const { return last_timestamp_; }
|
||||
|
||||
void TaskExecutionSpecification::SetLastTimeStamp(int64_t new_timestamp) {
|
||||
last_timestamp_ = new_timestamp;
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,75 @@
|
||||
#ifndef RAY_RAYLET_TASK_EXECUTION_SPECIFICATION_H
|
||||
#define RAY_RAYLET_TASK_EXECUTION_SPECIFICATION_H
|
||||
|
||||
#include <vector>
|
||||
|
||||
#include "ray/id.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// \class TaskExecutionSpecification
|
||||
///
|
||||
/// The task execution specification encapsulates all mutable information about
|
||||
/// the task. These fields may change at execution time, converse to the
|
||||
/// TaskSpecification that is determined at submission time.
|
||||
class TaskExecutionSpecification {
|
||||
public:
|
||||
/// Create a task execution specification.
|
||||
///
|
||||
/// \param execution_dependencies The task's dependencies, determined at
|
||||
/// execution time.
|
||||
TaskExecutionSpecification(const std::vector<ObjectID> &&execution_dependencies);
|
||||
|
||||
/// Create a task execution specification.
|
||||
///
|
||||
/// \param execution_dependencies The task's dependencies, determined at
|
||||
/// execution time.
|
||||
/// \param spillback_count The number of times this task was spilled back by
|
||||
/// local schedulers.
|
||||
TaskExecutionSpecification(const std::vector<ObjectID> &&execution_dependencies,
|
||||
int spillback_count);
|
||||
|
||||
/// Get the task's execution dependencies.
|
||||
///
|
||||
/// \return A vector of object IDs representing this task's execution
|
||||
/// dependencies.
|
||||
const std::vector<ObjectID> &ExecutionDependencies() const;
|
||||
|
||||
/// Set the task's execution dependencies.
|
||||
///
|
||||
/// \param dependencies The value to set the execution dependencies to.
|
||||
void SetExecutionDependencies(const std::vector<ObjectID> &dependencies);
|
||||
|
||||
/// Get the task's spillback count, which tracks the number of times
|
||||
/// this task was spilled back from local to the global scheduler.
|
||||
///
|
||||
/// \return The spillback count for this task.
|
||||
int SpillbackCount() const;
|
||||
|
||||
/// Increment the spillback count for this task.
|
||||
void IncrementSpillbackCount();
|
||||
|
||||
/// Get the task's last timestamp.
|
||||
///
|
||||
/// \return The timestamp when this task was last received for scheduling.
|
||||
int64_t LastTimeStamp() const;
|
||||
|
||||
/// Set the task's last timestamp to the specified value.
|
||||
///
|
||||
/// \param new_timestamp The new timestamp in millisecond to set the task's
|
||||
/// time stamp to. Tracks the last time this task entered a local
|
||||
/// scheduler.
|
||||
void SetLastTimeStamp(int64_t new_timestamp);
|
||||
|
||||
private:
|
||||
/// A list of object IDs representing the dependencies of this task that may
|
||||
/// change at execution time.
|
||||
std::vector<ObjectID> execution_dependencies_;
|
||||
/// The last time this task was received for scheduling.
|
||||
int64_t last_timestamp_;
|
||||
/// The number of times this task was spilled back by local schedulers.
|
||||
int spillback_count_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
#endif // RAY_RAYLET_TASK_EXECUTION_SPECIFICATION_H
|
||||
@@ -0,0 +1,182 @@
|
||||
#include "task_spec.h"
|
||||
|
||||
#include "common.h"
|
||||
#include "common_protocol.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
TaskArgument::~TaskArgument() {}
|
||||
|
||||
TaskArgumentByReference::TaskArgumentByReference(const std::vector<ObjectID> &references)
|
||||
: references_(references) {}
|
||||
|
||||
flatbuffers::Offset<Arg> TaskArgumentByReference::ToFlatbuffer(
|
||||
flatbuffers::FlatBufferBuilder &fbb) const {
|
||||
return CreateArg(fbb, to_flatbuf(fbb, references_));
|
||||
}
|
||||
|
||||
const BYTE *TaskArgumentByReference::HashData() const {
|
||||
return reinterpret_cast<const BYTE *>(references_.data());
|
||||
}
|
||||
|
||||
size_t TaskArgumentByReference::HashDataLength() const {
|
||||
return references_.size() * sizeof(ObjectID);
|
||||
}
|
||||
|
||||
TaskArgumentByValue::TaskArgumentByValue(const uint8_t *value, size_t length) {
|
||||
value_.assign(value, value + length);
|
||||
}
|
||||
|
||||
flatbuffers::Offset<Arg> TaskArgumentByValue::ToFlatbuffer(
|
||||
flatbuffers::FlatBufferBuilder &fbb) const {
|
||||
auto arg =
|
||||
fbb.CreateString(reinterpret_cast<const char *>(value_.data()), value_.size());
|
||||
auto empty_ids = fbb.CreateVectorOfStrings({});
|
||||
return CreateArg(fbb, empty_ids, arg);
|
||||
}
|
||||
|
||||
const BYTE *TaskArgumentByValue::HashData() const { return value_.data(); }
|
||||
|
||||
size_t TaskArgumentByValue::HashDataLength() const { return value_.size(); }
|
||||
|
||||
static const ObjectID task_compute_return_id(TaskID task_id, int64_t return_index) {
|
||||
// Here, return_indices need to be >= 0, so we can use negative indices for put.
|
||||
RAY_DCHECK(return_index >= 0);
|
||||
// TODO(rkn): This line requires object and task IDs to be the same size.
|
||||
ObjectID return_id = task_id;
|
||||
int64_t *first_bytes = (int64_t *)&return_id;
|
||||
// XOR the first bytes of the object ID with the return index.
|
||||
// We add one so the first return ID is not the same as the task ID.
|
||||
*first_bytes = *first_bytes ^ (return_index + 1);
|
||||
return return_id;
|
||||
}
|
||||
|
||||
TaskSpecification::TaskSpecification(const uint8_t *spec, size_t spec_size)
|
||||
: spec_(spec, spec + spec_size) {}
|
||||
|
||||
TaskSpecification::TaskSpecification(const flatbuffers::String &string)
|
||||
: TaskSpecification(reinterpret_cast<const uint8_t *>(string.data()), string.size()) {
|
||||
}
|
||||
|
||||
TaskSpecification::TaskSpecification(
|
||||
UniqueID driver_id, TaskID parent_task_id, int64_t parent_counter,
|
||||
// UniqueID actor_id,
|
||||
// UniqueID actor_handle_id,
|
||||
// int64_t actor_counter,
|
||||
FunctionID function_id, const std::vector<TaskArgument> &task_arguments,
|
||||
int64_t num_returns,
|
||||
const std::unordered_map<std::string, double> &required_resources) {
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
|
||||
// Compute hashes.
|
||||
SHA256_CTX ctx;
|
||||
sha256_init(&ctx);
|
||||
sha256_update(&ctx, (BYTE *)&driver_id, sizeof(driver_id));
|
||||
sha256_update(&ctx, (BYTE *)&parent_task_id, sizeof(parent_task_id));
|
||||
sha256_update(&ctx, (BYTE *)&parent_counter, sizeof(parent_counter));
|
||||
// sha256_update(&ctx, (BYTE *) &actor_id, sizeof(actor_id));
|
||||
// sha256_update(&ctx, (BYTE *) &actor_counter, sizeof(actor_counter));
|
||||
// sha256_update(&ctx, (BYTE *) &is_actor_checkpoint_method,
|
||||
// sizeof(is_actor_checkpoint_method));
|
||||
sha256_update(&ctx, (BYTE *)&function_id, sizeof(function_id));
|
||||
|
||||
// Serialize and hash the arguments.
|
||||
std::vector<flatbuffers::Offset<Arg>> arguments;
|
||||
for (auto &argument : task_arguments) {
|
||||
arguments.push_back(argument.ToFlatbuffer(fbb));
|
||||
sha256_update(&ctx, (BYTE *)argument.HashData(), argument.HashDataLength());
|
||||
}
|
||||
|
||||
// Compute the final task ID from the hash.
|
||||
BYTE buff[DIGEST_SIZE];
|
||||
sha256_final(&ctx, buff);
|
||||
TaskID task_id;
|
||||
RAY_DCHECK(sizeof(task_id) <= DIGEST_SIZE);
|
||||
memcpy(&task_id, buff, sizeof(task_id));
|
||||
|
||||
// Add return object IDs.
|
||||
std::vector<flatbuffers::Offset<flatbuffers::String>> returns;
|
||||
for (int64_t i = 0; i < num_returns; i++) {
|
||||
ObjectID return_id = task_compute_return_id(task_id, i);
|
||||
returns.push_back(to_flatbuf(fbb, return_id));
|
||||
}
|
||||
|
||||
// Serialize the TaskSpecification.
|
||||
auto spec = CreateTaskInfo(
|
||||
fbb, to_flatbuf(fbb, driver_id), to_flatbuf(fbb, task_id),
|
||||
to_flatbuf(fbb, parent_task_id), parent_counter, to_flatbuf(fbb, WorkerID::nil()),
|
||||
to_flatbuf(fbb, ActorHandleID::nil()), 0, false, to_flatbuf(fbb, function_id),
|
||||
fbb.CreateVector(arguments), fbb.CreateVector(returns),
|
||||
map_to_flatbuf(fbb, required_resources));
|
||||
fbb.Finish(spec);
|
||||
TaskSpecification(fbb.GetBufferPointer(), fbb.GetSize());
|
||||
}
|
||||
|
||||
flatbuffers::Offset<flatbuffers::String> TaskSpecification::ToFlatbuffer(
|
||||
flatbuffers::FlatBufferBuilder &fbb) const {
|
||||
return fbb.CreateString(reinterpret_cast<const char *>(data()), size());
|
||||
}
|
||||
|
||||
// TODO(atumanov): copy/paste most TaskSpec_* methods from task.h and make them
|
||||
// methods of this class.
|
||||
const uint8_t *TaskSpecification::data() const { return spec_.data(); }
|
||||
|
||||
size_t TaskSpecification::size() const { return spec_.size(); }
|
||||
|
||||
// Task specification getter methods.
|
||||
TaskID TaskSpecification::TaskId() const {
|
||||
auto message = flatbuffers::GetRoot<TaskInfo>(spec_.data());
|
||||
return from_flatbuf(*message->task_id());
|
||||
}
|
||||
UniqueID TaskSpecification::DriverId() const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
TaskID TaskSpecification::ParentTaskId() const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
int64_t TaskSpecification::ParentCounter() const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
FunctionID TaskSpecification::FunctionId() const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
|
||||
int64_t TaskSpecification::NumArgs() const {
|
||||
auto message = flatbuffers::GetRoot<TaskInfo>(spec_.data());
|
||||
return message->args()->size();
|
||||
}
|
||||
|
||||
int64_t TaskSpecification::NumReturns() const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
|
||||
bool TaskSpecification::ArgByRef(int64_t arg_index) const {
|
||||
return (ArgIdCount(arg_index) != 0);
|
||||
}
|
||||
|
||||
int TaskSpecification::ArgIdCount(int64_t arg_index) const {
|
||||
auto message = flatbuffers::GetRoot<TaskInfo>(spec_.data());
|
||||
auto ids = message->args()->Get(arg_index)->object_ids();
|
||||
return ids->size();
|
||||
}
|
||||
|
||||
ObjectID TaskSpecification::ArgId(int64_t arg_index, int64_t id_index) const {
|
||||
auto message = flatbuffers::GetRoot<TaskInfo>(spec_.data());
|
||||
return from_flatbuf(*message->args()->Get(arg_index)->object_ids()->Get(id_index));
|
||||
}
|
||||
const uint8_t *TaskSpecification::ArgVal(int64_t arg_index) const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
size_t TaskSpecification::ArgValLength(int64_t arg_index) const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
double TaskSpecification::GetRequiredResource(const std::string &resource_name) const {
|
||||
throw std::runtime_error("Method not implemented");
|
||||
}
|
||||
const ResourceSet TaskSpecification::GetRequiredResources() const {
|
||||
auto message = flatbuffers::GetRoot<TaskInfo>(spec_.data());
|
||||
auto required_resources = map_from_flatbuf(*message->required_resources());
|
||||
return ResourceSet(required_resources);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,152 @@
|
||||
#ifndef RAY_RAYLET_TASK_SPECIFICATION_H
|
||||
#define RAY_RAYLET_TASK_SPECIFICATION_H
|
||||
|
||||
#include <cstddef>
|
||||
#include <string>
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
|
||||
#include "ray/../common/format/common_generated.h"
|
||||
#include "ray/id.h"
|
||||
#include "ray/raylet/scheduling_resources.h"
|
||||
|
||||
extern "C" {
|
||||
#include "sha256.h"
|
||||
}
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// \class TaskArgument
|
||||
///
|
||||
/// A virtual class that represents an argument to a task.
|
||||
class TaskArgument {
|
||||
public:
|
||||
/// Serialize the task argument to a flatbuffer.
|
||||
///
|
||||
/// \param fbb The flatbuffer builder to serialize with.
|
||||
/// \return An offset to the serialized task argument.
|
||||
virtual flatbuffers::Offset<Arg> ToFlatbuffer(
|
||||
flatbuffers::FlatBufferBuilder &fbb) const = 0;
|
||||
|
||||
/// Get the hashable byte data.
|
||||
///
|
||||
/// \return A pointer to the byte data.
|
||||
virtual const BYTE *HashData() const = 0;
|
||||
|
||||
/// Get the hashable byte data length.
|
||||
///
|
||||
/// \return The length of the hashable byte data.
|
||||
virtual size_t HashDataLength() const = 0;
|
||||
|
||||
virtual ~TaskArgument() = 0;
|
||||
};
|
||||
|
||||
/// \class TaskArgumentByReference
|
||||
///
|
||||
/// A task argument consisting of a list of object ID references.
|
||||
class TaskArgumentByReference : virtual public TaskArgument {
|
||||
/// Create a task argument by reference from a list of object IDs.
|
||||
///
|
||||
/// \param references A list of object ID references.
|
||||
TaskArgumentByReference(const std::vector<ObjectID> &references);
|
||||
|
||||
flatbuffers::Offset<Arg> ToFlatbuffer(flatbuffers::FlatBufferBuilder &fbb) const;
|
||||
const BYTE *HashData() const;
|
||||
size_t HashDataLength() const;
|
||||
|
||||
private:
|
||||
/// The object IDs.
|
||||
const std::vector<ObjectID> references_;
|
||||
};
|
||||
|
||||
/// \class TaskArgumentByValue
|
||||
///
|
||||
/// A task argument containing the raw value.
|
||||
class TaskArgumentByValue : public TaskArgument {
|
||||
/// Create a task argument from a raw value.
|
||||
///
|
||||
/// \param value A pointer to the raw value.
|
||||
/// \param length The size of the raw value.
|
||||
TaskArgumentByValue(const uint8_t *value, size_t length);
|
||||
|
||||
flatbuffers::Offset<Arg> ToFlatbuffer(flatbuffers::FlatBufferBuilder &fbb) const;
|
||||
const BYTE *HashData() const;
|
||||
size_t HashDataLength() const;
|
||||
|
||||
private:
|
||||
/// The raw value.
|
||||
std::vector<uint8_t> value_;
|
||||
};
|
||||
|
||||
/// \class TaskSpecification
|
||||
///
|
||||
/// The task specification encapsulates all immutable information about the
|
||||
/// task. These fields are determined at submission time, converse to the
|
||||
/// TaskExecutionSpecification that may change at execution time.
|
||||
class TaskSpecification {
|
||||
public:
|
||||
/// Deserialize a task specification from a flatbuffer.
|
||||
///
|
||||
/// \param string A serialized task specification flatbuffer.
|
||||
TaskSpecification(const flatbuffers::String &string);
|
||||
|
||||
// TODO(swang): Define an actor task constructor.
|
||||
/// Create a task specification from the raw fields.
|
||||
///
|
||||
/// \param driver_id The driver ID, representing the job that this task is a
|
||||
/// part of.
|
||||
/// \param parent_task_id The task ID of the task that spawned this task.
|
||||
/// \param parent_counter The number of tasks that this task's parent spawned
|
||||
/// before this task.
|
||||
/// \param function_id The ID of the function this task should execute.
|
||||
/// \param arguments The list of task arguments.
|
||||
/// \param num_returns The number of values returned by the task.
|
||||
/// \param required_resources The task's resource demands.
|
||||
TaskSpecification(UniqueID driver_id, TaskID parent_task_id, int64_t parent_counter,
|
||||
// UniqueID actor_id,
|
||||
// UniqueID actor_handle_id,
|
||||
// int64_t actor_counter,
|
||||
FunctionID function_id, const std::vector<TaskArgument> &arguments,
|
||||
int64_t num_returns,
|
||||
const std::unordered_map<std::string, double> &required_resources);
|
||||
|
||||
~TaskSpecification() {}
|
||||
|
||||
/// Serialize the TaskSpecification to a flatbuffer.
|
||||
///
|
||||
/// \param fbb The flatbuffer builder to serialize with.
|
||||
/// \return An offset to the serialized task specification.
|
||||
flatbuffers::Offset<flatbuffers::String> ToFlatbuffer(
|
||||
flatbuffers::FlatBufferBuilder &fbb) const;
|
||||
|
||||
// TODO(swang): Finalize and document these methods.
|
||||
TaskID TaskId() const;
|
||||
UniqueID DriverId() const;
|
||||
TaskID ParentTaskId() const;
|
||||
int64_t ParentCounter() const;
|
||||
FunctionID FunctionId() const;
|
||||
int64_t NumArgs() const;
|
||||
int64_t NumReturns() const;
|
||||
bool ArgByRef(int64_t arg_index) const;
|
||||
int ArgIdCount(int64_t arg_index) const;
|
||||
ObjectID ArgId(int64_t arg_index, int64_t id_index) const;
|
||||
const uint8_t *ArgVal(int64_t arg_index) const;
|
||||
size_t ArgValLength(int64_t arg_index) const;
|
||||
double GetRequiredResource(const std::string &resource_name) const;
|
||||
const ResourceSet GetRequiredResources() const;
|
||||
|
||||
private:
|
||||
/// Task specification constructor from a pointer.
|
||||
TaskSpecification(const uint8_t *spec, size_t spec_size);
|
||||
/// Get a pointer to the byte data.
|
||||
const uint8_t *data() const;
|
||||
/// Get the size in bytes of the task specification.
|
||||
size_t size() const;
|
||||
|
||||
/// The task specification data.
|
||||
std::vector<uint8_t> spec_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_TASK_SPECIFICATION_H
|
||||
@@ -0,0 +1,25 @@
|
||||
#include "worker.h"
|
||||
|
||||
#include <boost/bind.hpp>
|
||||
|
||||
#include "common.h"
|
||||
#include "ray/raylet/format/node_manager_generated.h"
|
||||
#include "ray/raylet/raylet.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// A constructor responsible for initializing the state of a worker.
|
||||
Worker::Worker(pid_t pid, std::shared_ptr<LocalClientConnection> connection)
|
||||
: pid_(pid), connection_(connection), assigned_task_id_(TaskID::nil()) {}
|
||||
|
||||
pid_t Worker::Pid() const { return pid_; }
|
||||
|
||||
void Worker::AssignTaskId(const TaskID &task_id) { assigned_task_id_ = task_id; }
|
||||
|
||||
const TaskID &Worker::GetAssignedTaskId() const { return assigned_task_id_; }
|
||||
|
||||
const std::shared_ptr<LocalClientConnection> Worker::Connection() const {
|
||||
return connection_;
|
||||
}
|
||||
|
||||
} // end namespace ray
|
||||
@@ -0,0 +1,37 @@
|
||||
#ifndef RAY_RAYLET_WORKER_H
|
||||
#define RAY_RAYLET_WORKER_H
|
||||
|
||||
#include <memory>
|
||||
|
||||
#include "ray/common/client_connection.h"
|
||||
#include "ray/id.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// Worker class encapsulates the implementation details of a worker. A worker
|
||||
/// is the execution container around a unit of Ray work, such as a task or an
|
||||
/// actor. Ray units of work execute in the context of a Worker.
|
||||
class Worker {
|
||||
public:
|
||||
/// A constructor that initializes a worker object.
|
||||
Worker(pid_t pid, std::shared_ptr<LocalClientConnection> connection);
|
||||
/// A destructor responsible for freeing all worker state.
|
||||
~Worker() {}
|
||||
/// Return the worker's PID.
|
||||
pid_t Pid() const;
|
||||
void AssignTaskId(const TaskID &task_id);
|
||||
const TaskID &GetAssignedTaskId() const;
|
||||
/// Return the worker's connection.
|
||||
const std::shared_ptr<LocalClientConnection> Connection() const;
|
||||
|
||||
private:
|
||||
/// The worker's PID.
|
||||
pid_t pid_;
|
||||
/// Connection state of a worker.
|
||||
std::shared_ptr<LocalClientConnection> connection_;
|
||||
TaskID assigned_task_id_;
|
||||
};
|
||||
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_WORKER_H
|
||||
@@ -0,0 +1,78 @@
|
||||
#include "ray/raylet/worker_pool.h"
|
||||
|
||||
#include "ray/status.h"
|
||||
#include "ray/util/logging.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
/// A constructor that initializes a worker pool with num_workers workers.
|
||||
WorkerPool::WorkerPool(int num_workers) {
|
||||
for (int i = 0; i < num_workers; i++) {
|
||||
StartWorker();
|
||||
}
|
||||
}
|
||||
|
||||
WorkerPool::~WorkerPool() {
|
||||
// TODO(swang): Kill registered workers.
|
||||
pool_.clear();
|
||||
registered_workers_.clear();
|
||||
}
|
||||
|
||||
/// Create a new worker and add it to the pool
|
||||
bool WorkerPool::StartWorker() {
|
||||
// TODO(swang): Start the worker.
|
||||
return true;
|
||||
}
|
||||
|
||||
uint32_t WorkerPool::PoolSize() const { return pool_.size(); }
|
||||
|
||||
void WorkerPool::RegisterWorker(std::shared_ptr<Worker> worker) {
|
||||
RAY_LOG(DEBUG) << "Registering worker with pid " << worker->Pid();
|
||||
registered_workers_.push_back(worker);
|
||||
}
|
||||
|
||||
const std::shared_ptr<Worker> WorkerPool::GetRegisteredWorker(
|
||||
std::shared_ptr<LocalClientConnection> connection) const {
|
||||
for (auto it = registered_workers_.begin(); it != registered_workers_.end(); it++) {
|
||||
if ((*it)->Connection() == connection) {
|
||||
return (*it);
|
||||
}
|
||||
}
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
void WorkerPool::PushWorker(std::shared_ptr<Worker> worker) {
|
||||
// Since the worker is now idle, unset its assigned task ID.
|
||||
worker->AssignTaskId(TaskID::nil());
|
||||
// Add the worker to the idle pool.
|
||||
pool_.push_back(std::move(worker));
|
||||
}
|
||||
|
||||
std::shared_ptr<Worker> WorkerPool::PopWorker() {
|
||||
if (pool_.empty()) {
|
||||
return nullptr;
|
||||
}
|
||||
std::shared_ptr<Worker> worker = std::move(pool_.back());
|
||||
pool_.pop_back();
|
||||
return worker;
|
||||
}
|
||||
|
||||
// A helper function to remove a worker from a list. Returns true if the worker
|
||||
// was found and removed.
|
||||
bool removeWorker(std::list<std::shared_ptr<Worker>> &worker_pool,
|
||||
std::shared_ptr<Worker> worker) {
|
||||
for (auto it = worker_pool.begin(); it != worker_pool.end(); it++) {
|
||||
if (*it == worker) {
|
||||
worker_pool.erase(it);
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
bool WorkerPool::DisconnectWorker(std::shared_ptr<Worker> worker) {
|
||||
RAY_CHECK(removeWorker(registered_workers_, worker));
|
||||
return removeWorker(pool_, worker);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
@@ -0,0 +1,85 @@
|
||||
#ifndef RAY_RAYLET_WORKER_POOL_H
|
||||
#define RAY_RAYLET_WORKER_POOL_H
|
||||
|
||||
#include <inttypes.h>
|
||||
#include <list>
|
||||
|
||||
#include "ray/common/client_connection.h"
|
||||
#include "ray/raylet/worker.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
class Worker;
|
||||
|
||||
/// \class WorkerPool
|
||||
///
|
||||
/// The WorkerPool is responsible for managing a pool of Workers. Each Worker
|
||||
/// is a container for a unit of work.
|
||||
class WorkerPool {
|
||||
public:
|
||||
/// Create a pool and asynchronously start the specified number of workers.
|
||||
/// Once each worker process has registered with an external server, the
|
||||
/// process should create and register a new Worker, then add itself to the
|
||||
/// pool.
|
||||
///
|
||||
/// \param num_workers The number of workers to start.
|
||||
WorkerPool(int num_workers);
|
||||
|
||||
/// Destructor responsible for freeing a set of workers owned by this class.
|
||||
~WorkerPool();
|
||||
|
||||
/// Get the number of idle workers in the pool.
|
||||
///
|
||||
/// \return The number of idle workers.
|
||||
uint32_t PoolSize() const;
|
||||
|
||||
/// Asynchronously start a new worker process. Once the worker process has
|
||||
/// registered with an external server, the process should create and
|
||||
/// register a new Worker, then add itself to the pool.
|
||||
///
|
||||
/// \return Whether the worker process was successfully started.
|
||||
bool StartWorker();
|
||||
|
||||
/// Register a new worker. The Worker should be added by the caller to the
|
||||
/// pool after it becomes idle (e.g., requests a work assignment).
|
||||
///
|
||||
/// \param The Worker to be registered.
|
||||
void RegisterWorker(std::shared_ptr<Worker> worker);
|
||||
|
||||
/// Get the client connection's registered worker.
|
||||
///
|
||||
/// \param The client connection owned by a registered worker.
|
||||
/// \return The Worker that owns the given client connection. Returns nullptr
|
||||
/// if the client has not registered a worker yet.
|
||||
const std::shared_ptr<Worker> GetRegisteredWorker(
|
||||
std::shared_ptr<LocalClientConnection> connection) const;
|
||||
|
||||
/// Disconnect a registered worker.
|
||||
///
|
||||
/// \param The worker to disconnect. The worker must be registered.
|
||||
/// \return Whether the given worker was in the pool of idle workers.
|
||||
bool DisconnectWorker(std::shared_ptr<Worker> worker);
|
||||
|
||||
/// Add an idle worker to the pool. The worker's task assignment will be
|
||||
/// reset.
|
||||
///
|
||||
/// \param The idle worker to add.
|
||||
void PushWorker(std::shared_ptr<Worker> worker);
|
||||
|
||||
/// Pop an idle worker from the pool. The caller is responsible for pushing
|
||||
/// the worker back onto the pool once the worker has completed its work.
|
||||
///
|
||||
/// \return An idle worker. Returns nullptr if the pool is empty.
|
||||
std::shared_ptr<Worker> PopWorker();
|
||||
|
||||
private:
|
||||
/// The pool of idle workers.
|
||||
std::list<std::shared_ptr<Worker>> pool_;
|
||||
/// All workers that have registered and are still connected, including both
|
||||
/// idle and executing.
|
||||
// TODO(swang): Make this a map to make GetRegisteredWorker faster.
|
||||
std::list<std::shared_ptr<Worker>> registered_workers_;
|
||||
};
|
||||
} // namespace ray
|
||||
|
||||
#endif // RAY_RAYLET_WORKER_POOL_H
|
||||
@@ -0,0 +1,74 @@
|
||||
#include "gmock/gmock.h"
|
||||
#include "gtest/gtest.h"
|
||||
|
||||
#include "ray/raylet/node_manager.h"
|
||||
#include "ray/raylet/worker_pool.h"
|
||||
|
||||
namespace ray {
|
||||
|
||||
class MockClientManager : public ClientManager<boost::asio::local::stream_protocol> {
|
||||
public:
|
||||
MOCK_METHOD3(ProcessClientMessage,
|
||||
void(std::shared_ptr<LocalClientConnection>, int64_t, const uint8_t *));
|
||||
MOCK_METHOD1(ProcessNewClient, void(std::shared_ptr<LocalClientConnection>));
|
||||
};
|
||||
|
||||
class WorkerPoolTest : public ::testing::Test {
|
||||
public:
|
||||
WorkerPoolTest() : worker_pool_(0), client_manager_(), io_service_() {}
|
||||
|
||||
std::shared_ptr<Worker> CreateWorker(pid_t pid) {
|
||||
boost::asio::local::stream_protocol::socket socket(io_service_);
|
||||
auto client = LocalClientConnection::Create(client_manager_, std::move(socket));
|
||||
return std::shared_ptr<Worker>(new Worker(pid, client));
|
||||
}
|
||||
|
||||
protected:
|
||||
WorkerPool worker_pool_;
|
||||
MockClientManager client_manager_;
|
||||
boost::asio::io_service io_service_;
|
||||
};
|
||||
|
||||
TEST_F(WorkerPoolTest, HandleWorkerRegistration) {
|
||||
auto worker = CreateWorker(1234);
|
||||
// Check that we cannot lookup the worker before it's registered.
|
||||
ASSERT_EQ(worker_pool_.GetRegisteredWorker(worker->Connection()), nullptr);
|
||||
worker_pool_.RegisterWorker(worker);
|
||||
// Check that we can lookup the worker after it's registered.
|
||||
ASSERT_EQ(worker_pool_.GetRegisteredWorker(worker->Connection()), worker);
|
||||
worker_pool_.DisconnectWorker(worker);
|
||||
// Check that we cannot lookup the worker after it's disconnected.
|
||||
ASSERT_EQ(worker_pool_.GetRegisteredWorker(worker->Connection()), nullptr);
|
||||
}
|
||||
|
||||
TEST_F(WorkerPoolTest, HandleWorkerPushPop) {
|
||||
// Try to pop a worker from the empty pool and make sure we don't get one.
|
||||
std::shared_ptr<Worker> popped_worker;
|
||||
popped_worker = worker_pool_.PopWorker();
|
||||
ASSERT_EQ(popped_worker, nullptr);
|
||||
|
||||
// Create some workers.
|
||||
std::unordered_set<std::shared_ptr<Worker>> workers;
|
||||
workers.insert(CreateWorker(1234));
|
||||
workers.insert(CreateWorker(5678));
|
||||
// Add the workers to the pool.
|
||||
for (auto &worker : workers) {
|
||||
worker_pool_.PushWorker(worker);
|
||||
}
|
||||
ASSERT_EQ(worker_pool_.PoolSize(), workers.size());
|
||||
|
||||
// Pop two workers and make sure they're one of the workers we created.
|
||||
popped_worker = worker_pool_.PopWorker();
|
||||
ASSERT_NE(popped_worker, nullptr);
|
||||
ASSERT_TRUE(workers.count(popped_worker) > 0);
|
||||
popped_worker = worker_pool_.PopWorker();
|
||||
ASSERT_NE(popped_worker, nullptr);
|
||||
ASSERT_TRUE(workers.count(popped_worker) > 0);
|
||||
}
|
||||
|
||||
} // namespace ray
|
||||
|
||||
int main(int argc, char **argv) {
|
||||
::testing::InitGoogleTest(&argc, argv);
|
||||
return RUN_ALL_TESTS();
|
||||
}
|
||||
@@ -19,6 +19,7 @@
|
||||
#include <iosfwd>
|
||||
#include <string>
|
||||
|
||||
#include "ray/util/logging.h"
|
||||
#include "ray/util/macros.h"
|
||||
#include "ray/util/visibility.h"
|
||||
|
||||
|
||||
@@ -38,8 +38,7 @@ namespace ray {
|
||||
|
||||
#define RAY_DCHECK(condition) \
|
||||
RAY_IGNORE_EXPR(condition) \
|
||||
while (false) \
|
||||
::ray::internal::NullLog()
|
||||
while (false) ::ray::internal::NullLog()
|
||||
|
||||
#else
|
||||
|
||||
|
||||
Reference in New Issue
Block a user