mirror of
https://github.com/wassname/ray.git
synced 2026-08-13 12:30:18 +08:00
[Java] Local and distributed ref counting in Java (#9371)
This commit is contained in:
@@ -164,6 +164,18 @@ void CoreWorkerProcess::EnsureInitialized() {
|
||||
<< "shutdown.";
|
||||
}
|
||||
|
||||
std::shared_ptr<CoreWorker> CoreWorkerProcess::TryGetWorker(const WorkerID &worker_id) {
|
||||
if (!instance_) {
|
||||
return nullptr;
|
||||
}
|
||||
absl::ReaderMutexLock workers_lock(&instance_->worker_map_mutex_);
|
||||
auto it = instance_->workers_.find(worker_id);
|
||||
if (it != instance_->workers_.end()) {
|
||||
return it->second;
|
||||
}
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
CoreWorker &CoreWorkerProcess::GetCoreWorker() {
|
||||
EnsureInitialized();
|
||||
if (instance_->options_.num_workers == 1) {
|
||||
|
||||
@@ -178,6 +178,13 @@ class CoreWorkerProcess {
|
||||
/// `CoreWorkerProcess` has full control of the destruction timing of `CoreWorker`.
|
||||
static CoreWorker &GetCoreWorker();
|
||||
|
||||
/// Try to get the `CoreWorker` instance by worker ID.
|
||||
/// If the current thread is not associated with a core worker, returns a null pointer.
|
||||
///
|
||||
/// \param[in] workerId The worker ID.
|
||||
/// \return The `CoreWorker` instance.
|
||||
static std::shared_ptr<CoreWorker> TryGetWorker(const WorkerID &worker_id);
|
||||
|
||||
/// Set the core worker associated with the current thread by worker ID.
|
||||
/// Currently used by Java worker only.
|
||||
///
|
||||
|
||||
@@ -18,10 +18,10 @@
|
||||
|
||||
#include <sstream>
|
||||
|
||||
#include "jni_utils.h"
|
||||
#include "ray/common/id.h"
|
||||
#include "ray/core_worker/actor_handle.h"
|
||||
#include "ray/core_worker/core_worker.h"
|
||||
#include "jni_utils.h"
|
||||
|
||||
thread_local JNIEnv *local_env = nullptr;
|
||||
jobject java_task_executor = nullptr;
|
||||
@@ -72,6 +72,19 @@ jobject ToJavaArgs(JNIEnv *env, jbooleanArray java_check_results,
|
||||
}
|
||||
}
|
||||
|
||||
JNIEnv *GetJNIEnv() {
|
||||
JNIEnv *env = local_env;
|
||||
if (!env) {
|
||||
// Attach the native thread to JVM.
|
||||
auto status =
|
||||
jvm->AttachCurrentThreadAsDaemon(reinterpret_cast<void **>(&env), nullptr);
|
||||
RAY_CHECK(status == JNI_OK) << "Failed to get JNIEnv. Return code: " << status;
|
||||
local_env = env;
|
||||
}
|
||||
RAY_CHECK(env);
|
||||
return env;
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
@@ -98,16 +111,7 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
|
||||
const std::vector<ObjectID> &arg_reference_ids,
|
||||
const std::vector<ObjectID> &return_ids,
|
||||
std::vector<std::shared_ptr<ray::RayObject>> *results) {
|
||||
JNIEnv *env = local_env;
|
||||
if (!env) {
|
||||
// Attach the native thread to JVM.
|
||||
auto status =
|
||||
jvm->AttachCurrentThreadAsDaemon(reinterpret_cast<void **>(&env), nullptr);
|
||||
RAY_CHECK(status == JNI_OK) << "Failed to get JNIEnv. Return code: " << status;
|
||||
local_env = env;
|
||||
}
|
||||
|
||||
RAY_CHECK(env);
|
||||
JNIEnv *env = GetJNIEnv();
|
||||
RAY_CHECK(java_task_executor);
|
||||
|
||||
// convert RayFunction
|
||||
@@ -141,6 +145,8 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
|
||||
env->CallObjectMethod(java_task_executor, java_task_executor_execute,
|
||||
ray_function_array_list, args_array_list);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
|
||||
// Process return objects.
|
||||
if (!return_ids.empty()) {
|
||||
std::vector<std::shared_ptr<ray::RayObject>> return_objects;
|
||||
JavaListToNativeVector<std::shared_ptr<ray::RayObject>>(
|
||||
@@ -148,8 +154,26 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
|
||||
[](JNIEnv *env, jobject java_native_ray_object) {
|
||||
return JavaNativeRayObjectToNativeRayObject(env, java_native_ray_object);
|
||||
});
|
||||
for (auto &obj : return_objects) {
|
||||
results->push_back(obj);
|
||||
std::vector<size_t> data_sizes;
|
||||
std::vector<std::shared_ptr<ray::Buffer>> metadatas;
|
||||
std::vector<std::vector<ray::ObjectID>> contained_object_ids;
|
||||
for (size_t i = 0; i < return_objects.size(); i++) {
|
||||
data_sizes.push_back(
|
||||
return_objects[i]->HasData() ? return_objects[i]->GetData()->Size() : 0);
|
||||
metadatas.push_back(return_objects[i]->GetMetadata());
|
||||
contained_object_ids.push_back(return_objects[i]->GetNestedIds());
|
||||
}
|
||||
RAY_CHECK_OK(ray::CoreWorkerProcess::GetCoreWorker().AllocateReturnObjects(
|
||||
return_ids, data_sizes, metadatas, contained_object_ids, results));
|
||||
for (size_t i = 0; i < data_sizes.size(); i++) {
|
||||
auto result = (*results)[i];
|
||||
// A nullptr is returned if the object already exists.
|
||||
if (result != nullptr) {
|
||||
if (result->HasData()) {
|
||||
memcpy(result->GetData()->Data(), return_objects[i]->GetData()->Data(),
|
||||
data_sizes[i]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -159,6 +183,26 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
|
||||
return ray::Status::OK();
|
||||
};
|
||||
|
||||
auto gc_collect = []() {
|
||||
// A Java worker process usually contains more than one worker.
|
||||
// A LocalGC request is likely to be received by multiple workers in a short time.
|
||||
// Here we ensure that the 1 second interval of `System.gc()` execution is
|
||||
// guaranteed no matter how frequent the requests are received and how many workers
|
||||
// the process has.
|
||||
static absl::Mutex mutex;
|
||||
static int64_t last_gc_time_ms = 0;
|
||||
absl::MutexLock lock(&mutex);
|
||||
int64_t start = current_time_ms();
|
||||
if (last_gc_time_ms + 1000 < start) {
|
||||
JNIEnv *env = GetJNIEnv();
|
||||
RAY_LOG(INFO) << "Calling System.gc() ...";
|
||||
env->CallStaticObjectMethod(java_system_class, java_system_gc);
|
||||
last_gc_time_ms = current_time_ms();
|
||||
RAY_LOG(INFO) << "GC finished in " << (double) (last_gc_time_ms - start) / 1000
|
||||
<< " seconds.";
|
||||
}
|
||||
};
|
||||
|
||||
ray::CoreWorkerOptions options = {
|
||||
static_cast<ray::WorkerType>(workerMode), // worker_type
|
||||
ray::Language::JAVA, // langauge
|
||||
@@ -178,10 +222,10 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_RayNativeRuntime_nativeInitialize(
|
||||
"", // stderr_file
|
||||
task_execution_callback, // task_execution_callback
|
||||
nullptr, // check_signals
|
||||
nullptr, // gc_collect
|
||||
gc_collect, // gc_collect
|
||||
nullptr, // get_lang_stack
|
||||
nullptr, // kill_main
|
||||
false, // ref_counting_enabled
|
||||
true, // ref_counting_enabled
|
||||
false, // is_local_mode
|
||||
static_cast<int>(numWorkersPerProcess), // num_workers
|
||||
};
|
||||
|
||||
@@ -14,10 +14,47 @@
|
||||
|
||||
#include "io_ray_runtime_object_NativeObjectStore.h"
|
||||
#include <jni.h>
|
||||
#include "jni_utils.h"
|
||||
#include "ray/common/id.h"
|
||||
#include "ray/core_worker/common.h"
|
||||
#include "ray/core_worker/core_worker.h"
|
||||
#include "jni_utils.h"
|
||||
|
||||
ray::Status PutSerializedObject(JNIEnv *env, jobject obj, ray::ObjectID object_id,
|
||||
ray::ObjectID *out_object_id, bool pin_object = true) {
|
||||
auto native_ray_object = JavaNativeRayObjectToNativeRayObject(env, obj);
|
||||
RAY_CHECK(native_ray_object != nullptr);
|
||||
|
||||
size_t data_size = 0;
|
||||
if (native_ray_object->HasData()) {
|
||||
data_size = native_ray_object->GetData()->Size();
|
||||
}
|
||||
std::shared_ptr<ray::Buffer> data;
|
||||
ray::Status status;
|
||||
if (object_id.IsNil()) {
|
||||
status = ray::CoreWorkerProcess::GetCoreWorker().Create(
|
||||
native_ray_object->GetMetadata(), data_size, native_ray_object->GetNestedIds(),
|
||||
out_object_id, &data);
|
||||
} else {
|
||||
status = ray::CoreWorkerProcess::GetCoreWorker().Create(
|
||||
native_ray_object->GetMetadata(), data_size, object_id, &data);
|
||||
*out_object_id = object_id;
|
||||
}
|
||||
if (!status.ok()) {
|
||||
return status;
|
||||
}
|
||||
|
||||
// If data is nullptr, that means the ObjectID already existed, which we ignore.
|
||||
// TODO(edoakes): this is hacky, we should return the error instead and deal with it
|
||||
// here.
|
||||
if (data != nullptr) {
|
||||
if (data->Size() > 0) {
|
||||
memcpy(data->Data(), native_ray_object->GetData()->Data(), data->Size());
|
||||
}
|
||||
RAY_CHECK_OK(ray::CoreWorkerProcess::GetCoreWorker().Seal(
|
||||
*out_object_id, pin_object && object_id.IsNil()));
|
||||
}
|
||||
return ray::Status::OK();
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
@@ -26,10 +63,9 @@ extern "C" {
|
||||
JNIEXPORT jbyteArray JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativePut__Lio_ray_runtime_object_NativeRayObject_2(
|
||||
JNIEnv *env, jclass, jobject obj) {
|
||||
auto ray_object = JavaNativeRayObjectToNativeRayObject(env, obj);
|
||||
RAY_CHECK(ray_object != nullptr);
|
||||
ray::ObjectID object_id;
|
||||
auto status = ray::CoreWorkerProcess::GetCoreWorker().Put(*ray_object, {}, &object_id);
|
||||
auto status = PutSerializedObject(env, obj, /*object_id=*/ray::ObjectID::Nil(),
|
||||
/*out_object_id=*/&object_id, /*pin_object=*/true);
|
||||
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, nullptr);
|
||||
return IdToJavaByteArray<ray::ObjectID>(env, object_id);
|
||||
}
|
||||
@@ -38,9 +74,10 @@ JNIEXPORT void JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativePut___3BLio_ray_runtime_object_NativeRayObject_2(
|
||||
JNIEnv *env, jclass, 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 = ray::CoreWorkerProcess::GetCoreWorker().Put(*ray_object, {}, object_id);
|
||||
ray::ObjectID dummy_object_id;
|
||||
auto status =
|
||||
PutSerializedObject(env, obj, object_id,
|
||||
/*out_object_id=*/&dummy_object_id, /*pin_object=*/true);
|
||||
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
|
||||
}
|
||||
|
||||
@@ -71,7 +108,10 @@ JNIEXPORT jobject JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeWai
|
||||
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);
|
||||
jobject java_item =
|
||||
env->NewObject(java_boolean_class, java_boolean_init, (jboolean)item);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
return java_item;
|
||||
});
|
||||
}
|
||||
|
||||
@@ -88,6 +128,49 @@ JNIEXPORT void JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeDelete
|
||||
THROW_EXCEPTION_AND_RETURN_IF_NOT_OK(env, status, (void)0);
|
||||
}
|
||||
|
||||
JNIEXPORT void JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativeAddLocalReference(
|
||||
JNIEnv *env, jclass, jbyteArray workerId, jbyteArray objectId) {
|
||||
auto worker_id = JavaByteArrayToId<ray::WorkerID>(env, workerId);
|
||||
auto object_id = JavaByteArrayToId<ray::ObjectID>(env, objectId);
|
||||
auto core_worker = ray::CoreWorkerProcess::TryGetWorker(worker_id);
|
||||
RAY_CHECK(core_worker);
|
||||
core_worker->AddLocalReference(object_id);
|
||||
}
|
||||
|
||||
JNIEXPORT void JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativeRemoveLocalReference(
|
||||
JNIEnv *env, jclass, jbyteArray workerId, jbyteArray objectId) {
|
||||
auto worker_id = JavaByteArrayToId<ray::WorkerID>(env, workerId);
|
||||
auto object_id = JavaByteArrayToId<ray::ObjectID>(env, objectId);
|
||||
// We can't control the timing of Java GC, so it's normal that this method is called but
|
||||
// core worker is shutting down (or already shut down). If we can't get a core worker
|
||||
// instance here, skip calling the `RemoveLocalReference` method.
|
||||
auto core_worker = ray::CoreWorkerProcess::TryGetWorker(worker_id);
|
||||
if (core_worker) {
|
||||
core_worker->RemoveLocalReference(object_id);
|
||||
}
|
||||
}
|
||||
|
||||
JNIEXPORT jobject JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativeGetAllReferenceCounts(JNIEnv *env,
|
||||
jclass) {
|
||||
auto reference_counts = ray::CoreWorkerProcess::GetCoreWorker().GetAllReferenceCounts();
|
||||
return NativeMapToJavaMap<ray::ObjectID, std::pair<size_t, size_t>>(
|
||||
env, reference_counts,
|
||||
[](JNIEnv *env, const ray::ObjectID &key) {
|
||||
return IdToJavaByteArray<ObjectID>(env, key);
|
||||
},
|
||||
[](JNIEnv *env, const std::pair<size_t, size_t> &value) {
|
||||
jlongArray array = env->NewLongArray(2);
|
||||
jlong *elements = env->GetLongArrayElements(array, nullptr);
|
||||
elements[0] = static_cast<jlong>(value.first);
|
||||
elements[1] = static_cast<jlong>(value.second);
|
||||
env->ReleaseLongArrayElements(array, elements, 0);
|
||||
return array;
|
||||
});
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
@@ -65,6 +65,35 @@ JNIEXPORT jobject JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeWai
|
||||
JNIEXPORT void JNICALL Java_io_ray_runtime_object_NativeObjectStore_nativeDelete(
|
||||
JNIEnv *, jclass, jobject, jboolean, jboolean);
|
||||
|
||||
/*
|
||||
* Class: io_ray_runtime_object_NativeObjectStore
|
||||
* Method: nativeAddLocalReference
|
||||
* Signature: ([B)V
|
||||
*/
|
||||
JNIEXPORT void JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativeAddLocalReference(JNIEnv *, jclass,
|
||||
jbyteArray,
|
||||
jbyteArray);
|
||||
|
||||
/*
|
||||
* Class: io_ray_runtime_object_NativeObjectStore
|
||||
* Method: nativeRemoveLocalReference
|
||||
* Signature: ([B)V
|
||||
*/
|
||||
JNIEXPORT void JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativeRemoveLocalReference(JNIEnv *, jclass,
|
||||
jbyteArray,
|
||||
jbyteArray);
|
||||
|
||||
/*
|
||||
* Class: io_ray_runtime_object_NativeObjectStore
|
||||
* Method: nativeGetAllReferenceCounts
|
||||
* Signature: ()Ljava/util/Map;
|
||||
*/
|
||||
JNIEXPORT jobject JNICALL
|
||||
Java_io_ray_runtime_object_NativeObjectStore_nativeGetAllReferenceCounts(JNIEnv *,
|
||||
jclass);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
@@ -76,7 +76,6 @@ inline std::vector<std::unique_ptr<ray::TaskArg>> ToTaskArgs(JNIEnv *env, jobjec
|
||||
|
||||
inline std::unordered_map<std::string, double> ToResources(JNIEnv *env,
|
||||
jobject java_resources) {
|
||||
std::unordered_map<std::string, double> resources;
|
||||
return JavaMapToNativeMap<std::string, double>(
|
||||
env, java_resources,
|
||||
[](JNIEnv *env, jobject java_key) {
|
||||
|
||||
@@ -34,6 +34,10 @@ jmethodID java_array_list_init_with_capacity;
|
||||
|
||||
jclass java_map_class;
|
||||
jmethodID java_map_entry_set;
|
||||
jmethodID java_map_put;
|
||||
|
||||
jclass java_hash_map_class;
|
||||
jmethodID java_hash_map_init;
|
||||
|
||||
jclass java_set_class;
|
||||
jmethodID java_set_iterator;
|
||||
@@ -46,6 +50,9 @@ jclass java_map_entry_class;
|
||||
jmethodID java_map_entry_get_key;
|
||||
jmethodID java_map_entry_get_value;
|
||||
|
||||
jclass java_system_class;
|
||||
jmethodID java_system_gc;
|
||||
|
||||
jclass java_ray_exception_class;
|
||||
|
||||
jclass java_jni_exception_util_class;
|
||||
@@ -86,6 +93,7 @@ jclass java_native_ray_object_class;
|
||||
jmethodID java_native_ray_object_init;
|
||||
jfieldID java_native_ray_object_data;
|
||||
jfieldID java_native_ray_object_metadata;
|
||||
jfieldID java_native_ray_object_contained_object_ids;
|
||||
|
||||
jclass java_task_executor_class;
|
||||
jmethodID java_task_executor_parse_function_arguments;
|
||||
@@ -135,6 +143,11 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
|
||||
|
||||
java_map_class = LoadClass(env, "java/util/Map");
|
||||
java_map_entry_set = env->GetMethodID(java_map_class, "entrySet", "()Ljava/util/Set;");
|
||||
java_map_put = env->GetMethodID(
|
||||
java_map_class, "put", "(Ljava/lang/Object;Ljava/lang/Object;)Ljava/lang/Object;");
|
||||
|
||||
java_hash_map_class = LoadClass(env, "java/util/HashMap");
|
||||
java_hash_map_init = env->GetMethodID(java_hash_map_class, "<init>", "()V");
|
||||
|
||||
java_set_class = LoadClass(env, "java/util/Set");
|
||||
java_set_iterator =
|
||||
@@ -151,6 +164,9 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
|
||||
java_map_entry_get_value =
|
||||
env->GetMethodID(java_map_entry_class, "getValue", "()Ljava/lang/Object;");
|
||||
|
||||
java_system_class = LoadClass(env, "java/lang/System");
|
||||
java_system_gc = env->GetStaticMethodID(java_system_class, "gc", "()V");
|
||||
|
||||
java_ray_exception_class = LoadClass(env, "io/ray/api/exception/RayException");
|
||||
|
||||
java_jni_exception_util_class = LoadClass(env, "io/ray/runtime/util/JniExceptionUtil");
|
||||
@@ -220,6 +236,8 @@ jint JNI_OnLoad(JavaVM *vm, void *reserved) {
|
||||
env->GetFieldID(java_native_ray_object_class, "data", "[B");
|
||||
java_native_ray_object_metadata =
|
||||
env->GetFieldID(java_native_ray_object_class, "metadata", "[B");
|
||||
java_native_ray_object_contained_object_ids = env->GetFieldID(
|
||||
java_native_ray_object_class, "containedObjectIds", "Ljava/util/List;");
|
||||
|
||||
java_task_executor_class = LoadClass(env, "io/ray/runtime/task/TaskExecutor");
|
||||
java_task_executor_parse_function_arguments = env->GetMethodID(
|
||||
@@ -241,9 +259,11 @@ void JNI_OnUnload(JavaVM *vm, void *reserved) {
|
||||
env->DeleteGlobalRef(java_list_class);
|
||||
env->DeleteGlobalRef(java_array_list_class);
|
||||
env->DeleteGlobalRef(java_map_class);
|
||||
env->DeleteGlobalRef(java_hash_map_class);
|
||||
env->DeleteGlobalRef(java_set_class);
|
||||
env->DeleteGlobalRef(java_iterator_class);
|
||||
env->DeleteGlobalRef(java_map_entry_class);
|
||||
env->DeleteGlobalRef(java_system_class);
|
||||
env->DeleteGlobalRef(java_ray_exception_class);
|
||||
env->DeleteGlobalRef(java_jni_exception_util_class);
|
||||
env->DeleteGlobalRef(java_base_id_class);
|
||||
|
||||
@@ -58,6 +58,13 @@ extern jmethodID java_array_list_init_with_capacity;
|
||||
extern jclass java_map_class;
|
||||
/// entrySet method of Map interface
|
||||
extern jmethodID java_map_entry_set;
|
||||
/// put method of Map interface
|
||||
extern jmethodID java_map_put;
|
||||
|
||||
/// HashMap class
|
||||
extern jclass java_hash_map_class;
|
||||
/// Constructor of HashMap class
|
||||
extern jmethodID java_hash_map_init;
|
||||
|
||||
/// Set interface
|
||||
extern jclass java_set_class;
|
||||
@@ -78,6 +85,11 @@ extern jmethodID java_map_entry_get_key;
|
||||
/// getValue method of Map.Entry interface
|
||||
extern jmethodID java_map_entry_get_value;
|
||||
|
||||
/// System class
|
||||
extern jclass java_system_class;
|
||||
/// gc method of System class
|
||||
extern jmethodID java_system_gc;
|
||||
|
||||
/// RayException class
|
||||
extern jclass java_ray_exception_class;
|
||||
|
||||
@@ -149,6 +161,8 @@ extern jmethodID java_native_ray_object_init;
|
||||
extern jfieldID java_native_ray_object_data;
|
||||
/// metadata field of NativeRayObject class
|
||||
extern jfieldID java_native_ray_object_metadata;
|
||||
// containedObjectIds field of NativeRayObject class
|
||||
extern jfieldID java_native_ray_object_contained_object_ids;
|
||||
|
||||
/// TaskExecutor class
|
||||
extern jclass java_task_executor_class;
|
||||
@@ -323,6 +337,7 @@ inline jobject NativeVectorToJavaList(
|
||||
jobject java_list =
|
||||
env->NewObject(java_array_list_class, java_array_list_init_with_capacity,
|
||||
(jint)native_vector.size());
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
for (const auto &item : native_vector) {
|
||||
auto element = element_converter(env, item);
|
||||
env->CallVoidMethod(java_list, java_list_add, element);
|
||||
@@ -364,10 +379,11 @@ inline std::unordered_map<key_type, value_type> JavaMapToNativeMap(
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
jobject map_entry = env->CallObjectMethod(iterator, java_iterator_next);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
auto java_key = (jstring)env->CallObjectMethod(map_entry, java_map_entry_get_key);
|
||||
auto java_key = env->CallObjectMethod(map_entry, java_map_entry_get_key);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
key_type key = key_converter(env, java_key);
|
||||
auto java_value = env->CallObjectMethod(map_entry, java_map_entry_get_value);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
value_type value = value_converter(env, java_value);
|
||||
native_map.emplace(key, value);
|
||||
env->DeleteLocalRef(java_key);
|
||||
@@ -381,6 +397,25 @@ inline std::unordered_map<key_type, value_type> JavaMapToNativeMap(
|
||||
return native_map;
|
||||
}
|
||||
|
||||
/// Convert a C++ std::unordered_map<?, ?> to a Java Map<?, ?>
|
||||
template <typename key_type, typename value_type>
|
||||
inline jobject NativeMapToJavaMap(
|
||||
JNIEnv *env, const std::unordered_map<key_type, value_type> &native_map,
|
||||
const std::function<jobject(JNIEnv *, const key_type &)> &key_converter,
|
||||
const std::function<jobject(JNIEnv *, const value_type &)> &value_converter) {
|
||||
jobject java_map = env->NewObject(java_hash_map_class, java_hash_map_init);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
for (const auto &entry : native_map) {
|
||||
jobject java_key = key_converter(env, entry.first);
|
||||
jobject java_value = value_converter(env, entry.second);
|
||||
env->CallObjectMethod(java_map, java_map_put, java_key, java_value);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
env->DeleteLocalRef(java_key);
|
||||
env->DeleteLocalRef(java_value);
|
||||
}
|
||||
return java_map;
|
||||
}
|
||||
|
||||
/// Convert a C++ ray::Buffer to a Java byte array.
|
||||
inline jbyteArray NativeBufferToJavaByteArray(JNIEnv *env,
|
||||
const std::shared_ptr<ray::Buffer> buffer) {
|
||||
@@ -423,9 +458,16 @@ inline std::shared_ptr<ray::RayObject> JavaNativeRayObjectToNativeRayObject(
|
||||
if (metadata_buffer && metadata_buffer->Size() == 0) {
|
||||
metadata_buffer = nullptr;
|
||||
}
|
||||
// TODO: Support nested IDs for Java.
|
||||
|
||||
auto java_contained_ids =
|
||||
env->GetObjectField(java_obj, java_native_ray_object_contained_object_ids);
|
||||
std::vector<ray::ObjectID> contained_object_ids;
|
||||
JavaListToNativeVector<ray::ObjectID>(
|
||||
env, java_contained_ids, &contained_object_ids, [](JNIEnv *env, jobject id) {
|
||||
return JavaByteArrayToId<ray::ObjectID>(env, static_cast<jbyteArray>(id));
|
||||
});
|
||||
return std::make_shared<ray::RayObject>(data_buffer, metadata_buffer,
|
||||
std::vector<ray::ObjectID>());
|
||||
contained_object_ids);
|
||||
}
|
||||
|
||||
/// Convert a C++ ray::RayObject to a Java NativeRayObject.
|
||||
@@ -438,6 +480,7 @@ inline jobject NativeRayObjectToJavaNativeRayObject(
|
||||
auto java_metadata = NativeBufferToJavaByteArray(env, rayObject->GetMetadata());
|
||||
auto java_obj = env->NewObject(java_native_ray_object_class,
|
||||
java_native_ray_object_init, java_data, java_metadata);
|
||||
RAY_CHECK_JAVA_EXCEPTION(env);
|
||||
env->DeleteLocalRef(java_metadata);
|
||||
env->DeleteLocalRef(java_data);
|
||||
return java_obj;
|
||||
|
||||
Reference in New Issue
Block a user