mirror of
https://github.com/wassname/ray.git
synced 2026-08-02 13:01:01 +08:00
[Java] Add runtime context (#4194)
This commit is contained in:
@@ -120,4 +120,11 @@ public final class Ray extends RayCall {
|
||||
public static RayRuntime internal() {
|
||||
return runtime;
|
||||
}
|
||||
|
||||
/**
|
||||
* Get the runtime context.
|
||||
*/
|
||||
public static RuntimeContext getRuntimeContext() {
|
||||
return runtime.getRuntimeContext();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
package org.ray.api;
|
||||
|
||||
import org.ray.api.id.UniqueId;
|
||||
|
||||
/**
|
||||
* A class used for getting information of Ray runtime.
|
||||
*/
|
||||
public interface RuntimeContext {
|
||||
|
||||
/**
|
||||
* Get the current Driver ID.
|
||||
*
|
||||
* If called in a driver, this returns the driver ID. If called in a worker, this returns the ID
|
||||
* of the associated driver.
|
||||
*/
|
||||
UniqueId getCurrentDriverId();
|
||||
|
||||
/**
|
||||
* Get the current actor ID.
|
||||
*
|
||||
* Note, this can only be called in actors.
|
||||
*/
|
||||
UniqueId getCurrentActorId();
|
||||
|
||||
/**
|
||||
* Returns true if the current actor was reconstructed, false if it's created for the first time.
|
||||
*
|
||||
* Note, this method should only be called from an actor creation task.
|
||||
*/
|
||||
boolean wasCurrentActorReconstructed();
|
||||
|
||||
/**
|
||||
* Get the raylet socket name.
|
||||
*/
|
||||
String getRayletSocketName();
|
||||
|
||||
/**
|
||||
* Get the object store socket name.
|
||||
*/
|
||||
String getObjectStoreSocketName();
|
||||
|
||||
/**
|
||||
* Return true if Ray is running in single-process mode, false if Ray is running in cluster mode.
|
||||
*/
|
||||
boolean isSingleProcess();
|
||||
}
|
||||
@@ -3,6 +3,7 @@ package org.ray.api.runtime;
|
||||
import java.util.List;
|
||||
import org.ray.api.RayActor;
|
||||
import org.ray.api.RayObject;
|
||||
import org.ray.api.RuntimeContext;
|
||||
import org.ray.api.WaitResult;
|
||||
import org.ray.api.function.RayFunc;
|
||||
import org.ray.api.id.UniqueId;
|
||||
@@ -93,4 +94,6 @@ public interface RayRuntime {
|
||||
*/
|
||||
<T> RayActor<T> createActor(RayFunc actorFactoryFunc, Object[] args,
|
||||
ActorCreationOptions options);
|
||||
|
||||
RuntimeContext getRuntimeContext();
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import java.util.Map;
|
||||
import java.util.stream.Collectors;
|
||||
import org.ray.api.RayActor;
|
||||
import org.ray.api.RayObject;
|
||||
import org.ray.api.RuntimeContext;
|
||||
import org.ray.api.WaitResult;
|
||||
import org.ray.api.exception.RayException;
|
||||
import org.ray.api.function.RayFunc;
|
||||
@@ -61,6 +62,7 @@ public abstract class AbstractRayRuntime implements RayRuntime {
|
||||
protected RayletClient rayletClient;
|
||||
protected ObjectStoreProxy objectStoreProxy;
|
||||
protected FunctionManager functionManager;
|
||||
protected RuntimeContext runtimeContext;
|
||||
|
||||
public AbstractRayRuntime(RayConfig rayConfig) {
|
||||
this.rayConfig = rayConfig;
|
||||
@@ -68,6 +70,7 @@ public abstract class AbstractRayRuntime implements RayRuntime {
|
||||
worker = new Worker(this);
|
||||
workerContext = new WorkerContext(rayConfig.workerMode,
|
||||
rayConfig.driverId, rayConfig.runMode);
|
||||
runtimeContext = new RuntimeContextImpl(this);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -346,4 +349,9 @@ public abstract class AbstractRayRuntime implements RayRuntime {
|
||||
public RayConfig getRayConfig() {
|
||||
return rayConfig;
|
||||
}
|
||||
|
||||
public RuntimeContext getRuntimeContext() {
|
||||
return runtimeContext;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -173,4 +173,22 @@ public final class RayNativeRuntime extends AbstractRayRuntime {
|
||||
checkpoints.sort((x, y) -> Long.compare(y.timestamp, x.timestamp));
|
||||
return checkpoints;
|
||||
}
|
||||
|
||||
|
||||
/**
|
||||
* Query whether the actor exists in Gcs.
|
||||
*/
|
||||
boolean actorExistsInGcs(UniqueId actorId) {
|
||||
byte[] key = ArrayUtils.addAll("ACTOR".getBytes(), actorId.getBytes());
|
||||
|
||||
// TODO(qwang): refactor this with `GlobalState` after this issue
|
||||
// getting finished. https://github.com/ray-project/ray/issues/3933
|
||||
for (RedisClient client : redisClients) {
|
||||
if (client.exists(key)) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package org.ray.runtime;
|
||||
|
||||
import com.google.common.base.Preconditions;
|
||||
import org.ray.api.RuntimeContext;
|
||||
import org.ray.api.id.UniqueId;
|
||||
import org.ray.runtime.config.RunMode;
|
||||
import org.ray.runtime.config.WorkerMode;
|
||||
import org.ray.runtime.task.TaskSpec;
|
||||
|
||||
public class RuntimeContextImpl implements RuntimeContext {
|
||||
|
||||
private AbstractRayRuntime runtime;
|
||||
|
||||
public RuntimeContextImpl(AbstractRayRuntime runtime) {
|
||||
this.runtime = runtime;
|
||||
}
|
||||
|
||||
@Override
|
||||
public UniqueId getCurrentDriverId() {
|
||||
return runtime.getWorkerContext().getCurrentDriverId();
|
||||
}
|
||||
|
||||
@Override
|
||||
public UniqueId getCurrentActorId() {
|
||||
Preconditions.checkState(runtime.rayConfig.workerMode == WorkerMode.WORKER);
|
||||
return runtime.getWorker().getCurrentActorId();
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean wasCurrentActorReconstructed() {
|
||||
TaskSpec currentTask = runtime.getWorkerContext().getCurrentTask();
|
||||
Preconditions.checkState(currentTask != null && currentTask.isActorCreationTask(),
|
||||
"This method can only be called from an actor creation task.");
|
||||
if (isSingleProcess()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return ((RayNativeRuntime) runtime).actorExistsInGcs(getCurrentActorId());
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getRayletSocketName() {
|
||||
return runtime.getRayConfig().rayletSocketName;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getObjectStoreSocketName() {
|
||||
return runtime.getRayConfig().objectStoreSocketName;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean isSingleProcess() {
|
||||
return RunMode.SINGLE_PROCESS == runtime.getRayConfig().runMode;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -63,6 +63,10 @@ public class Worker {
|
||||
this.runtime = runtime;
|
||||
}
|
||||
|
||||
public UniqueId getCurrentActorId() {
|
||||
return currentActorId;
|
||||
}
|
||||
|
||||
public void loop() {
|
||||
while (true) {
|
||||
LOGGER.info("Fetching new task in thread {}.", Thread.currentThread().getName());
|
||||
@@ -86,6 +90,11 @@ public class Worker {
|
||||
// Set context
|
||||
runtime.getWorkerContext().setCurrentTask(spec, rayFunction.classLoader);
|
||||
Thread.currentThread().setContextClassLoader(rayFunction.classLoader);
|
||||
|
||||
if (spec.isActorCreationTask()) {
|
||||
currentActorId = returnId;
|
||||
}
|
||||
|
||||
// Get local actor object and arguments.
|
||||
Object actor = null;
|
||||
if (spec.isActorTask()) {
|
||||
@@ -94,6 +103,7 @@ public class Worker {
|
||||
throw actorCreationException;
|
||||
}
|
||||
actor = currentActor;
|
||||
|
||||
}
|
||||
Object[] args = ArgumentsBuilder.unwrap(spec, rayFunction.classLoader);
|
||||
// Execute the task.
|
||||
@@ -112,7 +122,6 @@ public class Worker {
|
||||
} else {
|
||||
maybeLoadCheckpoint(result, returnId);
|
||||
currentActor = result;
|
||||
currentActorId = returnId;
|
||||
}
|
||||
LOGGER.info("Finished executing task {}", spec.taskId);
|
||||
} catch (Exception e) {
|
||||
@@ -121,7 +130,6 @@ public class Worker {
|
||||
runtime.put(returnId, new RayTaskException("Error executing task " + spec, e));
|
||||
} else {
|
||||
actorCreationException = e;
|
||||
currentActorId = returnId;
|
||||
}
|
||||
} finally {
|
||||
Thread.currentThread().setContextClassLoader(oldLoader);
|
||||
|
||||
@@ -26,6 +26,8 @@ public class WorkerContext {
|
||||
*/
|
||||
private ThreadLocal<Integer> taskIndex;
|
||||
|
||||
private ThreadLocal<TaskSpec> currentTask;
|
||||
|
||||
private UniqueId currentDriverId;
|
||||
|
||||
private ClassLoader currentClassLoader;
|
||||
@@ -46,6 +48,7 @@ public class WorkerContext {
|
||||
putIndex = ThreadLocal.withInitial(() -> 0);
|
||||
currentTaskId = ThreadLocal.withInitial(UniqueId::randomId);
|
||||
this.runMode = runMode;
|
||||
currentTask = ThreadLocal.withInitial(() -> null);
|
||||
currentClassLoader = null;
|
||||
if (workerMode == WorkerMode.DRIVER) {
|
||||
workerId = driverId;
|
||||
@@ -83,6 +86,7 @@ public class WorkerContext {
|
||||
this.currentDriverId = task.driverId;
|
||||
taskIndex.set(0);
|
||||
putIndex.set(0);
|
||||
this.currentTask.set(task);
|
||||
currentClassLoader = classLoader;
|
||||
}
|
||||
|
||||
@@ -124,4 +128,10 @@ public class WorkerContext {
|
||||
return currentClassLoader;
|
||||
}
|
||||
|
||||
/**
|
||||
* Get the current task.
|
||||
*/
|
||||
public TaskSpec getCurrentTask() {
|
||||
return this.currentTask.get();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -85,4 +85,14 @@ public class RedisClient {
|
||||
return jedis.lrange(key, start, end);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Whether the key exists in Redis.
|
||||
*/
|
||||
public boolean exists(byte[] key) {
|
||||
try (Jedis jedis = jedisPool.getResource()) {
|
||||
return jedis.exists(key);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -268,6 +268,10 @@ public class RunManager {
|
||||
cmd.add("-Dray.logging.file.path=" + logFile);
|
||||
}
|
||||
|
||||
// socket names
|
||||
cmd.add("-Dray.raylet.socket-name=" + rayConfig.rayletSocketName);
|
||||
cmd.add("-Dray.object-store.socket-name=" + rayConfig.objectStoreSocketName);
|
||||
|
||||
// Config overwrite
|
||||
cmd.add("-Dray.redis.address=" + rayConfig.getRedisAddress());
|
||||
|
||||
|
||||
@@ -24,6 +24,16 @@ public class ActorReconstructionTest extends BaseTest {
|
||||
|
||||
protected int value = 0;
|
||||
|
||||
private boolean wasCurrentActorReconstructed = false;
|
||||
|
||||
public Counter() {
|
||||
wasCurrentActorReconstructed = Ray.getRuntimeContext().wasCurrentActorReconstructed();
|
||||
}
|
||||
|
||||
public boolean wasCurrentActorReconstructed() {
|
||||
return wasCurrentActorReconstructed;
|
||||
}
|
||||
|
||||
public int increase() {
|
||||
value += 1;
|
||||
return value;
|
||||
@@ -48,6 +58,8 @@ public class ActorReconstructionTest extends BaseTest {
|
||||
Ray.call(Counter::increase, actor).get();
|
||||
}
|
||||
|
||||
Assert.assertFalse(Ray.call(Counter::wasCurrentActorReconstructed, actor).get());
|
||||
|
||||
// Kill the actor process.
|
||||
int pid = Ray.call(Counter::getPid, actor).get();
|
||||
Runtime.getRuntime().exec("kill -9 " + pid);
|
||||
@@ -58,6 +70,8 @@ public class ActorReconstructionTest extends BaseTest {
|
||||
int value = Ray.call(Counter::increase, actor).get();
|
||||
Assert.assertEquals(value, 4);
|
||||
|
||||
Assert.assertTrue(Ray.call(Counter::wasCurrentActorReconstructed, actor).get());
|
||||
|
||||
// Kill the actor process again.
|
||||
pid = Ray.call(Counter::getPid, actor).get();
|
||||
Runtime.getRuntime().exec("kill -9 " + pid);
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
package org.ray.api.test;
|
||||
|
||||
import org.ray.api.Ray;
|
||||
import org.ray.api.RayActor;
|
||||
import org.ray.api.annotation.RayRemote;
|
||||
import org.ray.api.id.UniqueId;
|
||||
import org.testng.Assert;
|
||||
import org.testng.annotations.Test;
|
||||
|
||||
public class RuntimeContextTest extends BaseTest {
|
||||
|
||||
private static UniqueId DRIVER_ID =
|
||||
UniqueId.fromHexString("0011223344556677889900112233445566778899");
|
||||
private static String RAYLET_SOCKET_NAME = "/tmp/ray/test/raylet_socket";
|
||||
private static String OBJECT_STORE_SOCKET_NAME = "/tmp/ray/test/object_store_socket";
|
||||
|
||||
@Override
|
||||
public void beforeInitRay() {
|
||||
System.setProperty("ray.driver.id", DRIVER_ID.toString());
|
||||
System.setProperty("ray.raylet.socket-name", RAYLET_SOCKET_NAME);
|
||||
System.setProperty("ray.object-store.socket-name", OBJECT_STORE_SOCKET_NAME);
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testRuntimeContextInDriver() {
|
||||
Assert.assertEquals(DRIVER_ID, Ray.getRuntimeContext().getCurrentDriverId());
|
||||
Assert.assertEquals(RAYLET_SOCKET_NAME, Ray.getRuntimeContext().getRayletSocketName());
|
||||
Assert.assertEquals(OBJECT_STORE_SOCKET_NAME,
|
||||
Ray.getRuntimeContext().getObjectStoreSocketName());
|
||||
}
|
||||
|
||||
@RayRemote
|
||||
public static class RuntimeContextTester {
|
||||
|
||||
public String testRuntimeContext(UniqueId actorId) {
|
||||
Assert.assertEquals(DRIVER_ID, Ray.getRuntimeContext().getCurrentDriverId());
|
||||
Assert.assertEquals(actorId, Ray.getRuntimeContext().getCurrentActorId());
|
||||
Assert.assertEquals(RAYLET_SOCKET_NAME, Ray.getRuntimeContext().getRayletSocketName());
|
||||
Assert.assertEquals(OBJECT_STORE_SOCKET_NAME,
|
||||
Ray.getRuntimeContext().getObjectStoreSocketName());
|
||||
return "ok";
|
||||
}
|
||||
}
|
||||
|
||||
@Test
|
||||
public void testRuntimeContextInActor() {
|
||||
RayActor<RuntimeContextTester> actor = Ray.createActor(RuntimeContextTester::new);
|
||||
Assert.assertEquals("ok",
|
||||
Ray.call(RuntimeContextTester::testRuntimeContext, actor, actor.getId()).get());
|
||||
}
|
||||
|
||||
}
|
||||
Reference in New Issue
Block a user