mirror of
https://github.com/wassname/ray.git
synced 2026-08-16 11:27:09 +08:00
[xray] Implements ray.wait (#2162)
Implements ray.wait for xray. Fixes #1128.
This commit is contained in:
@@ -179,6 +179,58 @@ static PyObject *PyLocalSchedulerClient_set_actor_frontier(PyObject *self,
|
||||
Py_RETURN_NONE;
|
||||
}
|
||||
|
||||
static PyObject *PyLocalSchedulerClient_wait(PyObject *self, PyObject *args) {
|
||||
PyObject *py_object_ids;
|
||||
int num_returns;
|
||||
int64_t timeout_ms;
|
||||
PyObject *py_wait_local;
|
||||
|
||||
if (!PyArg_ParseTuple(args, "OilO", &py_object_ids, &num_returns, &timeout_ms,
|
||||
&py_wait_local)) {
|
||||
return NULL;
|
||||
}
|
||||
|
||||
bool wait_local = PyObject_IsTrue(py_wait_local);
|
||||
|
||||
// Convert object ids.
|
||||
PyObject *iter = PyObject_GetIter(py_object_ids);
|
||||
if (!iter) {
|
||||
return NULL;
|
||||
}
|
||||
std::vector<ObjectID> object_ids;
|
||||
while (true) {
|
||||
PyObject *next = PyIter_Next(iter);
|
||||
ObjectID object_id;
|
||||
if (!next) {
|
||||
break;
|
||||
}
|
||||
if (!PyObjectToUniqueID(next, &object_id)) {
|
||||
// Error parsing object id.
|
||||
return NULL;
|
||||
}
|
||||
object_ids.push_back(object_id);
|
||||
}
|
||||
|
||||
// Invoke wait.
|
||||
std::pair<std::vector<ObjectID>, std::vector<ObjectID>> result =
|
||||
local_scheduler_wait(reinterpret_cast<PyLocalSchedulerClient *>(self)
|
||||
->local_scheduler_connection,
|
||||
object_ids, num_returns, timeout_ms,
|
||||
static_cast<bool>(wait_local));
|
||||
|
||||
// Convert result to py object.
|
||||
PyObject *py_found = PyList_New(static_cast<Py_ssize_t>(result.first.size()));
|
||||
for (uint i = 0; i < result.first.size(); ++i) {
|
||||
PyList_SetItem(py_found, i, PyObjectID_make(result.first[i]));
|
||||
}
|
||||
PyObject *py_remaining =
|
||||
PyList_New(static_cast<Py_ssize_t>(result.second.size()));
|
||||
for (uint i = 0; i < result.second.size(); ++i) {
|
||||
PyList_SetItem(py_remaining, i, PyObjectID_make(result.second[i]));
|
||||
}
|
||||
return Py_BuildValue("(OO)", py_found, py_remaining);
|
||||
}
|
||||
|
||||
static PyMethodDef PyLocalSchedulerClient_methods[] = {
|
||||
{"disconnect", (PyCFunction) PyLocalSchedulerClient_disconnect, METH_NOARGS,
|
||||
"Notify the local scheduler that this client is exiting gracefully."},
|
||||
@@ -201,6 +253,8 @@ static PyMethodDef PyLocalSchedulerClient_methods[] = {
|
||||
(PyCFunction) PyLocalSchedulerClient_get_actor_frontier, METH_VARARGS, ""},
|
||||
{"set_actor_frontier",
|
||||
(PyCFunction) PyLocalSchedulerClient_set_actor_frontier, METH_VARARGS, ""},
|
||||
{"wait", (PyCFunction) PyLocalSchedulerClient_wait, METH_VARARGS,
|
||||
"Wait for a list of objects to be created."},
|
||||
{NULL} /* Sentinel */
|
||||
};
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
#include "common_protocol.h"
|
||||
#include "format/local_scheduler_generated.h"
|
||||
#include "ray/raylet/format/node_manager_generated.h"
|
||||
|
||||
#include "common/io.h"
|
||||
#include "common/task.h"
|
||||
@@ -207,3 +208,41 @@ void local_scheduler_set_actor_frontier(LocalSchedulerConnection *conn,
|
||||
ray::local_scheduler::protocol::MessageType_SetActorFrontier,
|
||||
frontier.size(), const_cast<uint8_t *>(frontier.data()));
|
||||
}
|
||||
|
||||
std::pair<std::vector<ObjectID>, std::vector<ObjectID>> local_scheduler_wait(
|
||||
LocalSchedulerConnection *conn,
|
||||
const std::vector<ObjectID> &object_ids,
|
||||
int num_returns,
|
||||
int64_t timeout_milliseconds,
|
||||
bool wait_local) {
|
||||
// Write request.
|
||||
flatbuffers::FlatBufferBuilder fbb;
|
||||
auto message = ray::protocol::CreateWaitRequest(
|
||||
fbb, to_flatbuf(fbb, object_ids), num_returns, timeout_milliseconds,
|
||||
wait_local);
|
||||
fbb.Finish(message);
|
||||
write_message(conn->conn, ray::protocol::MessageType_WaitRequest,
|
||||
fbb.GetSize(), fbb.GetBufferPointer());
|
||||
// Read result.
|
||||
int64_t type;
|
||||
int64_t reply_size;
|
||||
uint8_t *reply;
|
||||
read_message(conn->conn, &type, &reply_size, &reply);
|
||||
RAY_CHECK(type == ray::protocol::MessageType_WaitReply);
|
||||
auto reply_message = flatbuffers::GetRoot<ray::protocol::WaitReply>(reply);
|
||||
// Convert result.
|
||||
std::pair<std::vector<ObjectID>, std::vector<ObjectID>> result;
|
||||
auto found = reply_message->found();
|
||||
for (uint i = 0; i < found->size(); i++) {
|
||||
ObjectID object_id = ObjectID::from_binary(found->Get(i)->str());
|
||||
result.first.push_back(object_id);
|
||||
}
|
||||
auto remaining = reply_message->remaining();
|
||||
for (uint i = 0; i < remaining->size(); i++) {
|
||||
ObjectID object_id = ObjectID::from_binary(remaining->Get(i)->str());
|
||||
result.second.push_back(object_id);
|
||||
}
|
||||
/* Free the original message from the local scheduler. */
|
||||
free(reply);
|
||||
return result;
|
||||
}
|
||||
|
||||
@@ -169,4 +169,22 @@ const std::vector<uint8_t> local_scheduler_get_actor_frontier(
|
||||
void local_scheduler_set_actor_frontier(LocalSchedulerConnection *conn,
|
||||
const std::vector<uint8_t> &frontier);
|
||||
|
||||
/// Wait for the given objects until timeout expires or num_return objects are
|
||||
/// found.
|
||||
///
|
||||
/// \param conn The connection information.
|
||||
/// \param object_ids The objects to wait for.
|
||||
/// \param num_returns The number of objects to wait for.
|
||||
/// \param timeout_milliseconds Duration, in milliseconds, to wait before
|
||||
/// returning.
|
||||
/// \param wait_local Whether to wait for objects to appear on this node.
|
||||
/// \return A pair with the first element containing the object ids that were
|
||||
/// found, and the second element the objects that were not found.
|
||||
std::pair<std::vector<ObjectID>, std::vector<ObjectID>> local_scheduler_wait(
|
||||
LocalSchedulerConnection *conn,
|
||||
const std::vector<ObjectID> &object_ids,
|
||||
int num_returns,
|
||||
int64_t timeout_milliseconds,
|
||||
bool wait_local);
|
||||
|
||||
#endif
|
||||
|
||||
Reference in New Issue
Block a user