diff --git a/src/local_scheduler/lib/java/org_ray_spi_impl_DefaultLocalSchedulerClient.cc b/src/local_scheduler/lib/java/org_ray_spi_impl_DefaultLocalSchedulerClient.cc index 7eef8596d..70de5786c 100644 --- a/src/local_scheduler/lib/java/org_ray_spi_impl_DefaultLocalSchedulerClient.cc +++ b/src/local_scheduler/lib/java/org_ray_spi_impl_DefaultLocalSchedulerClient.cc @@ -50,8 +50,9 @@ Java_org_ray_spi_impl_DefaultLocalSchedulerClient__1init(JNIEnv *env, UniqueIdFromJByteArray worker_id(env, wid); UniqueIdFromJByteArray driver_id(env, driverId); const char *nativeString = env->GetStringUTFChars(sockName, JNI_FALSE); - auto client = LocalSchedulerConnection_init( - nativeString, *worker_id.PID, isWorker, *driver_id.PID, useRaylet); + auto client = + LocalSchedulerConnection_init(nativeString, *worker_id.PID, isWorker, + *driver_id.PID, useRaylet, Language::JAVA); env->ReleaseStringUTFChars(sockName, nativeString); return reinterpret_cast(client); } diff --git a/src/local_scheduler/lib/python/local_scheduler_extension.cc b/src/local_scheduler/lib/python/local_scheduler_extension.cc index 07853aaf8..d473a4eb9 100644 --- a/src/local_scheduler/lib/python/local_scheduler_extension.cc +++ b/src/local_scheduler/lib/python/local_scheduler_extension.cc @@ -31,7 +31,8 @@ static int PyLocalSchedulerClient_init(PyLocalSchedulerClient *self, /* Connect to the local scheduler. */ self->local_scheduler_connection = LocalSchedulerConnection_init( socket_name, client_id, static_cast(PyObject_IsTrue(is_worker)), - driver_id, static_cast(PyObject_IsTrue(use_raylet))); + driver_id, static_cast(PyObject_IsTrue(use_raylet)), + Language::PYTHON); return 0; } diff --git a/src/local_scheduler/local_scheduler_client.cc b/src/local_scheduler/local_scheduler_client.cc index 994c0f176..91b5fa9c9 100644 --- a/src/local_scheduler/local_scheduler_client.cc +++ b/src/local_scheduler/local_scheduler_client.cc @@ -17,7 +17,8 @@ LocalSchedulerConnection *LocalSchedulerConnection_init( const UniqueID &client_id, bool is_worker, const JobID &driver_id, - bool use_raylet) { + bool use_raylet, + const Language &language) { LocalSchedulerConnection *result = new LocalSchedulerConnection(); result->use_raylet = use_raylet; result->conn = connect_ipc_sock_retry(local_scheduler_socket, -1, -1); @@ -26,10 +27,17 @@ LocalSchedulerConnection *LocalSchedulerConnection_init( * NOTE(swang): If the local scheduler exits and we are registered as a * worker, we will get killed. */ flatbuffers::FlatBufferBuilder fbb; - auto message = ray::local_scheduler::protocol::CreateRegisterClientRequest( - fbb, is_worker, to_flatbuf(fbb, client_id), getpid(), - to_flatbuf(fbb, driver_id)); - fbb.Finish(message); + if (use_raylet) { + auto message = ray::protocol::CreateRegisterClientRequest( + fbb, is_worker, to_flatbuf(fbb, client_id), getpid(), + to_flatbuf(fbb, driver_id), language); + fbb.Finish(message); + } else { + auto message = ray::local_scheduler::protocol::CreateRegisterClientRequest( + fbb, is_worker, to_flatbuf(fbb, client_id), getpid(), + to_flatbuf(fbb, driver_id)); + fbb.Finish(message); + } /* Register the process ID with the local scheduler. */ int success = write_message( result->conn, static_cast(MessageType::RegisterClientRequest), diff --git a/src/local_scheduler/local_scheduler_client.h b/src/local_scheduler/local_scheduler_client.h index 22eb51484..bb4fdb345 100644 --- a/src/local_scheduler/local_scheduler_client.h +++ b/src/local_scheduler/local_scheduler_client.h @@ -46,7 +46,8 @@ LocalSchedulerConnection *LocalSchedulerConnection_init( const UniqueID &worker_id, bool is_worker, const JobID &driver_id, - bool use_raylet); + bool use_raylet, + const Language &language); /** * Disconnect from the local scheduler. diff --git a/src/local_scheduler/test/local_scheduler_tests.cc b/src/local_scheduler/test/local_scheduler_tests.cc index d7096b274..b155ea949 100644 --- a/src/local_scheduler/test/local_scheduler_tests.cc +++ b/src/local_scheduler/test/local_scheduler_tests.cc @@ -126,7 +126,7 @@ LocalSchedulerMock *LocalSchedulerMock_init(int num_workers, for (int i = 0; i < num_mock_workers; ++i) { mock->conns[i] = LocalSchedulerConnection_init( local_scheduler_socket_name.c_str(), WorkerID::nil(), true, - JobID::nil(), false); + JobID::nil(), false, Language::PYTHON); } background_thread.join(); diff --git a/src/ray/gcs/format/util.h b/src/ray/gcs/format/util.h new file mode 100644 index 000000000..85bc0511d --- /dev/null +++ b/src/ray/gcs/format/util.h @@ -0,0 +1,24 @@ +#ifndef RAY_RAYLET_GCS_FORMAT_UTIL_H +#define RAY_RAYLET_GCS_FORMAT_UTIL_H + +#include "ray/gcs/format/gcs_generated.h" + +namespace std { + +template <> +struct hash { + size_t operator()(const Language &language) const { + return std::hash()(static_cast(language)); + } +}; + +template <> +struct hash { + size_t operator()(const Language &language) const { + return std::hash()(static_cast(language)); + } +}; + +} // namespace std + +#endif // RAY_RAYLET_GCS_FORMAT_UTIL_H diff --git a/src/ray/raylet/format/node_manager.fbs b/src/ray/raylet/format/node_manager.fbs index f339a08f4..4ede8f2b3 100644 --- a/src/ray/raylet/format/node_manager.fbs +++ b/src/ray/raylet/format/node_manager.fbs @@ -123,6 +123,8 @@ table RegisterClientRequest { worker_pid: long; // The driver ID. This is non-nil if the client is a driver. driver_id: string; + // Language of this worker. + language: Language; } table RegisterClientReply { diff --git a/src/ray/raylet/main.cc b/src/ray/raylet/main.cc index dfb7be79b..1932198a3 100644 --- a/src/ray/raylet/main.cc +++ b/src/ray/raylet/main.cc @@ -5,6 +5,15 @@ #include "ray/status.h" #ifndef RAYLET_TEST + +/// A helper function that parse the worker command string into a vector of arguments. +static std::vector parse_worker_command(std::string worker_command) { + std::istringstream iss(worker_command); + std::vector result(std::istream_iterator{iss}, + std::istream_iterator()); + return result; +} + int main(int argc, char *argv[]) { RayLog::StartRayLog(argv[0], RAY_INFO); RAY_CHECK(argc == 10); @@ -39,23 +48,20 @@ int main(int argc, char *argv[]) { node_manager_config.num_initial_workers = num_initial_workers; node_manager_config.num_workers_per_process = RayConfig::instance().num_workers_per_process(); - // Use a default worker that can execute empty tasks with dependencies. - std::string worker_command; if (!python_worker_command.empty()) { - worker_command = python_worker_command; - } else if (!java_worker_command.empty()) { - worker_command = java_worker_command; - } else { + node_manager_config.worker_commands.emplace( + make_pair(Language::PYTHON, parse_worker_command(python_worker_command))); + } + if (!java_worker_command.empty()) { + node_manager_config.worker_commands.emplace( + make_pair(Language::JAVA, parse_worker_command(java_worker_command))); + } + if (python_worker_command.empty() && java_worker_command.empty()) { RAY_CHECK(0) << "Either Python worker command or Java worker command should be provided."; } - std::istringstream iss(worker_command); - std::vector results(std::istream_iterator{iss}, - std::istream_iterator()); - node_manager_config.worker_command.swap(results); - node_manager_config.heartbeat_period_ms = RayConfig::instance().heartbeat_timeout_milliseconds(); node_manager_config.max_lineage_size = RayConfig::instance().max_lineage_size(); diff --git a/src/ray/raylet/node_manager.cc b/src/ray/raylet/node_manager.cc index 1d5d3da2f..99230827a 100644 --- a/src/ray/raylet/node_manager.cc +++ b/src/ray/raylet/node_manager.cc @@ -87,7 +87,7 @@ NodeManager::NodeManager(boost::asio::io_service &io_service, local_available_resources_(config.resource_config), worker_pool_(config.num_initial_workers, config.num_workers_per_process, static_cast(config.resource_config.GetNumCpus()), - config.worker_command), + config.worker_commands), local_queues_(SchedulingQueue()), scheduling_policy_(local_queues_), reconstruction_policy_( @@ -485,7 +485,8 @@ void NodeManager::ProcessClientMessage( case protocol::MessageType::RegisterClientRequest: { auto message = flatbuffers::GetRoot(message_data); client->SetClientID(from_flatbuf(*message->client_id())); - auto worker = std::make_shared(message->worker_pid(), client); + auto worker = + std::make_shared(message->worker_pid(), message->language(), client); if (message->is_worker()) { // Register the new worker. worker_pool_.RegisterWorker(std::move(worker)); @@ -1030,13 +1031,13 @@ void NodeManager::AssignTask(Task &task) { } // Try to get an idle worker that can execute this task. - std::shared_ptr worker = worker_pool_.PopWorker(spec.ActorId()); + std::shared_ptr worker = worker_pool_.PopWorker(spec); if (worker == nullptr) { // There are no workers that can execute this task. if (!spec.IsActorTask()) { // There are no more non-actor workers available to execute this task. // Start a new worker. - worker_pool_.StartWorkerProcess(); + worker_pool_.StartWorkerProcess(spec.GetLanguage()); } // Queue this task for future assignment. The task will be assigned to a // worker once one becomes available. diff --git a/src/ray/raylet/node_manager.h b/src/ray/raylet/node_manager.h index 85d4f14fd..809ab02c2 100644 --- a/src/ray/raylet/node_manager.h +++ b/src/ray/raylet/node_manager.h @@ -7,6 +7,7 @@ #include "ray/raylet/task.h" #include "ray/object_manager/object_manager.h" #include "ray/common/client_connection.h" +#include "ray/gcs/format/util.h" #include "ray/raylet/actor_registration.h" #include "ray/raylet/lineage_cache.h" #include "ray/raylet/scheduling_policy.h" @@ -25,7 +26,8 @@ struct NodeManagerConfig { ResourceSet resource_config; int num_initial_workers; int num_workers_per_process; - std::vector worker_command; + /// The commands used to start the worker process, grouped by language. + std::unordered_map> worker_commands; uint64_t heartbeat_period_ms; uint64_t max_lineage_size; /// The store socket name. diff --git a/src/ray/raylet/object_manager_integration_test.cc b/src/ray/raylet/object_manager_integration_test.cc index 7d2ad6c7c..c22c05c15 100644 --- a/src/ray/raylet/object_manager_integration_test.cc +++ b/src/ray/raylet/object_manager_integration_test.cc @@ -39,11 +39,12 @@ class TestObjectManagerBase : public ::testing::Test { node_manager_config.num_initial_workers = 0; node_manager_config.num_workers_per_process = 1; // Use a default worker that can execute empty tasks with dependencies. - node_manager_config.worker_command.push_back("python"); - node_manager_config.worker_command.push_back( - "../python/ray/workers/default_worker.py"); - node_manager_config.worker_command.push_back(raylet_socket_name.c_str()); - node_manager_config.worker_command.push_back(store_socket_name.c_str()); + std::vector py_worker_command; + py_worker_command.push_back("python"); + py_worker_command.push_back("../python/ray/workers/default_worker.py"); + py_worker_command.push_back(raylet_socket_name.c_str()); + py_worker_command.push_back(store_socket_name.c_str()); + node_manager_config.worker_commands[Language::PYTHON] = py_worker_command; return node_manager_config; }; diff --git a/src/ray/raylet/worker.cc b/src/ray/raylet/worker.cc index 17289d480..b7a70297f 100644 --- a/src/ray/raylet/worker.cc +++ b/src/ray/raylet/worker.cc @@ -11,8 +11,10 @@ namespace ray { namespace raylet { /// A constructor responsible for initializing the state of a worker. -Worker::Worker(pid_t pid, std::shared_ptr connection) +Worker::Worker(pid_t pid, const Language &language, + std::shared_ptr connection) : pid_(pid), + language_(language), connection_(connection), assigned_task_id_(TaskID::nil()), actor_id_(ActorID::nil()), @@ -26,6 +28,8 @@ bool Worker::IsBlocked() const { return blocked_; } pid_t Worker::Pid() const { return pid_; } +Language Worker::GetLanguage() const { return language_; } + void Worker::AssignTaskId(const TaskID &task_id) { assigned_task_id_ = task_id; } const TaskID &Worker::GetAssignedTaskId() const { return assigned_task_id_; } diff --git a/src/ray/raylet/worker.h b/src/ray/raylet/worker.h index cb1be04d8..fe200a2ca 100644 --- a/src/ray/raylet/worker.h +++ b/src/ray/raylet/worker.h @@ -17,7 +17,8 @@ namespace raylet { class Worker { public: /// A constructor that initializes a worker object. - Worker(pid_t pid, std::shared_ptr connection); + Worker(pid_t pid, const Language &language, + std::shared_ptr connection); /// A destructor responsible for freeing all worker state. ~Worker() {} void MarkBlocked(); @@ -25,6 +26,7 @@ class Worker { bool IsBlocked() const; /// Return the worker's PID. pid_t Pid() const; + Language GetLanguage() const; void AssignTaskId(const TaskID &task_id); const TaskID &GetAssignedTaskId() const; void AssignActorId(const ActorID &actor_id); @@ -45,6 +47,8 @@ class Worker { private: /// The worker's PID. pid_t pid_; + /// The language type of this worker. + Language language_; /// Connection state of a worker. std::shared_ptr connection_; /// The worker's currently assigned task. diff --git a/src/ray/raylet/worker_pool.cc b/src/ray/raylet/worker_pool.cc index e1dd34a0e..b5d9aa12c 100644 --- a/src/ray/raylet/worker_pool.cc +++ b/src/ray/raylet/worker_pool.cc @@ -7,8 +7,7 @@ namespace { -// A helper function to remove a worker from a list. Returns true if the worker -// was found and removed. +// A helper function to get a worker from a list. std::shared_ptr GetWorker( const std::list> &worker_pool, const std::shared_ptr &connection) { @@ -40,28 +39,36 @@ namespace ray { namespace raylet { /// A constructor that initializes a worker pool with -/// (num_worker_processes * num_workers_per_process) workers -WorkerPool::WorkerPool(int num_worker_processes, int num_workers_per_process, - int num_cpus, const std::vector &worker_command) - : num_workers_per_process_(num_workers_per_process), - num_cpus_(num_cpus), - worker_command_(worker_command) { +/// (num_worker_processes * num_workers_per_process) workers for each language. +WorkerPool::WorkerPool( + int num_worker_processes, int num_workers_per_process, int num_cpus, + const std::unordered_map> &worker_commands) + : num_workers_per_process_(num_workers_per_process), num_cpus_(num_cpus) { RAY_CHECK(num_workers_per_process > 0) << "num_workers_per_process must be positive."; // Ignore SIGCHLD signals. If we don't do this, then worker processes will // become zombies instead of dying gracefully. signal(SIGCHLD, SIG_IGN); - for (int i = 0; i < num_worker_processes; i++) { - // Force-start num_workers workers. - StartWorkerProcess(true); + for (const auto &entry : worker_commands) { + // Initialize the pool state for this language. + auto &state = states_by_lang_[entry.first]; + // Set worker command for this language. + state.worker_command = entry.second; + RAY_CHECK(!state.worker_command.empty()) << "Worker command must not be empty."; + // Force-start num_workers worker processes for this language. + for (int i = 0; i < num_worker_processes; i++) { + StartWorkerProcess(entry.first, true); + } } } WorkerPool::~WorkerPool() { std::unordered_set pids_to_kill; - // Kill all registered workers. NOTE(swang): This assumes that the registered - // workers were started by the pool. - for (const auto &worker : registered_workers_) { - pids_to_kill.insert(worker->Pid()); + for (const auto &entry : states_by_lang_) { + // Kill all registered workers. NOTE(swang): This assumes that the registered + // workers were started by the pool. + for (const auto &worker : entry.second.registered_workers) { + pids_to_kill.insert(worker->Pid()); + } } // Kill all the workers that have been started but not registered. for (const auto &entry : starting_worker_processes_) { @@ -77,12 +84,17 @@ WorkerPool::~WorkerPool() { } } -uint32_t WorkerPool::Size() const { - return static_cast(actor_pool_.size() + pool_.size()); +uint32_t WorkerPool::Size(const Language &language) const { + const auto state = states_by_lang_.find(language); + if (state == states_by_lang_.end()) { + return 0; + } else { + return static_cast(state->second.idle.size() + + state->second.idle_actor.size()); + } } -void WorkerPool::StartWorkerProcess(bool force_start) { - RAY_CHECK(!worker_command_.empty()) << "No worker command provided"; +void WorkerPool::StartWorkerProcess(const Language &language, bool force_start) { // The first condition makes sure that we are always starting up to // num_cpus_ number of processes in parallel. if (static_cast(starting_worker_processes_.size()) >= num_cpus_ && !force_start) { @@ -91,9 +103,11 @@ void WorkerPool::StartWorkerProcess(bool force_start) { << " worker processes pending registration"; return; } + auto &state = GetStateForLanguage(language); // Either there are no workers pending registration or the worker start is being forced. - RAY_LOG(DEBUG) << "starting worker, actor pool " << actor_pool_.size() << " task pool " - << pool_.size(); + RAY_LOG(DEBUG) << "Starting new worker process, current pool has " + << state.idle_actor.size() << " actor workers, and " << state.idle.size() + << " non-actor workers"; // Launch the process to create the worker. pid_t pid = fork(); @@ -108,7 +122,7 @@ void WorkerPool::StartWorkerProcess(bool force_start) { // Extract pointers from the worker command to pass into execvp. std::vector worker_command_args; - for (auto const &token : worker_command_) { + for (auto const &token : state.worker_command) { worker_command_args.push_back(token.c_str()); } worker_command_args.push_back(nullptr); @@ -123,7 +137,8 @@ void WorkerPool::StartWorkerProcess(bool force_start) { void WorkerPool::RegisterWorker(std::shared_ptr worker) { auto pid = worker->Pid(); RAY_LOG(DEBUG) << "Registering worker with pid " << pid; - registered_workers_.push_back(std::move(worker)); + auto &state = GetStateForLanguage(worker->GetLanguage()); + state.registered_workers.push_back(std::move(worker)); auto it = starting_worker_processes_.find(pid); RAY_CHECK(it != starting_worker_processes_.end()); @@ -135,55 +150,79 @@ void WorkerPool::RegisterWorker(std::shared_ptr worker) { void WorkerPool::RegisterDriver(std::shared_ptr driver) { RAY_CHECK(!driver->GetAssignedTaskId().is_nil()); - registered_drivers_.push_back(driver); + auto &state = GetStateForLanguage(driver->GetLanguage()); + state.registered_drivers.push_back(driver); } std::shared_ptr WorkerPool::GetRegisteredWorker( const std::shared_ptr &connection) const { - return GetWorker(registered_workers_, connection); + for (const auto &entry : states_by_lang_) { + auto worker = GetWorker(entry.second.registered_workers, connection); + if (worker != nullptr) { + return worker; + } + } + return nullptr; } std::shared_ptr WorkerPool::GetRegisteredDriver( const std::shared_ptr &connection) const { - return GetWorker(registered_drivers_, connection); + for (const auto &entry : states_by_lang_) { + auto driver = GetWorker(entry.second.registered_drivers, connection); + if (driver != nullptr) { + return driver; + } + } + return nullptr; } void WorkerPool::PushWorker(std::shared_ptr worker) { // Since the worker is now idle, unset its assigned task ID. RAY_CHECK(worker->GetAssignedTaskId().is_nil()) << "Idle workers cannot have an assigned task ID"; + auto &state = GetStateForLanguage(worker->GetLanguage()); // Add the worker to the idle pool. if (worker->GetActorId().is_nil()) { - pool_.push_back(std::move(worker)); + state.idle.push_back(std::move(worker)); } else { - actor_pool_[worker->GetActorId()] = std::move(worker); + state.idle_actor[worker->GetActorId()] = std::move(worker); } } -std::shared_ptr WorkerPool::PopWorker(const ActorID &actor_id) { +std::shared_ptr WorkerPool::PopWorker(const TaskSpecification &task_spec) { + auto &state = GetStateForLanguage(task_spec.GetLanguage()); + const auto &actor_id = task_spec.ActorId(); std::shared_ptr worker = nullptr; if (actor_id.is_nil()) { - if (!pool_.empty()) { - worker = std::move(pool_.back()); - pool_.pop_back(); + if (!state.idle.empty()) { + worker = std::move(state.idle.back()); + state.idle.pop_back(); } } else { - auto actor_entry = actor_pool_.find(actor_id); - if (actor_entry != actor_pool_.end()) { + auto actor_entry = state.idle_actor.find(actor_id); + if (actor_entry != state.idle_actor.end()) { worker = std::move(actor_entry->second); - actor_pool_.erase(actor_entry); + state.idle_actor.erase(actor_entry); } } return worker; } bool WorkerPool::DisconnectWorker(std::shared_ptr worker) { - RAY_CHECK(RemoveWorker(registered_workers_, worker)); - return RemoveWorker(pool_, worker); + auto &state = GetStateForLanguage(worker->GetLanguage()); + RAY_CHECK(RemoveWorker(state.registered_workers, worker)); + return RemoveWorker(state.idle, worker); } void WorkerPool::DisconnectDriver(std::shared_ptr driver) { - RAY_CHECK(RemoveWorker(registered_drivers_, driver)); + auto &state = GetStateForLanguage(driver->GetLanguage()); + RAY_CHECK(RemoveWorker(state.registered_drivers, driver)); +} + +inline WorkerPool::State &WorkerPool::GetStateForLanguage(const Language &language) { + auto state = states_by_lang_.find(language); + RAY_CHECK(state != states_by_lang_.end()) << "Required Language isn't supported."; + return state->second; } } // namespace raylet diff --git a/src/ray/raylet/worker_pool.h b/src/ray/raylet/worker_pool.h index 7850ae54e..d943bc6cc 100644 --- a/src/ray/raylet/worker_pool.h +++ b/src/ray/raylet/worker_pool.h @@ -7,6 +7,8 @@ #include #include "ray/common/client_connection.h" +#include "ray/gcs/format/util.h" +#include "ray/raylet/task.h" #include "ray/raylet/worker.h" namespace ray { @@ -26,11 +28,13 @@ class WorkerPool { /// the process should create and register the specified number of workers, /// and add them to the pool. /// - /// \param num_worker_processes The number of worker processes to start. + /// \param num_worker_processes The number of worker processes to start, per language. /// \param num_workers_per_process The number of workers per process. - /// \param worker_command The command used to start the worker process. - WorkerPool(int num_worker_processes, int num_workers_per_process, int num_cpus, - const std::vector &worker_command); + /// \param worker_commands The commands used to start the worker process, grouped by + /// language. + WorkerPool( + int num_worker_processes, int num_workers_per_process, int num_cpus, + const std::unordered_map> &worker_commands); /// Destructor responsible for freeing a set of workers owned by this class. virtual ~WorkerPool(); @@ -42,9 +46,10 @@ class WorkerPool { /// This function will start up to num_cpus many workers in parallel /// if it is called multiple times. /// + /// \param language Which language this worker process should be. /// \param force_start Controls whether to force starting a worker regardless of any /// workers that have already been started but not yet registered. - void StartWorkerProcess(bool force_start = false); + void StartWorkerProcess(const Language &language, bool force_start = false); /// 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). @@ -92,16 +97,17 @@ class WorkerPool { /// 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. /// - /// \param actor_id The returned worker must have this actor ID. - /// \return An idle worker with the requested actor ID. Returns nullptr if no + /// \param task_spec The returned worker must be able to execute this task. + /// \return An idle worker with the requested task spec. Returns nullptr if no /// such worker exists. - std::shared_ptr PopWorker(const ActorID &actor_id); + std::shared_ptr PopWorker(const TaskSpecification &task_spec); - /// Return the current size of the worker pool. Counts only the workers that registered - /// and requested a task. + /// Return the current size of the worker pool for the requested language. Counts only + /// idle workers. /// + /// \param language The requested language. /// \return The total count of all workers (actor and non-actor) in the pool. - uint32_t Size() const; + uint32_t Size(const Language &language) const; protected: /// A map from the pids of starting worker processes @@ -111,20 +117,31 @@ class WorkerPool { int num_workers_per_process_; private: + /// An internal data structure that maintains the pool state per language. + struct State { + /// The commands and arguments used to start the worker process + std::vector worker_command; + /// The pool of idle non-actor workers. + std::list> idle; + /// The pool of idle actor workers. + std::unordered_map> idle_actor; + /// 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> registered_workers; + /// All drivers that have registered and are still connected. + std::list> registered_drivers; + }; + + /// A helper function that returns the reference of the pool state + /// for a given language. + inline State &GetStateForLanguage(const Language &language); + /// The number of CPUs this Raylet has available. int num_cpus_; - /// The command and arguments used to start the worker. - std::vector worker_command_; - /// The pool of idle workers. - std::list> pool_; - /// The pool of idle actor workers. - std::unordered_map> actor_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> registered_workers_; - /// All drivers that have registered and are still connected. - std::list> registered_drivers_; + + /// Pool states per language. + std::unordered_map states_by_lang_; }; } // namespace raylet diff --git a/src/ray/raylet/worker_pool_test.cc b/src/ray/raylet/worker_pool_test.cc index b73f9ae3c..17b45b809 100644 --- a/src/ray/raylet/worker_pool_test.cc +++ b/src/ray/raylet/worker_pool_test.cc @@ -12,9 +12,13 @@ int NUM_WORKERS_PER_PROCESS = 3; class WorkerPoolMock : public WorkerPool { public: - WorkerPoolMock() : WorkerPool(0, NUM_WORKERS_PER_PROCESS, 0, {}) {} + WorkerPoolMock() + : WorkerPool(0, NUM_WORKERS_PER_PROCESS, 0, + {{Language::PYTHON, {"dummy_py_worker_command"}}, + {Language::JAVA, {"dummy_java_worker_command"}}}) {} - void StartWorkerProcess(pid_t pid, bool force_start = false) { + void StartWorkerProcess(pid_t pid, const Language &language = Language::PYTHON, + bool force_start = false) { if (starting_worker_processes_.size() > 0 && !force_start) { // Workers have been started, but not registered. Force start disabled -- returning. RAY_LOG(DEBUG) << starting_worker_processes_.size() @@ -22,7 +26,7 @@ class WorkerPoolMock : public WorkerPool { return; } // Either no workers are pending registration or the worker start is being forced. - RAY_LOG(DEBUG) << "starting new worker process, worker pool size " << Size(); + RAY_LOG(DEBUG) << "starting new worker process, worker pool size " << Size(language); starting_worker_processes_.emplace(std::make_pair(pid, num_workers_per_process_)); } @@ -33,7 +37,8 @@ class WorkerPoolTest : public ::testing::Test { public: WorkerPoolTest() : worker_pool_(), io_service_() {} - std::shared_ptr CreateWorker(pid_t pid) { + std::shared_ptr CreateWorker(pid_t pid, + const Language &language = Language::PYTHON) { std::function client_handler = [this](LocalClientConnection &client) { HandleNewClient(client); }; std::function, int64_t, const uint8_t *)> @@ -44,7 +49,7 @@ class WorkerPoolTest : public ::testing::Test { boost::asio::local::stream_protocol::socket socket(io_service_); auto client = LocalClientConnection::Create(client_handler, message_handler, std::move(socket), "worker"); - return std::shared_ptr(new Worker(pid, client)); + return std::shared_ptr(new Worker(pid, language, client)); } protected: @@ -56,6 +61,14 @@ class WorkerPoolTest : public ::testing::Test { void HandleMessage(std::shared_ptr, int64_t, const uint8_t *){}; }; +static inline TaskSpecification ExampleTaskSpec( + const ActorID actor_id = ActorID::nil(), + const Language &language = Language::PYTHON) { + return TaskSpecification(UniqueID::nil(), UniqueID::nil(), 0, ActorID::nil(), + ObjectID::nil(), actor_id, ActorHandleID::nil(), 0, + FunctionID::nil(), {}, 0, {}, language); +} + TEST_F(WorkerPoolTest, HandleWorkerRegistration) { pid_t pid = 1234; worker_pool_.StartWorkerProcess(pid); @@ -85,7 +98,8 @@ TEST_F(WorkerPoolTest, HandleWorkerRegistration) { TEST_F(WorkerPoolTest, HandleWorkerPushPop) { // Try to pop a worker from the empty pool and make sure we don't get one. std::shared_ptr popped_worker; - popped_worker = worker_pool_.PopWorker(ActorID::nil()); + const auto task_spec = ExampleTaskSpec(); + popped_worker = worker_pool_.PopWorker(task_spec); ASSERT_EQ(popped_worker, nullptr); // Create some workers. @@ -98,13 +112,13 @@ TEST_F(WorkerPoolTest, HandleWorkerPushPop) { } // Pop two workers and make sure they're one of the workers we created. - popped_worker = worker_pool_.PopWorker(ActorID::nil()); + popped_worker = worker_pool_.PopWorker(task_spec); ASSERT_NE(popped_worker, nullptr); ASSERT_TRUE(workers.count(popped_worker) > 0); - popped_worker = worker_pool_.PopWorker(ActorID::nil()); + popped_worker = worker_pool_.PopWorker(task_spec); ASSERT_NE(popped_worker, nullptr); ASSERT_TRUE(workers.count(popped_worker) > 0); - popped_worker = worker_pool_.PopWorker(ActorID::nil()); + popped_worker = worker_pool_.PopWorker(task_spec); ASSERT_EQ(popped_worker, nullptr); } @@ -115,19 +129,39 @@ TEST_F(WorkerPoolTest, PopActorWorker) { worker_pool_.PushWorker(worker); // Assign an actor ID to the worker. - auto actor = worker_pool_.PopWorker(ActorID::nil()); + const auto task_spec = ExampleTaskSpec(); + auto actor = worker_pool_.PopWorker(task_spec); auto actor_id = ActorID::from_random(); actor->AssignActorId(actor_id); worker_pool_.PushWorker(actor); // Check that there are no more non-actor workers. - ASSERT_EQ(worker_pool_.PopWorker(ActorID::nil()), nullptr); + ASSERT_EQ(worker_pool_.PopWorker(task_spec), nullptr); // Check that we can pop the actor worker. - actor = worker_pool_.PopWorker(actor_id); + const auto actor_task_spec = ExampleTaskSpec(actor_id); + actor = worker_pool_.PopWorker(actor_task_spec); ASSERT_EQ(actor, worker); ASSERT_EQ(actor->GetActorId(), actor_id); } +TEST_F(WorkerPoolTest, PopWorkersOfMultipleLanguages) { + // Create a Python Worker, and add it to the pool + auto py_worker = CreateWorker(1234, Language::PYTHON); + worker_pool_.PushWorker(py_worker); + // Check that no worker will be popped if the given task is a Java task + const auto java_task_spec = ExampleTaskSpec(ActorID::nil(), Language::JAVA); + ASSERT_EQ(worker_pool_.PopWorker(java_task_spec), nullptr); + // Check that the worker can be popped if the given task is a Python task + const auto py_task_spec = ExampleTaskSpec(ActorID::nil(), Language::PYTHON); + ASSERT_NE(worker_pool_.PopWorker(py_task_spec), nullptr); + + // Create a Java Worker, and add it to the pool + auto java_worker = CreateWorker(1234, Language::JAVA); + worker_pool_.PushWorker(java_worker); + // Check that the worker will be popped now for Java task + ASSERT_NE(worker_pool_.PopWorker(java_task_spec), nullptr); +} + } // namespace raylet } // namespace ray