[Java worker] Migrate task execution and submission on top of core worker (#5370)

This commit is contained in:
Kai Yang
2019-08-16 13:52:13 +08:00
committed by Hao Chen
parent 3a853121b9
commit b1aae0e398
95 changed files with 3069 additions and 2991 deletions
-10
View File
@@ -64,16 +64,6 @@ class TaskArg {
const std::shared_ptr<Buffer> data_;
};
/// Information of a task
struct TaskInfo {
/// The ID of task.
const TaskID task_id;
/// The job ID.
const JobID job_id;
/// The type of task.
const TaskType task_type;
};
enum class StoreProviderType { LOCAL_PLASMA, PLASMA, MEMORY };
enum class TaskTransportType { RAYLET, DIRECT_ACTOR };
+20 -1
View File
@@ -6,7 +6,10 @@ namespace ray {
/// per-thread context for core worker.
struct WorkerThreadContext {
WorkerThreadContext()
: current_task_id_(TaskID::ForFakeTask()), task_index_(0), put_index_(0) {}
: current_task_id_(TaskID::ForFakeTask()),
current_actor_id_(ActorID::Nil()),
task_index_(0),
put_index_(0) {}
int GetNextTaskIndex() { return ++task_index_; }
@@ -18,6 +21,8 @@ struct WorkerThreadContext {
return current_task_;
}
const ActorID &GetCurrentActorID() const { return current_actor_id_; }
void SetCurrentTaskId(const TaskID &task_id) {
current_task_id_ = task_id;
task_index_ = 0;
@@ -27,12 +32,22 @@ struct WorkerThreadContext {
void SetCurrentTask(const TaskSpecification &task_spec) {
SetCurrentTaskId(task_spec.TaskId());
current_task_ = std::make_shared<const TaskSpecification>(task_spec);
if (task_spec.IsActorCreationTask()) {
RAY_CHECK(current_actor_id_.IsNil());
current_actor_id_ = task_spec.ActorCreationId();
}
if (task_spec.IsActorTask()) {
RAY_CHECK(current_actor_id_ == task_spec.ActorId());
}
}
private:
/// The task ID for current task.
TaskID current_task_id_;
/// ID of current actor.
ActorID current_actor_id_;
/// The current task.
std::shared_ptr<const TaskSpecification> current_task_;
@@ -81,6 +96,10 @@ std::shared_ptr<const TaskSpecification> WorkerContext::GetCurrentTask() const {
return GetThreadContext().GetCurrentTask();
}
const ActorID &WorkerContext::GetCurrentActorID() const {
return GetThreadContext().GetCurrentActorID();
}
WorkerThreadContext &WorkerContext::GetThreadContext() {
if (thread_context_ == nullptr) {
thread_context_ = std::unique_ptr<WorkerThreadContext>(new WorkerThreadContext());
+2
View File
@@ -24,6 +24,8 @@ class WorkerContext {
std::shared_ptr<const TaskSpecification> GetCurrentTask() const;
const ActorID &GetCurrentActorID() const;
int GetNextTaskIndex();
int GetNextPutIndex();
+6
View File
@@ -47,6 +47,12 @@ CoreWorker::~CoreWorker() {
gcs_client_->Disconnect();
io_service_.stop();
io_thread_.join();
if (task_execution_interface_) {
task_execution_interface_->Stop();
}
if (raylet_client_) {
RAY_IGNORE_EXPR(raylet_client_->Disconnect());
}
}
void CoreWorker::StartIOService() { io_service_.run(); }
+4
View File
@@ -38,6 +38,10 @@ class CoreWorker {
/// Language of this worker.
Language GetLanguage() const { return language_; }
WorkerContext &GetWorkerContext() { return worker_context_; }
RayletClient &GetRayletClient() { return *raylet_client_; }
/// Return the `CoreWorkerTaskInterface` that contains the methods related to task
/// submisson.
CoreWorkerTaskInterface &Tasks() { return *task_interface_; }
+130 -6
View File
@@ -3,6 +3,9 @@
jclass java_boolean_class;
jmethodID java_boolean_init;
jclass java_double_class;
jmethodID java_double_double_value;
jclass java_list_class;
jmethodID java_list_size;
jmethodID java_list_get;
@@ -12,18 +15,62 @@ jclass java_array_list_class;
jmethodID java_array_list_init;
jmethodID java_array_list_init_with_capacity;
jclass java_map_class;
jmethodID java_map_entry_set;
jclass java_set_class;
jmethodID java_set_iterator;
jclass java_iterator_class;
jmethodID java_iterator_has_next;
jmethodID java_iterator_next;
jclass java_map_entry_class;
jmethodID java_map_entry_get_key;
jmethodID java_map_entry_get_value;
jclass java_ray_exception_class;
jclass java_base_id_class;
jmethodID java_base_id_get_bytes;
jclass java_function_descriptor_class;
jmethodID java_function_descriptor_get_language;
jmethodID java_function_descriptor_to_list;
jclass java_language_class;
jmethodID java_language_get_number;
jclass java_function_arg_class;
jfieldID java_function_arg_id;
jfieldID java_function_arg_data;
jclass java_base_task_options_class;
jfieldID java_base_task_options_resources;
jclass java_actor_creation_options_class;
jfieldID java_actor_creation_options_max_reconstructions;
jfieldID java_actor_creation_options_jvm_options;
jclass java_gcs_client_options_class;
jfieldID java_gcs_client_options_ip;
jfieldID java_gcs_client_options_port;
jfieldID java_gcs_client_options_password;
jclass java_native_ray_object_class;
jmethodID java_native_ray_object_init;
jfieldID java_native_ray_object_data;
jfieldID java_native_ray_object_metadata;
jint JNI_VERSION = JNI_VERSION_1_8;
jclass java_task_executor_class;
jmethodID java_task_executor_execute;
JavaVM *jvm;
inline jclass LoadClass(JNIEnv *env, const char *class_name) {
jclass tempLocalClassRef = env->FindClass(class_name);
jclass ret = (jclass)env->NewGlobalRef(tempLocalClassRef);
RAY_CHECK(ret) << "Can't load Java class " << class_name;
env->DeleteLocalRef(tempLocalClassRef);
return ret;
}
@@ -31,13 +78,18 @@ inline jclass LoadClass(JNIEnv *env, const char *class_name) {
/// Load and cache frequently-used Java classes and methods
jint JNI_OnLoad(JavaVM *vm, void *reserved) {
JNIEnv *env;
if (vm->GetEnv(reinterpret_cast<void **>(&env), JNI_VERSION) != JNI_OK) {
if (vm->GetEnv(reinterpret_cast<void **>(&env), CURRENT_JNI_VERSION) != JNI_OK) {
return JNI_ERR;
}
jvm = vm;
java_boolean_class = LoadClass(env, "java/lang/Boolean");
java_boolean_init = env->GetMethodID(java_boolean_class, "<init>", "(Z)V");
java_double_class = LoadClass(env, "java/lang/Double");
java_double_double_value = env->GetMethodID(java_double_class, "doubleValue", "()D");
java_list_class = LoadClass(env, "java/util/List");
java_list_size = env->GetMethodID(java_list_class, "size", "()I");
java_list_get = env->GetMethodID(java_list_class, "get", "(I)Ljava/lang/Object;");
@@ -48,10 +100,65 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
java_array_list_init_with_capacity =
env->GetMethodID(java_array_list_class, "<init>", "(I)V");
java_map_class = LoadClass(env, "java/util/Map");
java_map_entry_set = env->GetMethodID(java_map_class, "entrySet", "()Ljava/util/Set;");
java_set_class = LoadClass(env, "java/util/Set");
java_set_iterator =
env->GetMethodID(java_set_class, "iterator", "()Ljava/util/Iterator;");
java_iterator_class = LoadClass(env, "java/util/Iterator");
java_iterator_has_next = env->GetMethodID(java_iterator_class, "hasNext", "()Z");
java_iterator_next =
env->GetMethodID(java_iterator_class, "next", "()Ljava/lang/Object;");
java_map_entry_class = LoadClass(env, "java/util/Map$Entry");
java_map_entry_get_key =
env->GetMethodID(java_map_entry_class, "getKey", "()Ljava/lang/Object;");
java_map_entry_get_value =
env->GetMethodID(java_map_entry_class, "getValue", "()Ljava/lang/Object;");
java_ray_exception_class = LoadClass(env, "org/ray/api/exception/RayException");
java_native_ray_object_class =
LoadClass(env, "org/ray/runtime/objectstore/NativeRayObject");
java_base_id_class = LoadClass(env, "org/ray/api/id/BaseId");
java_base_id_get_bytes = env->GetMethodID(java_base_id_class, "getBytes", "()[B");
java_function_descriptor_class =
LoadClass(env, "org/ray/runtime/functionmanager/FunctionDescriptor");
java_function_descriptor_get_language =
env->GetMethodID(java_function_descriptor_class, "getLanguage",
"()Lorg/ray/runtime/generated/Common$Language;");
java_function_descriptor_to_list =
env->GetMethodID(java_function_descriptor_class, "toList", "()Ljava/util/List;");
java_language_class = LoadClass(env, "org/ray/runtime/generated/Common$Language");
java_language_get_number = env->GetMethodID(java_language_class, "getNumber", "()I");
java_function_arg_class = LoadClass(env, "org/ray/runtime/task/FunctionArg");
java_function_arg_id =
env->GetFieldID(java_function_arg_class, "id", "Lorg/ray/api/id/ObjectId;");
java_function_arg_data = env->GetFieldID(java_function_arg_class, "data", "[B");
java_base_task_options_class = LoadClass(env, "org/ray/api/options/BaseTaskOptions");
java_base_task_options_resources =
env->GetFieldID(java_base_task_options_class, "resources", "Ljava/util/Map;");
java_actor_creation_options_class =
LoadClass(env, "org/ray/api/options/ActorCreationOptions");
java_actor_creation_options_max_reconstructions =
env->GetFieldID(java_actor_creation_options_class, "maxReconstructions", "I");
java_actor_creation_options_jvm_options = env->GetFieldID(
java_actor_creation_options_class, "jvmOptions", "Ljava/lang/String;");
java_gcs_client_options_class = LoadClass(env, "org/ray/runtime/gcs/GcsClientOptions");
java_gcs_client_options_ip =
env->GetFieldID(java_gcs_client_options_class, "ip", "Ljava/lang/String;");
java_gcs_client_options_port =
env->GetFieldID(java_gcs_client_options_class, "port", "I");
java_gcs_client_options_password =
env->GetFieldID(java_gcs_client_options_class, "password", "Ljava/lang/String;");
java_native_ray_object_class = LoadClass(env, "org/ray/runtime/object/NativeRayObject");
java_native_ray_object_init =
env->GetMethodID(java_native_ray_object_class, "<init>", "([B[B)V");
java_native_ray_object_data =
@@ -59,17 +166,34 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
java_native_ray_object_metadata =
env->GetFieldID(java_native_ray_object_class, "metadata", "[B");
return JNI_VERSION;
java_task_executor_class = LoadClass(env, "org/ray/runtime/task/TaskExecutor");
java_task_executor_execute =
env->GetMethodID(java_task_executor_class, "execute",
"(Ljava/util/List;Ljava/util/List;)Ljava/util/List;");
return CURRENT_JNI_VERSION;
}
/// Unload java classes
void JNI_OnUnload(JavaVM *vm, void *reserved) {
JNIEnv *env;
vm->GetEnv(reinterpret_cast<void **>(&env), JNI_VERSION);
vm->GetEnv(reinterpret_cast<void **>(&env), CURRENT_JNI_VERSION);
env->DeleteGlobalRef(java_boolean_class);
env->DeleteGlobalRef(java_double_class);
env->DeleteGlobalRef(java_list_class);
env->DeleteGlobalRef(java_array_list_class);
env->DeleteGlobalRef(java_map_class);
env->DeleteGlobalRef(java_set_class);
env->DeleteGlobalRef(java_iterator_class);
env->DeleteGlobalRef(java_map_entry_class);
env->DeleteGlobalRef(java_ray_exception_class);
env->DeleteGlobalRef(java_base_id_class);
env->DeleteGlobalRef(java_function_descriptor_class);
env->DeleteGlobalRef(java_language_class);
env->DeleteGlobalRef(java_function_arg_class);
env->DeleteGlobalRef(java_base_task_options_class);
env->DeleteGlobalRef(java_actor_creation_options_class);
env->DeleteGlobalRef(java_native_ray_object_class);
env->DeleteGlobalRef(java_task_executor_class);
}
+161 -37
View File
@@ -1,5 +1,5 @@
#ifndef RAY_COMMON_JAVA_JNI_HELPER_H
#define RAY_COMMON_JAVA_JNI_HELPER_H
#ifndef RAY_COMMON_JAVA_JNI_UTILS_H
#define RAY_COMMON_JAVA_JNI_UTILS_H
#include <jni.h>
#include "ray/common/buffer.h"
@@ -12,6 +12,11 @@ extern jclass java_boolean_class;
/// Constructor of Boolean class
extern jmethodID java_boolean_init;
/// Double class
extern jclass java_double_class;
/// doubleValue method of Double class
extern jmethodID java_double_double_value;
/// List class
extern jclass java_list_class;
/// size method of List class
@@ -28,9 +33,78 @@ extern jmethodID java_array_list_init;
/// Constructor of ArrayList class with single parameter capacity
extern jmethodID java_array_list_init_with_capacity;
/// Map interface
extern jclass java_map_class;
/// entrySet method of Map interface
extern jmethodID java_map_entry_set;
/// Set interface
extern jclass java_set_class;
/// iterator method of Set interface
extern jmethodID java_set_iterator;
/// Iterator interface
extern jclass java_iterator_class;
/// hasNext method of Iterator interface
extern jmethodID java_iterator_has_next;
/// next method of Iterator interface
extern jmethodID java_iterator_next;
/// Map.Entry interface
extern jclass java_map_entry_class;
/// getKey method of Map.Entry interface
extern jmethodID java_map_entry_get_key;
/// getValue method of Map.Entry interface
extern jmethodID java_map_entry_get_value;
/// RayException class
extern jclass java_ray_exception_class;
/// BaseId class
extern jclass java_base_id_class;
/// getBytes method of BaseId class
extern jmethodID java_base_id_get_bytes;
/// FunctionDescriptor interface
extern jclass java_function_descriptor_class;
/// getLanguage method of FunctionDescriptor interface
extern jmethodID java_function_descriptor_get_language;
/// toList method of FunctionDescriptor interface
extern jmethodID java_function_descriptor_to_list;
/// Language class
extern jclass java_language_class;
/// getNumber of Language class
extern jmethodID java_language_get_number;
/// NativeTaskArg class
extern jclass java_function_arg_class;
/// id field of NativeTaskArg class
extern jfieldID java_function_arg_id;
/// data field of NativeTaskArg class
extern jfieldID java_function_arg_data;
/// BaseTaskOptions class
extern jclass java_base_task_options_class;
/// resources field of BaseTaskOptions class
extern jfieldID java_base_task_options_resources;
/// ActorCreationOptions class
extern jclass java_actor_creation_options_class;
/// maxReconstructions field of ActorCreationOptions class
extern jfieldID java_actor_creation_options_max_reconstructions;
/// jvmOptions field of ActorCreationOptions class
extern jfieldID java_actor_creation_options_jvm_options;
/// GcsClientOptions class
extern jclass java_gcs_client_options_class;
/// ip field of GcsClientOptions class
extern jfieldID java_gcs_client_options_ip;
/// port field of GcsClientOptions class
extern jfieldID java_gcs_client_options_port;
/// password field of GcsClientOptions class
extern jfieldID java_gcs_client_options_password;
/// NativeRayObject class
extern jclass java_native_ray_object_class;
/// Constructor of NativeRayObject class
@@ -40,6 +114,15 @@ extern jfieldID java_native_ray_object_data;
/// metadata field of NativeRayObject class
extern jfieldID java_native_ray_object_metadata;
/// TaskExecutor class
extern jclass java_task_executor_class;
/// execute method of TaskExecutor class
extern jmethodID java_task_executor_execute;
#define CURRENT_JNI_VERSION JNI_VERSION_1_8
extern JavaVM *jvm;
/// Throws a Java RayException if the status is not OK.
#define THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, ret) \
{ \
@@ -49,6 +132,32 @@ extern jfieldID java_native_ray_object_metadata;
} \
}
/// Represents a byte buffer of Java byte array.
/// The destructor will automatically call ReleaseByteArrayElements.
/// NOTE: Instances of this class cannot be used across threads.
class JavaByteArrayBuffer : public ray::Buffer {
public:
JavaByteArrayBuffer(JNIEnv *env, jbyteArray java_byte_array)
: env_(env), java_byte_array_(java_byte_array) {
native_bytes_ = env_->GetByteArrayElements(java_byte_array_, nullptr);
}
uint8_t *Data() const override { return reinterpret_cast<uint8_t *>(native_bytes_); }
size_t Size() const override { return env_->GetArrayLength(java_byte_array_); }
bool OwnsData() const override { return true; }
~JavaByteArrayBuffer() {
env_->ReleaseByteArrayElements(java_byte_array_, native_bytes_, JNI_ABORT);
}
private:
JNIEnv *env_;
jbyteArray java_byte_array_;
jbyte *native_bytes_;
};
/// Convert a Java byte array to a C++ UniqueID.
template <typename ID>
inline ID JavaByteArrayToId(JNIEnv *env, const jbyteArray &bytes) {
@@ -95,6 +204,15 @@ inline void JavaListToNativeVector(
}
}
/// Convert a Java List<String> to C++ std::vector<std::string>.
inline void JavaStringListToNativeStringVector(JNIEnv *env, jobject java_list,
std::vector<std::string> *native_vector) {
JavaListToNativeVector<std::string>(
env, java_list, native_vector, [](JNIEnv *env, jobject jstr) {
return JavaStringToNativeString(env, static_cast<jstring>(jstr));
});
}
/// Convert a C++ std::vector to a Java List.
template <typename NativeT>
inline jobject NativeVectorToJavaList(
@@ -109,6 +227,22 @@ inline jobject NativeVectorToJavaList(
return java_list;
}
/// Convert a C++ std::vector<std::string> to a Java List<String>
inline jobject NativeStringVectorToJavaStringList(
JNIEnv *env, const std::vector<std::string> &native_vector) {
return NativeVectorToJavaList<std::string>(
env, native_vector,
[](JNIEnv *env, const std::string &str) { return env->NewStringUTF(str.c_str()); });
}
template <typename ID>
inline jobject NativeIdVectorToJavaByteArrayList(JNIEnv *env,
const std::vector<ID> &native_vector) {
return NativeVectorToJavaList<ID>(env, native_vector, [](JNIEnv *env, const ID &id) {
return IdToJavaByteArray<ID>(env, id);
});
}
/// Convert a C++ ray::Buffer to a Java byte array.
inline jbyteArray NativeBufferToJavaByteArray(JNIEnv *env,
const std::shared_ptr<ray::Buffer> buffer) {
@@ -123,50 +257,40 @@ inline jbyteArray NativeBufferToJavaByteArray(JNIEnv *env,
return java_byte_array;
}
/// A helper method to help access a Java NativeRayObject instance and ensure memory
/// safety.
///
/// \param[in] java_obj The Java NativeRayObject object.
/// \param[in] reader The callback function to access a C++ ray::RayObject instance.
/// \return The return value of callback function.
template <typename ReturnT>
inline ReturnT ReadJavaNativeRayObject(
JNIEnv *env, const jobject &java_obj,
std::function<ReturnT(const std::shared_ptr<ray::RayObject> &)> reader) {
/// Convert a Java byte[] as a C++ std::shared_ptr<JavaByteArrayBuffer>.
inline std::shared_ptr<JavaByteArrayBuffer> JavaByteArrayToNativeBuffer(
JNIEnv *env, const jbyteArray &javaByteArray) {
if (!javaByteArray) {
return nullptr;
}
return std::make_shared<JavaByteArrayBuffer>(env, javaByteArray);
}
/// Convert a Java NativeRayObject to a C++ ray::RayObject.
/// NOTE: the returned ray::RayObject cannot be used across threads.
inline std::shared_ptr<ray::RayObject> JavaNativeRayObjectToNativeRayObject(
JNIEnv *env, const jobject &java_obj) {
if (!java_obj) {
return reader(nullptr);
return nullptr;
}
auto java_data = (jbyteArray)env->GetObjectField(java_obj, java_native_ray_object_data);
auto java_metadata =
(jbyteArray)env->GetObjectField(java_obj, java_native_ray_object_metadata);
auto data_size = env->GetArrayLength(java_data);
jbyte *data = data_size > 0 ? env->GetByteArrayElements(java_data, nullptr) : nullptr;
auto metadata_size = java_metadata ? env->GetArrayLength(java_metadata) : 0;
jbyte *metadata =
metadata_size > 0 ? env->GetByteArrayElements(java_metadata, nullptr) : nullptr;
auto data_buffer = std::make_shared<ray::LocalMemoryBuffer>(
reinterpret_cast<uint8_t *>(data), data_size);
auto metadata_buffer = java_metadata
? std::make_shared<ray::LocalMemoryBuffer>(
reinterpret_cast<uint8_t *>(metadata), metadata_size)
: nullptr;
auto native_obj = std::make_shared<ray::RayObject>(data_buffer, metadata_buffer);
auto result = reader(native_obj);
if (data) {
env->ReleaseByteArrayElements(java_data, data, JNI_ABORT);
std::shared_ptr<ray::Buffer> data_buffer = JavaByteArrayToNativeBuffer(env, java_data);
std::shared_ptr<ray::Buffer> metadata_buffer =
JavaByteArrayToNativeBuffer(env, java_metadata);
if (!data_buffer) {
data_buffer = std::make_shared<ray::LocalMemoryBuffer>(nullptr, 0);
}
if (metadata) {
env->ReleaseByteArrayElements(java_metadata, metadata, JNI_ABORT);
if (!metadata_buffer) {
metadata_buffer = std::make_shared<ray::LocalMemoryBuffer>(nullptr, 0);
}
return result;
return std::make_shared<ray::RayObject>(data_buffer, metadata_buffer);
}
/// Convert a C++ ray::RayObject to a Java NativeRayObject.
inline jobject ToJavaNativeRayObject(JNIEnv *env,
const std::shared_ptr<ray::RayObject> &rayObject) {
inline jobject NativeRayObjectToJavaNativeRayObject(
JNIEnv *env, const std::shared_ptr<ray::RayObject> &rayObject) {
if (!rayObject) {
return nullptr;
}
@@ -177,4 +301,4 @@ inline jobject ToJavaNativeRayObject(JNIEnv *env,
return java_obj;
}
#endif // RAY_COMMON_JAVA_JNI_HELPER_H
#endif // RAY_COMMON_JAVA_JNI_UTILS_H
@@ -0,0 +1,134 @@
#include "ray/core_worker/lib/java/org_ray_runtime_RayNativeRuntime.h"
#include <jni.h>
#include <sstream>
#include "ray/common/id.h"
#include "ray/core_worker/core_worker.h"
#include "ray/core_worker/lib/java/jni_utils.h"
thread_local JNIEnv *local_env = nullptr;
thread_local jobject local_java_task_executor = nullptr;
inline ray::gcs::GcsClientOptions ToGcsClientOptions(JNIEnv *env,
jobject gcs_client_options) {
std::string ip = JavaStringToNativeString(
env, (jstring)env->GetObjectField(gcs_client_options, java_gcs_client_options_ip));
int port = env->GetIntField(gcs_client_options, java_gcs_client_options_port);
std::string password = JavaStringToNativeString(
env,
(jstring)env->GetObjectField(gcs_client_options, java_gcs_client_options_password));
return ray::gcs::GcsClientOptions(ip, port, password, /*is_test_client=*/false);
}
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeInitCoreWorker
* Signature:
* (ILjava/lang/String;Ljava/lang/String;[BLorg/ray/runtime/gcs/GcsClientOptions;)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeInitCoreWorker(
JNIEnv *env, jclass, jint workerMode, jstring storeSocket, jstring rayletSocket,
jbyteArray jobId, jobject gcsClientOptions) {
auto native_store_socket = JavaStringToNativeString(env, storeSocket);
auto native_raylet_socket = JavaStringToNativeString(env, rayletSocket);
auto job_id = JavaByteArrayToId<ray::JobID>(env, jobId);
auto gcs_client_options = ToGcsClientOptions(env, gcsClientOptions);
auto executor_func = [](const ray::RayFunction &ray_function,
const std::vector<std::shared_ptr<ray::RayObject>> &args,
int num_returns,
std::vector<std::shared_ptr<ray::RayObject>> *results) {
JNIEnv *env = local_env;
RAY_CHECK(env);
RAY_CHECK(local_java_task_executor);
// convert RayFunction
jobject ray_function_array_list =
NativeStringVectorToJavaStringList(env, ray_function.function_descriptor);
// convert args
// TODO (kfstorm): Avoid copying binary data from Java to C++
jobject args_array_list = NativeVectorToJavaList<std::shared_ptr<ray::RayObject>>(
env, args, NativeRayObjectToJavaNativeRayObject);
// invoke Java method
jobject java_return_objects =
env->CallObjectMethod(local_java_task_executor, java_task_executor_execute,
ray_function_array_list, args_array_list);
std::vector<std::shared_ptr<ray::RayObject>> return_objects;
JavaListToNativeVector<std::shared_ptr<ray::RayObject>>(
env, java_return_objects, &return_objects,
[](JNIEnv *env, jobject java_native_ray_object) {
return JavaNativeRayObjectToNativeRayObject(env, java_native_ray_object);
});
for (auto &obj : return_objects) {
results->push_back(obj);
}
return ray::Status::OK();
};
try {
auto core_worker = new ray::CoreWorker(
static_cast<ray::WorkerType>(workerMode), ::Language::JAVA, native_store_socket,
native_raylet_socket, job_id, gcs_client_options, executor_func);
return reinterpret_cast<jlong>(core_worker);
} catch (const std::exception &e) {
std::ostringstream oss;
oss << "Failed to construct core worker: " << e.what();
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, ray::Status::Invalid(oss.str()), 0);
return 0; // To make compiler no complain
}
}
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeRunTaskExecutor
* Signature: (JLorg/ray/runtime/task/TaskExecutor;)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeRunTaskExecutor(
JNIEnv *env, jclass o, jlong nativeCoreWorkerPointer, jobject javaTaskExecutor) {
local_env = env;
local_java_task_executor = javaTaskExecutor;
auto core_worker = reinterpret_cast<ray::CoreWorker *>(nativeCoreWorkerPointer);
core_worker->Execution().Run();
local_env = nullptr;
local_java_task_executor = nullptr;
}
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeDestroyCoreWorker
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeDestroyCoreWorker(
JNIEnv *env, jclass o, jlong nativeCoreWorkerPointer) {
delete reinterpret_cast<ray::CoreWorker *>(nativeCoreWorkerPointer);
}
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeSetup
* Signature: (Ljava/lang/String;)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeSetup(JNIEnv *env,
jclass,
jstring logDir) {
std::string log_dir = JavaStringToNativeString(env, logDir);
ray::RayLog::StartRayLog("java_worker", ray::RayLogLevel::INFO, log_dir);
// TODO (kfstorm): If we add InstallFailureSignalHandler here, Java test may crash.
}
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeShutdownHook
* Signature: ()V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeShutdownHook(JNIEnv *,
jclass) {
ray::RayLog::ShutDownRayLog();
}
#ifdef __cplusplus
}
#endif
@@ -0,0 +1,54 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_RayNativeRuntime */
#ifndef _Included_org_ray_runtime_RayNativeRuntime
#define _Included_org_ray_runtime_RayNativeRuntime
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeInitCoreWorker
* Signature:
* (ILjava/lang/String;Ljava/lang/String;[BLorg/ray/runtime/gcs/GcsClientOptions;)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeInitCoreWorker(
JNIEnv *, jclass, jint, jstring, jstring, jbyteArray, jobject);
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeRunTaskExecutor
* Signature: (JLorg/ray/runtime/task/TaskExecutor;)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeRunTaskExecutor(
JNIEnv *, jclass, jlong, jobject);
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeDestroyCoreWorker
* Signature: (J)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_RayNativeRuntime_nativeDestroyCoreWorker(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeSetup
* Signature: (Ljava/lang/String;)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeSetup(JNIEnv *, jclass,
jstring);
/*
* Class: org_ray_runtime_RayNativeRuntime
* Method: nativeShutdownHook
* Signature: ()V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_RayNativeRuntime_nativeShutdownHook(JNIEnv *,
jclass);
#ifdef __cplusplus
}
#endif
#endif
@@ -1,134 +0,0 @@
#include "ray/core_worker/lib/java/org_ray_runtime_WorkerContext.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/context.h"
#include "ray/core_worker/lib/java/jni_utils.h"
inline ray::WorkerContext *GetWorkerContextFromPointer(
jlong nativeWorkerContextFromPointer) {
return reinterpret_cast<ray::WorkerContext *>(nativeWorkerContextFromPointer);
}
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeCreateWorkerContext
* Signature: (I[B)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_WorkerContext_nativeCreateWorkerContext(
JNIEnv *env, jclass, jint workerType, jbyteArray jobId) {
return reinterpret_cast<jlong>(
new ray::WorkerContext(static_cast<ray::rpc::WorkerType>(workerType),
JavaByteArrayToId<ray::JobID>(env, jobId)));
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentTaskId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL Java_org_ray_runtime_WorkerContext_nativeGetCurrentTaskId(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer) {
auto task_id =
GetWorkerContextFromPointer(nativeWorkerContextFromPointer)->GetCurrentTaskID();
return IdToJavaByteArray<ray::TaskID>(env, task_id);
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeSetCurrentTask
* Signature: (J[B)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_WorkerContext_nativeSetCurrentTask(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer, jbyteArray taskSpec) {
jbyte *data = env->GetByteArrayElements(taskSpec, NULL);
jsize size = env->GetArrayLength(taskSpec);
ray::rpc::TaskSpec task_spec_message;
task_spec_message.ParseFromArray(data, size);
env->ReleaseByteArrayElements(taskSpec, data, JNI_ABORT);
ray::TaskSpecification spec(task_spec_message);
GetWorkerContextFromPointer(nativeWorkerContextFromPointer)->SetCurrentTask(spec);
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentTask
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL Java_org_ray_runtime_WorkerContext_nativeGetCurrentTask(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer) {
auto spec =
GetWorkerContextFromPointer(nativeWorkerContextFromPointer)->GetCurrentTask();
if (!spec) {
return nullptr;
}
auto task_message = spec->Serialize();
jbyteArray result = env->NewByteArray(task_message.size());
env->SetByteArrayRegion(
result, 0, task_message.size(),
reinterpret_cast<jbyte *>(const_cast<char *>(task_message.data())));
return result;
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentJobId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_WorkerContext_nativeGetCurrentJobId(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer) {
const auto &job_id =
GetWorkerContextFromPointer(nativeWorkerContextFromPointer)->GetCurrentJobID();
return IdToJavaByteBuffer<ray::JobID>(env, job_id);
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentWorkerId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL Java_org_ray_runtime_WorkerContext_nativeGetCurrentWorkerId(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer) {
auto worker_id =
GetWorkerContextFromPointer(nativeWorkerContextFromPointer)->GetWorkerID();
return IdToJavaByteArray<ray::WorkerID>(env, worker_id);
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetNextTaskIndex
* Signature: (J)I
*/
JNIEXPORT jint JNICALL Java_org_ray_runtime_WorkerContext_nativeGetNextTaskIndex(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer) {
return GetWorkerContextFromPointer(nativeWorkerContextFromPointer)->GetNextTaskIndex();
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetNextPutIndex
* Signature: (J)I
*/
JNIEXPORT jint JNICALL Java_org_ray_runtime_WorkerContext_nativeGetNextPutIndex(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer) {
return GetWorkerContextFromPointer(nativeWorkerContextFromPointer)->GetNextPutIndex();
}
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeDestroy
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_WorkerContext_nativeDestroy(
JNIEnv *env, jclass, jlong nativeWorkerContextFromPointer) {
delete GetWorkerContextFromPointer(nativeWorkerContextFromPointer);
}
#ifdef __cplusplus
}
#endif
@@ -1,87 +0,0 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_WorkerContext */
#ifndef _Included_org_ray_runtime_WorkerContext
#define _Included_org_ray_runtime_WorkerContext
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeCreateWorkerContext
* Signature: (I[B)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_WorkerContext_nativeCreateWorkerContext(
JNIEnv *, jclass, jint, jbyteArray);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentTaskId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_WorkerContext_nativeGetCurrentTaskId(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeSetCurrentTask
* Signature: (J[B)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_WorkerContext_nativeSetCurrentTask(
JNIEnv *, jclass, jlong, jbyteArray);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentTask
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_WorkerContext_nativeGetCurrentTask(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentJobId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_WorkerContext_nativeGetCurrentJobId(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetCurrentWorkerId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_WorkerContext_nativeGetCurrentWorkerId(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetNextTaskIndex
* Signature: (J)I
*/
JNIEXPORT jint JNICALL Java_org_ray_runtime_WorkerContext_nativeGetNextTaskIndex(JNIEnv *,
jclass,
jlong);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeGetNextPutIndex
* Signature: (J)I
*/
JNIEXPORT jint JNICALL Java_org_ray_runtime_WorkerContext_nativeGetNextPutIndex(JNIEnv *,
jclass,
jlong);
/*
* Class: org_ray_runtime_WorkerContext
* Method: nativeDestroy
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_WorkerContext_nativeDestroy(JNIEnv *, jclass,
jlong);
#ifdef __cplusplus
}
#endif
#endif
@@ -0,0 +1,113 @@
#include "ray/core_worker/lib/java/org_ray_runtime_actor_NativeRayActor.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/common.h"
#include "ray/core_worker/lib/java/jni_utils.h"
#include "ray/core_worker/task_interface.h"
inline ray::ActorHandle &GetActorHandle(jlong nativeActorHandle) {
return *(reinterpret_cast<ray::ActorHandle *>(nativeActorHandle));
}
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeFork
* Signature: (J)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeFork(
JNIEnv *env, jclass o, jlong nativeActorHandle) {
auto new_actor_handle = GetActorHandle(nativeActorHandle).Fork();
return reinterpret_cast<jlong>(new ray::ActorHandle(new_actor_handle));
}
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetActorId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeGetActorId(
JNIEnv *env, jclass o, jlong nativeActorHandle) {
return IdToJavaByteArray<ray::ActorID>(env,
GetActorHandle(nativeActorHandle).ActorID());
}
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetActorHandleId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeGetActorHandleId(
JNIEnv *env, jclass o, jlong nativeActorHandle) {
return IdToJavaByteArray<ray::ActorHandleID>(
env, GetActorHandle(nativeActorHandle).ActorHandleID());
}
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetLanguage
* Signature: (J)I
*/
JNIEXPORT jint JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeGetLanguage(
JNIEnv *env, jclass o, jlong nativeActorHandle) {
return (jint)GetActorHandle(nativeActorHandle).ActorLanguage();
}
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetActorCreationTaskFunctionDescriptor
* Signature: (J)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeGetActorCreationTaskFunctionDescriptor(
JNIEnv *env, jclass o, jlong nativeActorHandle) {
return NativeStringVectorToJavaStringList(
env, GetActorHandle(nativeActorHandle).ActorCreationTaskFunctionDescriptor());
}
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeSerialize
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeSerialize(
JNIEnv *env, jclass o, jlong nativeActorHandle) {
std::string output;
GetActorHandle(nativeActorHandle).Serialize(&output);
jbyteArray bytes = env->NewByteArray(output.size());
env->SetByteArrayRegion(bytes, 0, output.size(),
reinterpret_cast<const jbyte *>(output.c_str()));
return bytes;
}
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeDeserialize
* Signature: ([B)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeDeserialize(
JNIEnv *env, jclass o, jbyteArray data) {
auto buffer = JavaByteArrayToNativeBuffer(env, data);
RAY_CHECK(buffer->Size() > 0);
auto binary = std::string(reinterpret_cast<char *>(buffer->Data()), buffer->Size());
return reinterpret_cast<jlong>(
new ray::ActorHandle(ray::ActorHandle::Deserialize(binary)));
}
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeFree
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeFree(
JNIEnv *env, jclass o, jlong nativeActorHandle) {
delete &GetActorHandle(nativeActorHandle);
}
#ifdef __cplusplus
}
#endif
@@ -0,0 +1,80 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_actor_NativeRayActor */
#ifndef _Included_org_ray_runtime_actor_NativeRayActor
#define _Included_org_ray_runtime_actor_NativeRayActor
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeFork
* Signature: (J)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeFork(JNIEnv *,
jclass,
jlong);
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetActorId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeGetActorId(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetActorHandleId
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeGetActorHandleId(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetLanguage
* Signature: (J)I
*/
JNIEXPORT jint JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeGetLanguage(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeGetActorCreationTaskFunctionDescriptor
* Signature: (J)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeGetActorCreationTaskFunctionDescriptor(
JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeSerialize
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeSerialize(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeDeserialize
* Signature: ([B)J
*/
JNIEXPORT jlong JNICALL
Java_org_ray_runtime_actor_NativeRayActor_nativeDeserialize(JNIEnv *, jclass, jbyteArray);
/*
* Class: org_ray_runtime_actor_NativeRayActor
* Method: nativeFree
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_actor_NativeRayActor_nativeFree(JNIEnv *,
jclass,
jlong);
#ifdef __cplusplus
}
#endif
#endif
@@ -0,0 +1,83 @@
#include "ray/core_worker/lib/java/org_ray_runtime_context_NativeWorkerContext.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/context.h"
#include "ray/core_worker/core_worker.h"
#include "ray/core_worker/lib/java/jni_utils.h"
inline ray::WorkerContext &GetWorkerContextFromPointer(jlong nativeCoreWorkerPointer) {
return reinterpret_cast<ray::CoreWorker *>(nativeCoreWorkerPointer)->GetWorkerContext();
}
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentTaskType
* Signature: (J)I
*/
JNIEXPORT jint JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentTaskType(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer) {
auto task_spec = GetWorkerContextFromPointer(nativeCoreWorkerPointer).GetCurrentTask();
RAY_CHECK(task_spec) << "Current task is not set.";
return static_cast<int>(task_spec->GetMessage().type());
}
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentTaskId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentTaskId(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer) {
const ray::TaskID &task_id =
GetWorkerContextFromPointer(nativeCoreWorkerPointer).GetCurrentTaskID();
return IdToJavaByteBuffer<ray::TaskID>(env, task_id);
}
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentJobId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentJobId(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer) {
const auto &job_id =
GetWorkerContextFromPointer(nativeCoreWorkerPointer).GetCurrentJobID();
return IdToJavaByteBuffer<ray::JobID>(env, job_id);
}
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentWorkerId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentWorkerId(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer) {
const auto &worker_id =
GetWorkerContextFromPointer(nativeCoreWorkerPointer).GetWorkerID();
return IdToJavaByteBuffer<ray::WorkerID>(env, worker_id);
}
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentActorId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentActorId(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer) {
const auto &actor_id =
GetWorkerContextFromPointer(nativeCoreWorkerPointer).GetCurrentActorID();
return IdToJavaByteBuffer<ray::ActorID>(env, actor_id);
}
#ifdef __cplusplus
}
#endif
@@ -0,0 +1,58 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_context_NativeWorkerContext */
#ifndef _Included_org_ray_runtime_context_NativeWorkerContext
#define _Included_org_ray_runtime_context_NativeWorkerContext
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentTaskType
* Signature: (J)I
*/
JNIEXPORT jint JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentTaskType(JNIEnv *,
jclass, jlong);
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentTaskId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentTaskId(JNIEnv *, jclass,
jlong);
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentJobId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentJobId(JNIEnv *, jclass,
jlong);
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentWorkerId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentWorkerId(JNIEnv *,
jclass, jlong);
/*
* Class: org_ray_runtime_context_NativeWorkerContext
* Method: nativeGetCurrentActorId
* Signature: (J)Ljava/nio/ByteBuffer;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_context_NativeWorkerContext_nativeGetCurrentActorId(JNIEnv *, jclass,
jlong);
#ifdef __cplusplus
}
#endif
#endif
@@ -0,0 +1,119 @@
#include "ray/core_worker/lib/java/org_ray_runtime_object_NativeObjectStore.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/common.h"
#include "ray/core_worker/core_worker.h"
#include "ray/core_worker/lib/java/jni_utils.h"
#include "ray/core_worker/object_interface.h"
inline ray::CoreWorkerObjectInterface &GetObjectInterfaceFromPointer(
jlong nativeCoreWorkerPointer) {
return reinterpret_cast<ray::CoreWorker *>(nativeCoreWorkerPointer)->Objects();
}
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativePut
* Signature: (JLorg/ray/runtime/object/NativeRayObject;)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_object_NativeObjectStore_nativePut__JLorg_ray_runtime_object_NativeRayObject_2(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jobject obj) {
auto ray_object = JavaNativeRayObjectToNativeRayObject(env, obj);
RAY_CHECK(ray_object != nullptr);
ray::ObjectID object_id;
auto status =
GetObjectInterfaceFromPointer(nativeCoreWorkerPointer).Put(*ray_object, &object_id);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return IdToJavaByteArray<ray::ObjectID>(env, object_id);
}
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativePut
* Signature: (J[BLorg/ray/runtime/object/NativeRayObject;)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_object_NativeObjectStore_nativePut__J_3BLorg_ray_runtime_object_NativeRayObject_2(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jbyteArray objectId,
jobject obj) {
auto object_id = JavaByteArrayToId<ray::ObjectID>(env, objectId);
auto ray_object = JavaNativeRayObjectToNativeRayObject(env, obj);
RAY_CHECK(ray_object != nullptr);
auto status =
GetObjectInterfaceFromPointer(nativeCoreWorkerPointer).Put(*ray_object, object_id);
if (status.IsIOError() &&
status.message() == "object already exists in the plasma store") {
// Ignore duplicated put on the same object ID.
return;
}
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativeGet
* Signature: (JLjava/util/List;J)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_object_NativeObjectStore_nativeGet(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jobject ids, jlong timeoutMs) {
std::vector<ray::ObjectID> object_ids;
JavaListToNativeVector<ray::ObjectID>(
env, ids, &object_ids, [](JNIEnv *env, jobject id) {
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
});
std::vector<std::shared_ptr<ray::RayObject>> results;
auto status = GetObjectInterfaceFromPointer(nativeCoreWorkerPointer)
.Get(object_ids, (int64_t)timeoutMs, &results);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return NativeVectorToJavaList<std::shared_ptr<ray::RayObject>>(
env, results, NativeRayObjectToJavaNativeRayObject);
}
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativeWait
* Signature: (JLjava/util/List;IJ)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_object_NativeObjectStore_nativeWait(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jobject objectIds,
jint numObjects, jlong timeoutMs) {
std::vector<ray::ObjectID> object_ids;
JavaListToNativeVector<ray::ObjectID>(
env, objectIds, &object_ids, [](JNIEnv *env, jobject id) {
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
});
std::vector<bool> results;
auto status = GetObjectInterfaceFromPointer(nativeCoreWorkerPointer)
.Wait(object_ids, (int)numObjects, (int64_t)timeoutMs, &results);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return NativeVectorToJavaList<bool>(env, results, [](JNIEnv *env, const bool &item) {
return env->NewObject(java_boolean_class, java_boolean_init, (jboolean)item);
});
}
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativeDelete
* Signature: (JLjava/util/List;ZZ)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_object_NativeObjectStore_nativeDelete(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jobject objectIds,
jboolean localOnly, jboolean deleteCreatingTasks) {
std::vector<ray::ObjectID> object_ids;
JavaListToNativeVector<ray::ObjectID>(
env, objectIds, &object_ids, [](JNIEnv *env, jobject id) {
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
});
auto status = GetObjectInterfaceFromPointer(nativeCoreWorkerPointer)
.Delete(object_ids, (bool)localOnly, (bool)deleteCreatingTasks);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
#ifdef __cplusplus
}
#endif
@@ -0,0 +1,55 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_object_NativeObjectStore */
#ifndef _Included_org_ray_runtime_object_NativeObjectStore
#define _Included_org_ray_runtime_object_NativeObjectStore
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativePut
* Signature: (JLorg/ray/runtime/object/NativeRayObject;)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_object_NativeObjectStore_nativePut__JLorg_ray_runtime_object_NativeRayObject_2(
JNIEnv *, jclass, jlong, jobject);
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativePut
* Signature: (J[BLorg/ray/runtime/object/NativeRayObject;)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_object_NativeObjectStore_nativePut__J_3BLorg_ray_runtime_object_NativeRayObject_2(
JNIEnv *, jclass, jlong, jbyteArray, jobject);
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativeGet
* Signature: (JLjava/util/List;J)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_object_NativeObjectStore_nativeGet(
JNIEnv *, jclass, jlong, jobject, jlong);
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativeWait
* Signature: (JLjava/util/List;IJ)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_object_NativeObjectStore_nativeWait(
JNIEnv *, jclass, jlong, jobject, jint, jlong);
/*
* Class: org_ray_runtime_object_NativeObjectStore
* Method: nativeDelete
* Signature: (JLjava/util/List;ZZ)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_object_NativeObjectStore_nativeDelete(
JNIEnv *, jclass, jlong, jobject, jboolean, jboolean);
#ifdef __cplusplus
}
#endif
#endif
@@ -1,151 +0,0 @@
#include "ray/core_worker/lib/java/org_ray_runtime_objectstore_ObjectInterfaceImpl.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/common.h"
#include "ray/core_worker/lib/java/jni_utils.h"
#include "ray/core_worker/object_interface.h"
using ray::rpc::RayletClient;
inline ray::CoreWorkerObjectInterface *GetObjectInterfaceFromPointer(
jlong nativeObjectInterfacePointer) {
return reinterpret_cast<ray::CoreWorkerObjectInterface *>(nativeObjectInterfacePointer);
}
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeCreateObjectInterface
* Signature: (JJLjava/lang/String;)J
*/
JNIEXPORT jlong JNICALL
Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeCreateObjectInterface(
JNIEnv *env, jclass, jlong nativeWorkerContext, jlong nativeRayletClient,
jstring storeSocketName) {
return reinterpret_cast<jlong>(new ray::CoreWorkerObjectInterface(
*reinterpret_cast<ray::WorkerContext *>(nativeWorkerContext),
*reinterpret_cast<std::unique_ptr<RayletClient> *>(nativeRayletClient),
JavaStringToNativeString(env, storeSocketName)));
}
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativePut
* Signature: (JLorg/ray/runtime/objectstore/NativeRayObject;)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativePut__JLorg_ray_runtime_objectstore_NativeRayObject_2(
JNIEnv *env, jclass, jlong nativeObjectInterfacePointer, jobject obj) {
ray::Status status;
ray::ObjectID object_id = ReadJavaNativeRayObject<ray::ObjectID>(
env, obj,
[nativeObjectInterfacePointer,
&status](const std::shared_ptr<ray::RayObject> &rayObject) {
RAY_CHECK(rayObject != nullptr);
ray::ObjectID object_id;
status = GetObjectInterfaceFromPointer(nativeObjectInterfacePointer)
->Put(*rayObject, &object_id);
return object_id;
});
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return IdToJavaByteArray<ray::ObjectID>(env, object_id);
}
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativePut
* Signature: (J[BLorg/ray/runtime/objectstore/NativeRayObject;)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativePut__J_3BLorg_ray_runtime_objectstore_NativeRayObject_2(
JNIEnv *env, jclass, jlong nativeObjectInterfacePointer, jbyteArray objectId,
jobject obj) {
auto object_id = JavaByteArrayToId<ray::ObjectID>(env, objectId);
auto status = ReadJavaNativeRayObject<ray::Status>(
env, obj,
[nativeObjectInterfacePointer,
&object_id](const std::shared_ptr<ray::RayObject> &rayObject) {
RAY_CHECK(rayObject != nullptr);
return GetObjectInterfaceFromPointer(nativeObjectInterfacePointer)
->Put(*rayObject, object_id);
});
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeGet
* Signature: (JLjava/util/List;J)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeGet(
JNIEnv *env, jclass, jlong nativeObjectInterfacePointer, jobject ids,
jlong timeoutMs) {
std::vector<ray::ObjectID> object_ids;
JavaListToNativeVector<ray::ObjectID>(
env, ids, &object_ids, [](JNIEnv *env, jobject id) {
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
});
std::vector<std::shared_ptr<ray::RayObject>> results;
auto status = GetObjectInterfaceFromPointer(nativeObjectInterfacePointer)
->Get(object_ids, (int64_t)timeoutMs, &results);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return NativeVectorToJavaList<std::shared_ptr<ray::RayObject>>(env, results,
ToJavaNativeRayObject);
}
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeWait
* Signature: (JLjava/util/List;IJ)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeWait(
JNIEnv *env, jclass, jlong nativeObjectInterfacePointer, jobject objectIds,
jint numObjects, jlong timeoutMs) {
std::vector<ray::ObjectID> object_ids;
JavaListToNativeVector<ray::ObjectID>(
env, objectIds, &object_ids, [](JNIEnv *env, jobject id) {
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
});
std::vector<bool> results;
auto status = GetObjectInterfaceFromPointer(nativeObjectInterfacePointer)
->Wait(object_ids, (int)numObjects, (int64_t)timeoutMs, &results);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return NativeVectorToJavaList<bool>(env, results, [](JNIEnv *env, const bool &item) {
return env->NewObject(java_boolean_class, java_boolean_init, (jboolean)item);
});
}
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeDelete
* Signature: (JLjava/util/List;ZZ)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeDelete(
JNIEnv *env, jclass, jlong nativeObjectInterfacePointer, jobject objectIds,
jboolean localOnly, jboolean deleteCreatingTasks) {
std::vector<ray::ObjectID> object_ids;
JavaListToNativeVector<ray::ObjectID>(
env, objectIds, &object_ids, [](JNIEnv *env, jobject id) {
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
});
auto status = GetObjectInterfaceFromPointer(nativeObjectInterfacePointer)
->Delete(object_ids, (bool)localOnly, (bool)deleteCreatingTasks);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeDestroy
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeDestroy(
JNIEnv *env, jclass, jlong nativeObjectInterfacePointer) {
delete GetObjectInterfaceFromPointer(nativeObjectInterfacePointer);
}
#ifdef __cplusplus
}
#endif
@@ -1,72 +0,0 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_objectstore_ObjectInterfaceImpl */
#ifndef _Included_org_ray_runtime_objectstore_ObjectInterfaceImpl
#define _Included_org_ray_runtime_objectstore_ObjectInterfaceImpl
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeCreateObjectInterface
* Signature: (JJLjava/lang/String;)J
*/
JNIEXPORT jlong JNICALL
Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeCreateObjectInterface(
JNIEnv *, jclass, jlong, jlong, jstring);
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativePut
* Signature: (JLorg/ray/runtime/objectstore/NativeRayObject;)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativePut__JLorg_ray_runtime_objectstore_NativeRayObject_2(
JNIEnv *, jclass, jlong, jobject);
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativePut
* Signature: (J[BLorg/ray/runtime/objectstore/NativeRayObject;)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativePut__J_3BLorg_ray_runtime_objectstore_NativeRayObject_2(
JNIEnv *, jclass, jlong, jbyteArray, jobject);
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeGet
* Signature: (JLjava/util/List;J)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeGet(
JNIEnv *, jclass, jlong, jobject, jlong);
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeWait
* Signature: (JLjava/util/List;IJ)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeWait(
JNIEnv *, jclass, jlong, jobject, jint, jlong);
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeDelete
* Signature: (JLjava/util/List;ZZ)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeDelete(
JNIEnv *, jclass, jlong, jobject, jboolean, jboolean);
/*
* Class: org_ray_runtime_objectstore_ObjectInterfaceImpl
* Method: nativeDestroy
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_objectstore_ObjectInterfaceImpl_nativeDestroy(
JNIEnv *, jclass, jlong);
#ifdef __cplusplus
}
#endif
#endif
@@ -0,0 +1,74 @@
#include "ray/core_worker/lib/java/org_ray_runtime_raylet_NativeRayletClient.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/common.h"
#include "ray/core_worker/core_worker.h"
#include "ray/core_worker/lib/java/jni_utils.h"
#include "ray/rpc/raylet/raylet_client.h"
inline ray::RayletClient &GetRayletClientFromPointer(jlong nativeCoreWorkerPointer) {
return reinterpret_cast<ray::CoreWorker *>(nativeCoreWorkerPointer)->GetRayletClient();
}
#ifdef __cplusplus
extern "C" {
#endif
using ray::ClientID;
/*
* Class: org_ray_runtime_raylet_NativeRayletClient
* Method: nativePrepareCheckpoint
* Signature: (J[B)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_NativeRayletClient_nativePrepareCheckpoint(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jbyteArray actorId) {
const auto actor_id = JavaByteArrayToId<ActorID>(env, actorId);
ActorCheckpointID checkpoint_id;
auto status = GetRayletClientFromPointer(nativeCoreWorkerPointer)
.PrepareActorCheckpoint(actor_id, checkpoint_id);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
jbyteArray result = env->NewByteArray(checkpoint_id.Size());
env->SetByteArrayRegion(result, 0, checkpoint_id.Size(),
reinterpret_cast<const jbyte *>(checkpoint_id.Data()));
return result;
}
/*
* Class: org_ray_runtime_raylet_NativeRayletClient
* Method: nativeNotifyActorResumedFromCheckpoint
* Signature: (J[B[B)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_raylet_NativeRayletClient_nativeNotifyActorResumedFromCheckpoint(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jbyteArray actorId,
jbyteArray checkpointId) {
const auto actor_id = JavaByteArrayToId<ActorID>(env, actorId);
const auto checkpoint_id = JavaByteArrayToId<ActorCheckpointID>(env, checkpointId);
auto status = GetRayletClientFromPointer(nativeCoreWorkerPointer)
.NotifyActorResumedFromCheckpoint(actor_id, checkpoint_id);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
/*
* Class: org_ray_runtime_raylet_NativeRayletClient
* Method: nativeSetResource
* Signature: (JLjava/lang/String;D[B)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_raylet_NativeRayletClient_nativeSetResource(
JNIEnv *env, jclass, jlong nativeCoreWorkerPointer, jstring resourceName,
jdouble capacity, jbyteArray nodeId) {
const auto node_id = JavaByteArrayToId<ClientID>(env, nodeId);
const char *native_resource_name = env->GetStringUTFChars(resourceName, JNI_FALSE);
auto status =
GetRayletClientFromPointer(nativeCoreWorkerPointer)
.SetResource(native_resource_name, static_cast<double>(capacity), node_id);
env->ReleaseStringUTFChars(resourceName, native_resource_name);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
#ifdef __cplusplus
}
#endif
@@ -0,0 +1,39 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_raylet_NativeRayletClient */
#ifndef _Included_org_ray_runtime_raylet_NativeRayletClient
#define _Included_org_ray_runtime_raylet_NativeRayletClient
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_raylet_NativeRayletClient
* Method: nativePrepareCheckpoint
* Signature: (J[B)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_NativeRayletClient_nativePrepareCheckpoint(JNIEnv *, jclass,
jlong, jbyteArray);
/*
* Class: org_ray_runtime_raylet_NativeRayletClient
* Method: nativeNotifyActorResumedFromCheckpoint
* Signature: (J[B[B)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_raylet_NativeRayletClient_nativeNotifyActorResumedFromCheckpoint(
JNIEnv *, jclass, jlong, jbyteArray, jbyteArray);
/*
* Class: org_ray_runtime_raylet_NativeRayletClient
* Method: nativeSetResource
* Signature: (JLjava/lang/String;D[B)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_raylet_NativeRayletClient_nativeSetResource(
JNIEnv *, jclass, jlong, jstring, jdouble, jbyteArray);
#ifdef __cplusplus
}
#endif
#endif
@@ -0,0 +1,173 @@
#include "ray/core_worker/lib/java/org_ray_runtime_task_NativeTaskSubmitter.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/common.h"
#include "ray/core_worker/core_worker.h"
#include "ray/core_worker/lib/java/jni_utils.h"
#include "ray/core_worker/task_interface.h"
inline ray::CoreWorkerTaskInterface &GetTaskInterfaceFromPointer(
jlong nativeCoreWorkerPointer) {
return reinterpret_cast<ray::CoreWorker *>(nativeCoreWorkerPointer)->Tasks();
}
inline ray::RayFunction ToRayFunction(JNIEnv *env, jobject functionDescriptor) {
std::vector<std::string> function_descriptor;
JavaStringListToNativeStringVector(
env, env->CallObjectMethod(functionDescriptor, java_function_descriptor_to_list),
&function_descriptor);
jobject java_language =
env->CallObjectMethod(functionDescriptor, java_function_descriptor_get_language);
int language = env->CallIntMethod(java_language, java_language_get_number);
ray::RayFunction ray_function{static_cast<::Language>(language), function_descriptor};
return ray_function;
}
inline std::vector<ray::TaskArg> ToTaskArgs(JNIEnv *env, jobject args) {
std::vector<ray::TaskArg> task_args;
JavaListToNativeVector<ray::TaskArg>(
env, args, &task_args, [](JNIEnv *env, jobject arg) {
auto java_id = env->GetObjectField(arg, java_function_arg_id);
if (java_id) {
auto java_id_bytes = static_cast<jbyteArray>(
env->CallObjectMethod(java_id, java_base_id_get_bytes));
return ray::TaskArg::PassByReference(
JavaByteArrayToId<ray::ObjectID>(env, java_id_bytes));
}
auto java_data =
static_cast<jbyteArray>(env->GetObjectField(arg, java_function_arg_data));
RAY_CHECK(java_data) << "Both id and data of FunctionArg are null.";
return ray::TaskArg::PassByValue(JavaByteArrayToNativeBuffer(env, java_data));
});
return task_args;
}
inline std::unordered_map<std::string, double> ToResources(JNIEnv *env,
jobject java_resources) {
std::unordered_map<std::string, double> resources;
if (java_resources) {
jobject entry_set = env->CallObjectMethod(java_resources, java_map_entry_set);
jobject iterator = env->CallObjectMethod(entry_set, java_set_iterator);
while (env->CallBooleanMethod(iterator, java_iterator_has_next)) {
jobject map_entry = env->CallObjectMethod(iterator, java_iterator_next);
std::string key = JavaStringToNativeString(
env, (jstring)env->CallObjectMethod(map_entry, java_map_entry_get_key));
double value = env->CallDoubleMethod(
env->CallObjectMethod(map_entry, java_map_entry_get_value),
java_double_double_value);
resources.emplace(key, value);
}
}
return resources;
}
inline ray::TaskOptions ToTaskOptions(JNIEnv *env, jint numReturns, jobject callOptions) {
std::unordered_map<std::string, double> resources;
if (callOptions) {
jobject java_resources =
env->GetObjectField(callOptions, java_base_task_options_resources);
resources = ToResources(env, java_resources);
}
ray::TaskOptions task_options{numReturns, resources};
return task_options;
}
inline ray::ActorCreationOptions ToActorCreationOptions(JNIEnv *env,
jobject actorCreationOptions) {
uint64_t max_reconstructions = 0;
std::unordered_map<std::string, double> resources;
std::vector<std::string> dynamic_worker_options;
if (actorCreationOptions) {
max_reconstructions = static_cast<uint64_t>(env->GetIntField(
actorCreationOptions, java_actor_creation_options_max_reconstructions));
jobject java_resources =
env->GetObjectField(actorCreationOptions, java_base_task_options_resources);
resources = ToResources(env, java_resources);
std::string jvm_options = JavaStringToNativeString(
env, (jstring)env->GetObjectField(actorCreationOptions,
java_actor_creation_options_jvm_options));
dynamic_worker_options.emplace_back(jvm_options);
}
ray::ActorCreationOptions action_creation_options{
static_cast<uint64_t>(max_reconstructions), false, resources,
dynamic_worker_options};
return action_creation_options;
}
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_task_NativeTaskSubmitter
* Method: nativeSubmitTask
* Signature:
* (JLorg/ray/runtime/functionmanager/FunctionDescriptor;Ljava/util/List;ILorg/ray/api/options/CallOptions;)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_task_NativeTaskSubmitter_nativeSubmitTask(
JNIEnv *env, jclass p, jlong nativeCoreWorkerPointer, jobject functionDescriptor,
jobject args, jint numReturns, jobject callOptions) {
auto ray_function = ToRayFunction(env, functionDescriptor);
auto task_args = ToTaskArgs(env, args);
auto task_options = ToTaskOptions(env, numReturns, callOptions);
std::vector<ObjectID> return_ids;
auto status = GetTaskInterfaceFromPointer(nativeCoreWorkerPointer)
.SubmitTask(ray_function, task_args, task_options, &return_ids);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return NativeIdVectorToJavaByteArrayList(env, return_ids);
}
/*
* Class: org_ray_runtime_task_NativeTaskSubmitter
* Method: nativeCreateActor
* Signature:
* (JLorg/ray/runtime/functionmanager/FunctionDescriptor;Ljava/util/List;Lorg/ray/api/options/ActorCreationOptions;)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_task_NativeTaskSubmitter_nativeCreateActor(
JNIEnv *env, jclass p, jlong nativeCoreWorkerPointer, jobject functionDescriptor,
jobject args, jobject actorCreationOptions) {
auto ray_function = ToRayFunction(env, functionDescriptor);
auto task_args = ToTaskArgs(env, args);
auto actor_creation_options = ToActorCreationOptions(env, actorCreationOptions);
std::unique_ptr<ray::ActorHandle> actor_handle;
auto status =
GetTaskInterfaceFromPointer(nativeCoreWorkerPointer)
.CreateActor(ray_function, task_args, actor_creation_options, &actor_handle);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, 0);
return reinterpret_cast<jlong>(actor_handle.release());
}
/*
* Class: org_ray_runtime_task_NativeTaskSubmitter
* Method: nativeSubmitActorTask
* Signature:
* (JJLorg/ray/runtime/functionmanager/FunctionDescriptor;Ljava/util/List;ILorg/ray/api/options/CallOptions;)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_task_NativeTaskSubmitter_nativeSubmitActorTask(
JNIEnv *env, jclass p, jlong nativeCoreWorkerPointer, jlong nativeActorHandle,
jobject functionDescriptor, jobject args, jint numReturns, jobject callOptions) {
auto &actor_handle = *(reinterpret_cast<ray::ActorHandle *>(nativeActorHandle));
auto ray_function = ToRayFunction(env, functionDescriptor);
auto task_args = ToTaskArgs(env, args);
auto task_options = ToTaskOptions(env, numReturns, callOptions);
std::vector<ObjectID> return_ids;
auto status = GetTaskInterfaceFromPointer(nativeCoreWorkerPointer)
.SubmitActorTask(actor_handle, ray_function, task_args, task_options,
&return_ids);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
return NativeIdVectorToJavaByteArrayList(env, return_ids);
}
#ifdef __cplusplus
}
#endif
@@ -0,0 +1,43 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_task_NativeTaskSubmitter */
#ifndef _Included_org_ray_runtime_task_NativeTaskSubmitter
#define _Included_org_ray_runtime_task_NativeTaskSubmitter
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_task_NativeTaskSubmitter
* Method: nativeSubmitTask
* Signature:
* (JLorg/ray/runtime/functionmanager/FunctionDescriptor;Ljava/util/List;ILorg/ray/api/options/CallOptions;)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL Java_org_ray_runtime_task_NativeTaskSubmitter_nativeSubmitTask(
JNIEnv *, jclass, jlong, jobject, jobject, jint, jobject);
/*
* Class: org_ray_runtime_task_NativeTaskSubmitter
* Method: nativeCreateActor
* Signature:
* (JLorg/ray/runtime/functionmanager/FunctionDescriptor;Ljava/util/List;Lorg/ray/api/options/ActorCreationOptions;)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_task_NativeTaskSubmitter_nativeCreateActor(
JNIEnv *, jclass, jlong, jobject, jobject, jobject);
/*
* Class: org_ray_runtime_task_NativeTaskSubmitter
* Method: nativeSubmitActorTask
* Signature:
* (JJLorg/ray/runtime/functionmanager/FunctionDescriptor;Ljava/util/List;ILorg/ray/api/options/CallOptions;)Ljava/util/List;
*/
JNIEXPORT jobject JNICALL
Java_org_ray_runtime_task_NativeTaskSubmitter_nativeSubmitActorTask(JNIEnv *, jclass,
jlong, jlong, jobject,
jobject, jint,
jobject);
#ifdef __cplusplus
}
#endif
#endif
+15 -18
View File
@@ -13,7 +13,8 @@ CoreWorkerTaskExecutionInterface::CoreWorkerTaskExecutionInterface(
object_interface_(object_interface),
execution_callback_(executor),
worker_server_("Worker", 0 /* let grpc choose port */),
main_work_(main_service_) {
main_service_(std::make_shared<boost::asio::io_service>()),
main_work_(*main_service_) {
RAY_CHECK(execution_callback_ != nullptr);
auto func = std::bind(&CoreWorkerTaskExecutionInterface::ExecuteTask, this,
@@ -21,11 +22,11 @@ CoreWorkerTaskExecutionInterface::CoreWorkerTaskExecutionInterface(
task_receivers_.emplace(
TaskTransportType::RAYLET,
std::unique_ptr<CoreWorkerRayletTaskReceiver>(new CoreWorkerRayletTaskReceiver(
raylet_client, object_interface_, main_service_, worker_server_, func)));
raylet_client, object_interface_, *main_service_, worker_server_, func)));
task_receivers_.emplace(
TaskTransportType::DIRECT_ACTOR,
std::unique_ptr<CoreWorkerDirectActorTaskReceiver>(
new CoreWorkerDirectActorTaskReceiver(object_interface_, main_service_,
new CoreWorkerDirectActorTaskReceiver(object_interface_, *main_service_,
worker_server_, func)));
// Start RPC server after all the task receivers are properly initialized.
@@ -35,6 +36,8 @@ CoreWorkerTaskExecutionInterface::CoreWorkerTaskExecutionInterface(
Status CoreWorkerTaskExecutionInterface::ExecuteTask(
const TaskSpecification &task_spec,
std::vector<std::shared_ptr<RayObject>> *results) {
RAY_LOG(DEBUG) << "Executing task " << task_spec.TaskId();
worker_context_.SetCurrentTask(task_spec);
RayFunction func{task_spec.GetLanguage(), task_spec.FunctionDescriptor()};
@@ -42,17 +45,6 @@ Status CoreWorkerTaskExecutionInterface::ExecuteTask(
std::vector<std::shared_ptr<RayObject>> args;
RAY_CHECK_OK(BuildArgsForExecutor(task_spec, &args));
TaskType task_type;
if (task_spec.IsActorCreationTask()) {
task_type = TaskType::ACTOR_CREATION_TASK;
} else if (task_spec.IsActorTask()) {
task_type = TaskType::ACTOR_TASK;
} else {
task_type = TaskType::NORMAL_TASK;
}
TaskInfo task_info{task_spec.TaskId(), task_spec.JobId(), task_type};
auto num_returns = task_spec.NumReturns();
if (task_spec.IsActorCreationTask() || task_spec.IsActorTask()) {
RAY_CHECK(num_returns > 0);
@@ -60,7 +52,7 @@ Status CoreWorkerTaskExecutionInterface::ExecuteTask(
num_returns--;
}
auto status = execution_callback_(func, args, task_info, num_returns, results);
auto status = execution_callback_(func, args, num_returns, results);
// TODO(zhijunfu):
// 1. Check and handle failure.
// 2. Save or load checkpoint.
@@ -69,10 +61,15 @@ Status CoreWorkerTaskExecutionInterface::ExecuteTask(
void CoreWorkerTaskExecutionInterface::Run() {
// Run main IO service.
main_service_.run();
main_service_->run();
}
// should never reach here.
RAY_LOG(FATAL) << "should never reach here after running main io service";
void CoreWorkerTaskExecutionInterface::Stop() {
// Stop main IO service.
std::shared_ptr<boost::asio::io_service> main_service = main_service_;
// Delay the execution of io_service::stop() to avoid deadlock if
// CoreWorkerTaskExecutionInterface::Stop is called inside a task.
main_service_->post([main_service]() { main_service->stop(); });
}
Status CoreWorkerTaskExecutionInterface::BuildArgsForExecutor(
+8 -5
View File
@@ -27,23 +27,26 @@ class CoreWorkerTaskExecutionInterface {
///
/// \param ray_function[in] Information about the function to execute.
/// \param args[in] Arguments of the task.
/// \param task_info[in] Information of the task to execute.
/// \param results[out] Results of the task execution.
/// \return Status.
using TaskExecutor = std::function<Status(
const RayFunction &ray_function,
const std::vector<std::shared_ptr<RayObject>> &args, const TaskInfo &task_info,
int num_returns, std::vector<std::shared_ptr<RayObject>> *results)>;
const std::vector<std::shared_ptr<RayObject>> &args, int num_returns,
std::vector<std::shared_ptr<RayObject>> *results)>;
CoreWorkerTaskExecutionInterface(WorkerContext &worker_context,
std::unique_ptr<RayletClient> &raylet_client,
CoreWorkerObjectInterface &object_interface,
const TaskExecutor &executor);
/// Start receving and executes tasks in a infinite loop.
/// Start receiving and executing tasks.
/// \return void.
void Run();
/// Stop receiving and executing tasks.
/// \return void.
void Stop();
private:
/// Build arguments for task executor. This would loop through all the arguments
/// in task spec, and for each of them that's passed by reference (ObjectID),
@@ -80,7 +83,7 @@ class CoreWorkerTaskExecutionInterface {
rpc::GrpcServer worker_server_;
/// Event loop where tasks are processed.
boost::asio::io_service main_service_;
std::shared_ptr<boost::asio::io_service> main_service_;
/// The asio work to keep main_service_ alive.
boost::asio::io_service::work main_work_;
+1 -1
View File
@@ -171,7 +171,7 @@ Status CoreWorkerTaskInterface::CreateActor(
actor_creation_options.resources, actor_creation_options.resources,
TaskTransportType::RAYLET, &return_ids);
builder.SetActorCreationTaskSpec(actor_id, actor_creation_options.max_reconstructions,
{});
actor_creation_options.dynamic_worker_options);
*actor_handle = std::unique_ptr<ActorHandle>(new ActorHandle(
actor_id, ActorHandleID::Nil(), function.language,
+7 -2
View File
@@ -37,10 +37,12 @@ struct TaskOptions {
struct ActorCreationOptions {
ActorCreationOptions() {}
ActorCreationOptions(uint64_t max_reconstructions, bool is_direct_call,
const std::unordered_map<std::string, double> &resources)
const std::unordered_map<std::string, double> &resources,
const std::vector<std::string> &dynamic_worker_options)
: max_reconstructions(max_reconstructions),
is_direct_call(is_direct_call),
resources(resources) {}
resources(resources),
dynamic_worker_options(dynamic_worker_options) {}
/// Maximum number of times that the actor should be reconstructed when it dies
/// unexpectedly. It must be non-negative. If it's 0, the actor won't be reconstructed.
@@ -50,6 +52,9 @@ struct ActorCreationOptions {
const bool is_direct_call = false;
/// Resources required by the whole lifetime of this actor.
const std::unordered_map<std::string, double> resources;
/// The dynamic options used in the worker command when starting a worker process for
/// an actor creation task.
const std::vector<std::string> dynamic_worker_options;
};
/// A handle to an actor.
+3 -3
View File
@@ -62,7 +62,7 @@ std::unique_ptr<ActorHandle> CreateActorHelper(
std::vector<TaskArg> args;
args.emplace_back(TaskArg::PassByValue(buffer));
ActorCreationOptions actor_options{max_reconstructions, is_direct_call, resources};
ActorCreationOptions actor_options{max_reconstructions, is_direct_call, resources, {}};
// Create an actor.
RAY_CHECK_OK(worker.Tasks().CreateActor(func, args, actor_options, &actor_handle));
@@ -586,7 +586,7 @@ TEST_F(ZeroNodeTest, TestTaskSpecPerf) {
args.emplace_back(TaskArg::PassByValue(buffer));
std::unordered_map<std::string, double> resources;
ActorCreationOptions actor_options{0, /* is_direct_call */ true, resources};
ActorCreationOptions actor_options{0, /*is_direct_call*/ true, resources, {}};
const auto job_id = NextJobId();
ActorHandle actor_handle(ActorID::Of(job_id, TaskID::ForDriverTask(job_id), 1),
ActorHandleID::Nil(), function.language, true,
@@ -647,7 +647,7 @@ TEST_F(SingleNodeTest, TestDirectActorTaskSubmissionPerf) {
args.emplace_back(TaskArg::PassByValue(buffer));
std::unordered_map<std::string, double> resources;
ActorCreationOptions actor_options{0, /* is_direct_call */ true, resources};
ActorCreationOptions actor_options{0, /*is_direct_call*/ true, resources, {}};
// Create an actor.
RAY_CHECK_OK(driver.Tasks().CreateActor(func, args, actor_options, &actor_handle));
// wait for actor creation finish.
+2 -3
View File
@@ -24,7 +24,7 @@ class MockWorker {
const gcs::GcsClientOptions &gcs_options)
: worker_(WorkerType::WORKER, Language::PYTHON, store_socket, raylet_socket,
JobID::FromInt(1), gcs_options,
std::bind(&MockWorker::ExecuteTask, this, _1, _2, _3, _4, _5)) {}
std::bind(&MockWorker::ExecuteTask, this, _1, _2, _3, _4)) {}
void Run() {
// Start executing tasks.
@@ -33,8 +33,7 @@ class MockWorker {
private:
Status ExecuteTask(const RayFunction &ray_function,
const std::vector<std::shared_ptr<RayObject>> &args,
const TaskInfo &task_info, int num_returns,
const std::vector<std::shared_ptr<RayObject>> &args, int num_returns,
std::vector<std::shared_ptr<RayObject>> *results) {
// Note that this doesn't include dummy object id.
RAY_CHECK(num_returns >= 0);
@@ -29,6 +29,7 @@ void CoreWorkerRayletTaskReceiver::HandleAssignTask(
rpc::SendReplyCallback send_reply_callback) {
const Task task(request.task());
const auto &task_spec = task.GetTaskSpecification();
RAY_LOG(DEBUG) << "Received task " << task_spec.TaskId();
std::vector<std::shared_ptr<RayObject>> results;
auto status = task_handler_(task_spec, &results);
@@ -39,12 +40,23 @@ void CoreWorkerRayletTaskReceiver::HandleAssignTask(
num_returns--;
}
RAY_LOG(DEBUG) << "Assigned task " << task_spec.TaskId()
<< " finished execution. num_returns: " << num_returns;
RAY_CHECK(results.size() == num_returns);
for (size_t i = 0; i < num_returns; i++) {
ObjectID id = ObjectID::ForTaskReturn(
task_spec.TaskId(), /*index=*/i + 1,
/*transport_type=*/static_cast<int>(TaskTransportType::RAYLET));
RAY_CHECK_OK(object_interface_.Put(*results[i], id));
Status status = object_interface_.Put(*results[i], id);
if (!status.ok()) {
// TODO (kfstorm): RAY_LOG(FATAL) except the error is about the object to put
// already exists.
RAY_LOG(WARNING) << "Task " << task_spec.TaskId() << " failed to put object " << id
<< " in store: " << status.message();
} else {
RAY_LOG(DEBUG) << "Task " << task_spec.TaskId() << " put object " << id
<< " in store.";
}
}
// Notify raylet that current task is done via a `TaskDone` message. This is to
+28
View File
@@ -188,6 +188,17 @@ Status AuthenticateRedis(redisAsyncContext *context, const std::string &password
return Status::OK();
}
void RedisAsyncContextDisconnectCallback(const redisAsyncContext *context, int status) {
RAY_LOG(WARNING) << "Redis async context disconnected. Status: " << status;
reinterpret_cast<RedisContext *>(context->data)
->AsyncDisconnectCallback(context, status);
}
void SetDisconnectCallback(RedisContext *redis_context, redisAsyncContext *context) {
context->data = redis_context;
redisAsyncSetDisconnectCallback(context, RedisAsyncContextDisconnectCallback);
}
template <typename RedisContext, typename RedisConnectFunction>
Status ConnectWithRetries(const std::string &address, int port,
const RedisConnectFunction &connect_function,
@@ -216,6 +227,10 @@ Status ConnectWithRetries(const std::string &address, int port,
Status RedisContext::Connect(const std::string &address, int port, bool sharding,
const std::string &password = "") {
RAY_CHECK(!context_);
RAY_CHECK(!async_context_);
RAY_CHECK(!subscribe_context_);
RAY_CHECK_OK(ConnectWithRetries(address, port, redisConnect, &context_));
RAY_CHECK_OK(AuthenticateRedis(context_, password));
@@ -226,10 +241,12 @@ Status RedisContext::Connect(const std::string &address, int port, bool sharding
// Connect to async context
RAY_CHECK_OK(ConnectWithRetries(address, port, redisAsyncConnect, &async_context_));
SetDisconnectCallback(this, async_context_);
RAY_CHECK_OK(AuthenticateRedis(async_context_, password));
// Connect to subscribe context
RAY_CHECK_OK(ConnectWithRetries(address, port, redisAsyncConnect, &subscribe_context_));
SetDisconnectCallback(this, subscribe_context_);
RAY_CHECK_OK(AuthenticateRedis(subscribe_context_, password));
return Status::OK();
@@ -245,6 +262,7 @@ Status RedisContext::AttachToEventLoop(aeEventLoop *loop) {
}
Status RedisContext::RunArgvAsync(const std::vector<std::string> &args) {
RAY_CHECK(async_context_);
// Build the arguments.
std::vector<const char *> argv;
std::vector<size_t> argc;
@@ -268,6 +286,7 @@ Status RedisContext::SubscribeAsync(const ClientID &client_id,
int64_t *out_callback_index) {
RAY_CHECK(pubsub_channel != TablePubsub::NO_PUBLISH)
<< "Client requested subscribe on a table that does not support pubsub";
RAY_CHECK(subscribe_context_);
int64_t callback_index = RedisCallbackManager::instance().add(redisCallback, true);
RAY_CHECK(out_callback_index != nullptr);
@@ -294,6 +313,15 @@ Status RedisContext::SubscribeAsync(const ClientID &client_id,
return Status::OK();
}
void RedisContext::AsyncDisconnectCallback(const redisAsyncContext *context, int status) {
if (context == async_context_) {
async_context_ = nullptr;
}
if (context == subscribe_context_) {
subscribe_context_ = nullptr;
}
}
} // namespace gcs
} // namespace ray
+20 -3
View File
@@ -149,9 +149,25 @@ class RedisContext {
/// \return Status.
Status SubscribeAsync(const ClientID &client_id, const TablePubsub pubsub_channel,
const RedisCallback &redisCallback, int64_t *out_callback_index);
redisContext *sync_context() { return context_; }
redisAsyncContext *async_context() { return async_context_; }
redisAsyncContext *subscribe_context() { return subscribe_context_; };
/// Called when an instance of redisAsyncContext is disconnected.
///
/// \param context the redisAsyncContext instances
/// \param status The status code of disconnection
void AsyncDisconnectCallback(const redisAsyncContext *context, int status);
redisContext *sync_context() {
RAY_CHECK(context_);
return context_;
}
redisAsyncContext *async_context() {
RAY_CHECK(async_context_);
return async_context_;
}
redisAsyncContext *subscribe_context() {
RAY_CHECK(subscribe_context_);
return subscribe_context_;
};
private:
redisContext *context_;
@@ -164,6 +180,7 @@ Status RedisContext::RunAsync(const std::string &command, const ID &id, const vo
size_t length, const TablePrefix prefix,
const TablePubsub pubsub_channel,
RedisCallback redisCallback, int log_length) {
RAY_CHECK(async_context_);
int64_t callback_index = RedisCallbackManager::instance().add(redisCallback, false);
if (length > 0) {
if (log_length >= 0) {
+1 -1
View File
@@ -95,7 +95,7 @@ message ActorCreationTaskSpec {
// The max number of times this actor should be recontructed.
// If this number of 0 or negative, the actor won't be reconstructed on failure.
uint64 max_actor_reconstructions = 3;
// The dynamic options used in the worker command when starting the worker process for
// The dynamic options used in the worker command when starting a worker process for
// an actor creation task. If the list isn't empty, the options will be used to replace
// the placeholder strings (`RAY_WORKER_OPTION_0`, `RAY_WORKER_OPTION_1`, etc) in the
// worker command.
@@ -1,296 +0,0 @@
#include "ray/raylet/lib/java/org_ray_runtime_raylet_RayletClientImpl.h"
#include <jni.h>
#include "ray/common/id.h"
#include "ray/core_worker/lib/java/jni_utils.h"
#include "ray/rpc/raylet/raylet_client.h"
#include "ray/util/logging.h"
#ifdef __cplusplus
extern "C" {
#endif
using ray::ClientID;
using ray::WorkerID;
using ray::rpc::RayletClient;
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeInit
* Signature: (Ljava/lang/String;[BZ[B)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeInit(
JNIEnv *env, jclass, jstring sockName, jbyteArray workerId, jboolean isWorker,
jbyteArray jobId) {
const auto worker_id = JavaByteArrayToId<WorkerID>(env, workerId);
const auto job_id = JavaByteArrayToId<JobID>(env, jobId);
const char *nativeString = env->GetStringUTFChars(sockName, JNI_FALSE);
auto raylet_client = new std::unique_ptr<RayletClient>(
new RayletClient(nativeString, worker_id, isWorker, job_id, Language::JAVA));
env->ReleaseStringUTFChars(sockName, nativeString);
return reinterpret_cast<jlong>(raylet_client);
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeSubmitTask
* Signature: (J[BLjava/nio/ByteBuffer;II)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeSubmitTask(
JNIEnv *env, jclass, jlong client, jbyteArray taskSpec) {
auto &raylet_client = *reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
jbyte *data = env->GetByteArrayElements(taskSpec, NULL);
jsize size = env->GetArrayLength(taskSpec);
ray::rpc::TaskSpec task_spec_message;
task_spec_message.ParseFromArray(data, size);
env->ReleaseByteArrayElements(taskSpec, data, JNI_ABORT);
ray::TaskSpecification task_spec(task_spec_message);
auto status = raylet_client->SubmitTask(task_spec);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGetTask
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeGetTask(
JNIEnv *env, jclass, jlong client) {
auto &raylet_client = *reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
std::unique_ptr<ray::TaskSpecification> spec;
auto status = raylet_client->GetTask(&spec);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
// Serialize the task spec and copy to Java byte array.
auto task_data = spec->Serialize();
jbyteArray result = env->NewByteArray(task_data.size());
if (result == nullptr) {
return nullptr; /* out of memory error thrown */
}
env->SetByteArrayRegion(result, 0, task_data.size(),
reinterpret_cast<const jbyte *>(task_data.data()));
return result;
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeDestroy
* Signature: (J)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeDestroy(
JNIEnv *env, jclass, jlong client) {
auto raylet_client = reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
auto status = (*raylet_client)->Disconnect();
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
delete raylet_client;
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeWaitObject
* Signature: (J[[BIIZ[B)[Z
*/
JNIEXPORT jbooleanArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeWaitObject(
JNIEnv *env, jclass, jlong client, jobjectArray objectIds, jint numReturns,
jint timeoutMillis, jboolean isWaitLocal, jbyteArray currentTaskId) {
std::vector<ObjectID> object_ids;
auto len = env->GetArrayLength(objectIds);
for (int i = 0; i < len; i++) {
jbyteArray object_id_bytes =
static_cast<jbyteArray>(env->GetObjectArrayElement(objectIds, i));
const auto object_id = JavaByteArrayToId<ObjectID>(env, object_id_bytes);
object_ids.push_back(object_id);
env->DeleteLocalRef(object_id_bytes);
}
const auto current_task_id = JavaByteArrayToId<TaskID>(env, currentTaskId);
auto &raylet_client = *reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
// Invoke wait.
WaitResultPair result;
auto status =
raylet_client->Wait(object_ids, numReturns, timeoutMillis,
static_cast<bool>(isWaitLocal), current_task_id, &result);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
// Convert result to java object.
jboolean put_value = true;
jbooleanArray resultArray = env->NewBooleanArray(object_ids.size());
for (uint i = 0; i < result.first.size(); ++i) {
for (uint j = 0; j < object_ids.size(); ++j) {
if (result.first[i] == object_ids[j]) {
env->SetBooleanArrayRegion(resultArray, j, 1, &put_value);
break;
}
}
}
put_value = false;
for (uint i = 0; i < result.second.size(); ++i) {
for (uint j = 0; j < object_ids.size(); ++j) {
if (result.second[i] == object_ids[j]) {
env->SetBooleanArrayRegion(resultArray, j, 1, &put_value);
break;
}
}
}
return resultArray;
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGenerateActorCreationTaskId
* Signature: ([B[BI)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeGenerateActorCreationTaskId(
JNIEnv *env, jclass, jbyteArray jobId, jbyteArray parentTaskId,
jint parent_task_counter) {
const auto job_id = JavaByteArrayToId<JobID>(env, jobId);
const auto parent_task_id = JavaByteArrayToId<TaskID>(env, parentTaskId);
const ActorID actor_id = ray::ActorID::Of(job_id, parent_task_id, parent_task_counter);
const TaskID actor_creation_task_id = ray::TaskID::ForActorCreationTask(actor_id);
jbyteArray result = env->NewByteArray(actor_creation_task_id.Size());
if (nullptr == result) {
return nullptr;
}
env->SetByteArrayRegion(result, 0, actor_creation_task_id.Size(),
reinterpret_cast<const jbyte *>(actor_creation_task_id.Data()));
return result;
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGenerateActorTaskId
* Signature: ([B[BI[B)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeGenerateActorTaskId(
JNIEnv *env, jclass, jbyteArray jobId, jbyteArray parentTaskId,
jint parent_task_counter, jbyteArray actorId) {
const auto job_id = JavaByteArrayToId<JobID>(env, jobId);
const auto parent_task_id = JavaByteArrayToId<TaskID>(env, parentTaskId);
const auto actor_id = JavaByteArrayToId<ActorID>(env, actorId);
const TaskID actor_task_id =
ray::TaskID::ForActorTask(job_id, parent_task_id, parent_task_counter, actor_id);
jbyteArray result = env->NewByteArray(actor_task_id.Size());
if (nullptr == result) {
return nullptr;
}
env->SetByteArrayRegion(result, 0, actor_task_id.Size(),
reinterpret_cast<const jbyte *>(actor_task_id.Data()));
return result;
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGenerateNormalTaskId
* Signature: ([B[BI)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeGenerateNormalTaskId(
JNIEnv *env, jclass, jbyteArray jobId, jbyteArray parentTaskId,
jint parent_task_counter) {
const auto job_id = JavaByteArrayToId<JobID>(env, jobId);
const auto parent_task_id = JavaByteArrayToId<TaskID>(env, parentTaskId);
const TaskID task_id =
ray::TaskID::ForNormalTask(job_id, parent_task_id, parent_task_counter);
jbyteArray result = env->NewByteArray(task_id.Size());
if (nullptr == result) {
return nullptr;
}
env->SetByteArrayRegion(result, 0, task_id.Size(),
reinterpret_cast<const jbyte *>(task_id.Data()));
return result;
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeFreePlasmaObjects
* Signature: (J[[BZZ)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeFreePlasmaObjects(
JNIEnv *env, jclass, jlong client, jobjectArray objectIds, jboolean localOnly,
jboolean deleteCreatingTasks) {
std::vector<ObjectID> object_ids;
auto len = env->GetArrayLength(objectIds);
for (int i = 0; i < len; i++) {
jbyteArray object_id_bytes =
static_cast<jbyteArray>(env->GetObjectArrayElement(objectIds, i));
const auto object_id = JavaByteArrayToId<ObjectID>(env, object_id_bytes);
object_ids.push_back(object_id);
env->DeleteLocalRef(object_id_bytes);
}
auto &raylet_client = *reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
auto status = raylet_client->FreeObjects(object_ids, localOnly, deleteCreatingTasks);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativePrepareCheckpoint
* Signature: (J[B)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativePrepareCheckpoint(JNIEnv *env, jclass,
jlong client,
jbyteArray actorId) {
auto &raylet_client = *reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
const auto actor_id = JavaByteArrayToId<ActorID>(env, actorId);
ActorCheckpointID checkpoint_id;
auto status = raylet_client->PrepareActorCheckpoint(actor_id, checkpoint_id);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
jbyteArray result = env->NewByteArray(checkpoint_id.Size());
env->SetByteArrayRegion(result, 0, checkpoint_id.Size(),
reinterpret_cast<const jbyte *>(checkpoint_id.Data()));
return result;
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeNotifyActorResumedFromCheckpoint
* Signature: (J[B[B)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeNotifyActorResumedFromCheckpoint(
JNIEnv *env, jclass, jlong client, jbyteArray actorId, jbyteArray checkpointId) {
auto &raylet_client = *reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
const auto actor_id = JavaByteArrayToId<ActorID>(env, actorId);
const auto checkpoint_id = JavaByteArrayToId<ActorCheckpointID>(env, checkpointId);
auto status = raylet_client->NotifyActorResumedFromCheckpoint(actor_id, checkpoint_id);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeSetResource
* Signature: (JLjava/lang/String;D[B)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeSetResource(
JNIEnv *env, jclass, jlong client, jstring resourceName, jdouble capacity,
jbyteArray nodeId) {
auto &raylet_client = *reinterpret_cast<std::unique_ptr<RayletClient> *>(client);
const auto node_id = JavaByteArrayToId<ClientID>(env, nodeId);
const char *native_resource_name = env->GetStringUTFChars(resourceName, JNI_FALSE);
auto status = raylet_client->SetResource(native_resource_name,
static_cast<double>(capacity), node_id);
env->ReleaseStringUTFChars(resourceName, native_resource_name);
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
}
#ifdef __cplusplus
}
#endif
@@ -1,121 +0,0 @@
/* DO NOT EDIT THIS FILE - it is machine generated */
#include <jni.h>
/* Header for class org_ray_runtime_raylet_RayletClientImpl */
#ifndef _Included_org_ray_runtime_raylet_RayletClientImpl
#define _Included_org_ray_runtime_raylet_RayletClientImpl
#ifdef __cplusplus
extern "C" {
#endif
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeInit
* Signature: (Ljava/lang/String;[BZ[B)J
*/
JNIEXPORT jlong JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeInit(
JNIEnv *, jclass, jstring, jbyteArray, jboolean, jbyteArray);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeSubmitTask
* Signature: (J[B)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeSubmitTask(
JNIEnv *, jclass, jlong, jbyteArray);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGetTask
* Signature: (J)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeGetTask(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeDestroy
* Signature: (J)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeDestroy(JNIEnv *, jclass, jlong);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeWaitObject
* Signature: (J[[BIIZ[B)[Z
*/
JNIEXPORT jbooleanArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeWaitObject(JNIEnv *, jclass, jlong,
jobjectArray, jint, jint,
jboolean, jbyteArray);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGenerateActorCreationTaskId
* Signature: ([B[BI)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeGenerateActorCreationTaskId(
JNIEnv *, jclass, jbyteArray, jbyteArray, jint);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGenerateActorTaskId
* Signature: ([B[BI[B)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeGenerateActorTaskId(JNIEnv *, jclass,
jbyteArray,
jbyteArray, jint,
jbyteArray);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeGenerateNormalTaskId
* Signature: ([B[BI)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeGenerateNormalTaskId(JNIEnv *, jclass,
jbyteArray,
jbyteArray, jint);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeFreePlasmaObjects
* Signature: (J[[BZZ)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeFreePlasmaObjects(JNIEnv *, jclass,
jlong, jobjectArray,
jboolean, jboolean);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativePrepareCheckpoint
* Signature: (J[B)[B
*/
JNIEXPORT jbyteArray JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativePrepareCheckpoint(JNIEnv *, jclass,
jlong, jbyteArray);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeNotifyActorResumedFromCheckpoint
* Signature: (J[B[B)V
*/
JNIEXPORT void JNICALL
Java_org_ray_runtime_raylet_RayletClientImpl_nativeNotifyActorResumedFromCheckpoint(
JNIEnv *, jclass, jlong, jbyteArray, jbyteArray);
/*
* Class: org_ray_runtime_raylet_RayletClientImpl
* Method: nativeSetResource
* Signature: (JLjava/lang/String;D[B)V
*/
JNIEXPORT void JNICALL Java_org_ray_runtime_raylet_RayletClientImpl_nativeSetResource(
JNIEnv *, jclass, jlong, jstring, jdouble, jbyteArray);
#ifdef __cplusplus
}
#endif
#endif
+11 -8
View File
@@ -797,11 +797,12 @@ void NodeManager::HandleRegisterClientRequest(
<< ". Is worker: " << is_worker << ". Worker pid "
<< request.worker_pid();
Status status;
if (is_worker) {
// Register the new worker.
bool use_push_task = worker->UsePush();
worker_pool_.RegisterWorker(worker_id, std::move(worker));
if (use_push_task) {
status = worker_pool_.RegisterWorker(worker_id, std::move(worker));
if (status.ok() && use_push_task) {
// only call `HandleWorkerAvailable` when push mode is used.
HandleWorkerAvailable(worker_id);
}
@@ -811,13 +812,15 @@ void NodeManager::HandleRegisterClientRequest(
auto job_id = JobID::FromBinary(request.job_id());
worker->AssignTaskId(driver_task_id);
worker->AssignJobId(job_id);
worker_pool_.RegisterDriver(worker_id, std::move(worker));
local_queues_.AddDriverTaskId(driver_task_id);
RAY_CHECK_OK(gcs_client_->job_table().AppendJobData(
job_id, /*is_dead=*/false, std::time(nullptr),
initial_config_.node_manager_address, request.worker_pid()));
status = worker_pool_.RegisterDriver(worker_id, std::move(worker));
if (status.ok()) {
local_queues_.AddDriverTaskId(driver_task_id);
RAY_CHECK_OK(gcs_client_->job_table().AppendJobData(
job_id, /*is_dead=*/false, std::time(nullptr),
initial_config_.node_manager_address, request.worker_pid()));
}
}
send_reply_callback(Status::OK(), nullptr, nullptr);
send_reply_callback(status, nullptr, nullptr);
}
void NodeManager::HandleDisconnectedActor(const ActorID &actor_id, bool was_local,
+7
View File
@@ -141,6 +141,13 @@ void Worker::AssignTask(const Task &task, const ResourceIdSet &resource_id_set)
// and assigning new task will be done when raylet receives
// `TaskDone` message.
});
if (!status.ok()) {
RAY_LOG(ERROR) << "Failed to assign task " << task.GetTaskSpecification().TaskId()
<< " to worker " << worker_id_;
} else {
RAY_LOG(DEBUG) << "Assigned task " << task.GetTaskSpecification().TaskId()
<< " to worker " << worker_id_;
}
} else {
// Use pull mode. This corresponds to existing python/java workers that haven't been
// migrated to core worker architecture.
+9 -6
View File
@@ -166,30 +166,33 @@ pid_t WorkerPool::StartProcess(const std::vector<std::string> &worker_command_ar
return 0;
}
void WorkerPool::RegisterWorker(const WorkerID &worker_id,
const std::shared_ptr<Worker> &worker) {
Status WorkerPool::RegisterWorker(const WorkerID &worker_id,
const std::shared_ptr<Worker> &worker) {
const auto pid = worker->Pid();
const auto port = worker->Port();
RAY_LOG(DEBUG) << "Registering worker with pid " << pid << ", port: " << port;
auto &state = GetStateForLanguage(worker->GetLanguage());
state.registered_workers.emplace(worker_id, std::move(worker));
auto it = state.starting_worker_processes.find(pid);
if (it == state.starting_worker_processes.end()) {
RAY_LOG(WARNING) << "Received a register request from an unknown worker " << pid;
return;
return Status::Invalid("Unknown worker");
}
it->second--;
if (it->second == 0) {
state.starting_worker_processes.erase(it);
}
state.registered_workers.emplace(worker_id, std::move(worker));
return Status::OK();
}
void WorkerPool::RegisterDriver(const WorkerID &driver_id,
const std::shared_ptr<Worker> &driver) {
Status WorkerPool::RegisterDriver(const WorkerID &driver_id,
const std::shared_ptr<Worker> &driver) {
RAY_CHECK(!driver->GetAssignedTaskId().IsNil());
auto &state = GetStateForLanguage(driver->GetLanguage());
state.registered_drivers.emplace(driver_id, std::move(driver));
return Status::OK();
}
std::shared_ptr<Worker> WorkerPool::GetRegisteredWorker(const WorkerID &worker_id) const {
+4 -2
View File
@@ -51,13 +51,15 @@ class WorkerPool {
/// pool after it becomes idle (e.g., requests a work assignment).
///
/// \param The Worker to be registered.
void RegisterWorker(const WorkerID &worker_id, const std::shared_ptr<Worker> &worker);
/// \return If the registration is successful.
Status RegisterWorker(const WorkerID &worker_id, const std::shared_ptr<Worker> &worker);
/// Register a new driver.
/// Driver is a treated as a special worker, so use WorkerID as key here.
///
/// \param The driver to be registered.
void RegisterDriver(const WorkerID &driver_id, const std::shared_ptr<Worker> &worker);
/// \return If the registration is successful.
Status RegisterDriver(const WorkerID &driver_id, const std::shared_ptr<Worker> &worker);
/// Get the client connection's registered worker.
///
+1 -1
View File
@@ -123,7 +123,7 @@ TEST_F(WorkerPoolTest, HandleWorkerRegistration) {
ASSERT_EQ(worker_pool_.NumWorkerProcessesStarting(), 1);
// Check that we cannot lookup the worker before it's registered.
ASSERT_EQ(worker_pool_.GetRegisteredWorker(worker_id), nullptr);
worker_pool_.RegisterWorker(worker_id, worker);
RAY_CHECK_OK(worker_pool_.RegisterWorker(worker_id, worker));
// Check that we can lookup the worker after it's registered.
ASSERT_EQ(worker_pool_.GetRegisteredWorker(worker_id), worker);
}
+11 -13
View File
@@ -138,14 +138,15 @@ void RayLog::StartRayLog(const std::string &app_name, RayLogLevel severity_thres
log_dir_ = log_dir;
#ifdef RAY_USE_GLOG
google::InitGoogleLogging(app_name_.c_str());
google::SetStderrLogging(GetMappedSeverity(RayLogLevel::ERROR));
for (int i = static_cast<int>(severity_threshold_);
i <= static_cast<int>(RayLogLevel::FATAL); ++i) {
int level = GetMappedSeverity(static_cast<RayLogLevel>(i));
google::base::SetLogger(level, &stdout_logger_singleton);
}
// Enable log file if log_dir_ is not empty.
if (!log_dir_.empty()) {
if (log_dir_.empty()) {
google::SetStderrLogging(GetMappedSeverity(RayLogLevel::ERROR));
for (int i = static_cast<int>(severity_threshold_);
i <= static_cast<int>(RayLogLevel::FATAL); ++i) {
int level = GetMappedSeverity(static_cast<RayLogLevel>(i));
google::base::SetLogger(level, &stdout_logger_singleton);
}
} else {
// Enable log file if log_dir_ is not empty.
auto dir_ends_with_slash = log_dir_;
if (log_dir_[log_dir_.length() - 1] != '/') {
dir_ends_with_slash += "/";
@@ -161,11 +162,8 @@ void RayLog::StartRayLog(const std::string &app_name, RayLogLevel severity_thres
}
}
google::SetLogFilenameExtension(app_name_without_path.c_str());
for (int i = static_cast<int>(severity_threshold_);
i <= static_cast<int>(RayLogLevel::FATAL); ++i) {
int level = GetMappedSeverity(static_cast<RayLogLevel>(i));
google::SetLogDestination(level, dir_ends_with_slash.c_str());
}
int level = GetMappedSeverity(severity_threshold_);
google::SetLogDestination(level, dir_ends_with_slash.c_str());
}
#endif
}