diff --git a/src/ipc.cc b/src/ipc.cc index e5b525f79..22f941aea 100644 --- a/src/ipc.cc +++ b/src/ipc.cc @@ -1,5 +1,16 @@ #include "ipc.h" +#include + +#if defined(WIN32) || defined(_WIN32) +#include +#else +#include +#include +#endif + +#include "ray/ray.h" + using namespace arrow; ObjHandle::ObjHandle(SegmentId segmentid, size_t size, IpcPointer ipcpointer, size_t metadata_offset) @@ -26,6 +37,101 @@ int64_t BufferMemorySource::Size() const { return size_; } +MessageQueue<>::MessageQueue() : handle_(-1) { } + +MessageQueue<>::~MessageQueue() { close(); } + +MessageQueue<>::MessageQueue(MessageQueue&& other) { + handle_ = other.handle_; + other.handle_ = -1; +} + +MessageQueue<>& MessageQueue<>::operator=(MessageQueue<>&& other) { + close(); + handle_ = other.handle_; + other.handle_ = -1; + return *this; +} + +bool MessageQueue<>::connect(const std::string& name, bool create, size_t buffer_size) { + std::string name_translated = "ray-{BC200A09-2465-431D-AEC7-2F8530B04535}-" + name; +#if defined(WIN32) || defined(_WIN32) + name_translated.insert(0, "\\\\.\\pipe\\"); + std::replace(name_translated.begin(), name_translated.end(), '/', '\\'); + if (create) { + handle_ = reinterpret_cast(CreateNamedPipeA(name_translated.c_str(), (create ? FILE_FLAG_FIRST_PIPE_INSTANCE : 0) | PIPE_ACCESS_DUPLEX, PIPE_TYPE_MESSAGE | PIPE_READMODE_MESSAGE | PIPE_WAIT | PIPE_REJECT_REMOTE_CLIENTS, 1, static_cast(buffer_size), static_cast(buffer_size), INFINITE, NULL)); + } else { + handle_ = reinterpret_cast(CreateFileA(name_translated.c_str(), GENERIC_ALL, FILE_SHARE_READ | FILE_SHARE_WRITE, NULL, OPEN_EXISTING, 0, NULL)); + } +#else + name_translated.insert(0, "/tmp/"); + if (!create || mkfifo(name_translated.c_str(), S_IWUSR | S_IRUSR | S_IRGRP | S_IROTH) == 0 || errno == EEXIST) { + handle_ = open(name_translated.c_str(), O_RDWR); + if (handle_ == -1) { + unlink(name_translated.c_str()); + } + } +#endif + return handle_ != -1; +} + +bool MessageQueue<>::connected() { + return handle_ != -1; +} + +void MessageQueue<>::close() { + if (connected()) { +#if defined(WIN32) || defined(_WIN32) + CloseHandle(reinterpret_cast(handle_)); +#else + ::close(handle_); +#endif + handle_ = -1; + } +} + +bool MessageQueue<>::send(const unsigned char* object, size_t size) { + while (size > 0) { +#if defined(WIN32) || defined(_WIN32) + DWORD transmitted; + if (!WriteFile(reinterpret_cast(handle_), object, static_cast(size), &transmitted, NULL)) { + RAY_LOG(RAY_INFO, "GetLastError() == " << GetLastError()); + break; + } +#else + ssize_t transmitted = write(handle_, object, size); + if (transmitted < 0) { + RAY_LOG(RAY_INFO, "errno == " << errno); + break; + } +#endif + size -= static_cast(transmitted); + object += static_cast(transmitted); + } + return size == 0; +} + +bool MessageQueue<>::receive(unsigned char* object, size_t size) { + while (size > 0) { +#if defined(WIN32) || defined(_WIN32) + DWORD transmitted; + if (!ReadFile(reinterpret_cast(handle_), object, static_cast(size), &transmitted, NULL)) { + RAY_LOG(RAY_INFO, "GetLastError() == " << GetLastError()); + break; + } +#else + ssize_t transmitted = read(handle_, object, size); + if (transmitted < 0) { + RAY_LOG(RAY_INFO, "errno == " << errno); + break; + } +#endif + size -= static_cast(transmitted); + object += static_cast(transmitted); + } + return size == 0; +} + MemorySegmentPool::MemorySegmentPool(ObjStoreId objstoreid, bool create) : objstoreid_(objstoreid), create_mode_(create) { } // creates a memory segment if it is not already there; if the pool is in create mode, diff --git a/src/ipc.h b/src/ipc.h index aaf73603f..877c6e27c 100644 --- a/src/ipc.h +++ b/src/ipc.h @@ -5,7 +5,6 @@ #include #include -#include #include #include @@ -18,62 +17,38 @@ namespace bip = boost::interprocess; // Message Queues: Exchanging objects of type T between processes on a node -template -class MessageQueue { -public: - MessageQueue() {}; +template +class MessageQueue; - ~MessageQueue() { - bip::message_queue::remove(name_.c_str()); - } - - MessageQueue(MessageQueue&& other) noexcept - : name_(std::move(other.name_)), - queue_(std::move(other.queue_)) - { } - - bool connect(const std::string& name, bool create) { - name_ = name; - try { - if (create) { - bip::message_queue::remove(name.c_str()); // remove queue if it has not been properly removed from last run - queue_ = std::unique_ptr(new bip::message_queue(bip::create_only, name.c_str(), 100, sizeof(T))); - } else { - queue_ = std::unique_ptr(new bip::message_queue(bip::open_only, name.c_str())); - } - } catch(bip::interprocess_exception &ex) { - RAY_CHECK(false, "boost::interprocess exception: " << ex.what()); - } - return true; - }; - - bool connected() { - return queue_ != NULL; - } - - bool send(const T* object) { - try { - queue_->send(object, sizeof(T), 0); - } catch(bip::interprocess_exception &ex) { - RAY_CHECK(false, "boost::interprocess exception: " << ex.what()); - } - return true; - }; - - bool receive(T* object) { - unsigned int priority; - bip::message_queue::size_type recvd_size; - try { - queue_->receive(object, sizeof(T), recvd_size, priority); - } catch(bip::interprocess_exception &ex) { - RAY_CHECK(false, "boost::interprocess exception: " << ex.what()); - } - return true; - } +template<> +class MessageQueue<> { +protected: + bool connect(const std::string& name, bool create, size_t buffer_size); + bool connected(); + void close(); + ~MessageQueue(); + MessageQueue(); + MessageQueue(MessageQueue&& other); + MessageQueue& operator=(MessageQueue&& other); + bool send(const unsigned char* object, size_t size);; + bool receive(unsigned char* object, size_t size); private: - std::string name_; - std::unique_ptr queue_; +#if defined(WIN32) || defined(_WIN32) + int handle_; +#else + intptr_t handle_; +#endif +}; + +template +class MessageQueue : public MessageQueue<> { +public: + using MessageQueue<>::connected; + using MessageQueue<>::close; + bool connect(const std::string& name, bool create) { return MessageQueue<>::connect(name, create, sizeof(T)); } + bool send(const T* object) { return MessageQueue<>::send(reinterpret_cast(object), sizeof(*object)); } + bool receive(T* object) { return MessageQueue<>::receive(reinterpret_cast(object), sizeof(*object)); } }; // Object Queues diff --git a/src/objstore.cc b/src/objstore.cc index 0fd0ad882..42334e2e6 100644 --- a/src/objstore.cc +++ b/src/objstore.cc @@ -41,7 +41,7 @@ void ObjStoreService::get_data_from(ObjRef objref, ObjStore::Stub& stub) { ObjStoreService::ObjStoreService(const std::string& objstore_address, std::shared_ptr scheduler_channel) : scheduler_stub_(Scheduler::NewStub(scheduler_channel)), objstore_address_(objstore_address) { - recv_queue_.connect(std::string("queue:") + objstore_address + std::string(":obj"), true); + RAY_CHECK(recv_queue_.connect(std::string("queue:") + objstore_address + std::string(":obj"), true), "error connecting recv_queue_"); ClientContext context; RegisterObjStoreRequest request; request.set_objstore_address(objstore_address); @@ -150,7 +150,7 @@ Status ObjStoreService::NotifyAlias(ServerContext* context, const NotifyAliasReq ObjRequest done_request; done_request.type = ObjRequestType::ALIAS_DONE; done_request.objref = alias_objref; - recv_queue_.send(&done_request); + RAY_CHECK(recv_queue_.send(&done_request), "error sending over IPC"); return Status::OK; } @@ -195,7 +195,7 @@ void ObjStoreService::process_worker_request(const ObjRequest request) { } if (!send_queues_[request.workerid].connected()) { std::string queue_name = std::string("queue:") + objstore_address_ + std::string(":worker:") + std::to_string(request.workerid) + std::string(":obj"); - send_queues_[request.workerid].connect(queue_name, false); + RAY_CHECK(send_queues_[request.workerid].connect(queue_name, false), "error connecting receive_queue_"); } { std::lock_guard memory_lock(memory_lock_); @@ -206,7 +206,7 @@ void ObjStoreService::process_worker_request(const ObjRequest request) { switch (request.type) { case ObjRequestType::ALLOC: { ObjHandle handle = alloc(request.objref, request.size); // This method acquires memory_lock_ - send_queues_[request.workerid].send(&handle); + RAY_CHECK(send_queues_[request.workerid].send(&handle), "error sending over IPC"); } break; case ObjRequestType::GET: { @@ -214,7 +214,7 @@ void ObjStoreService::process_worker_request(const ObjRequest request) { std::pair& item = memory_[request.objref]; if (item.second == MemoryStatusType::READY) { RAY_LOG(RAY_DEBUG, "Responding to GET request: returning objref " << request.objref); - send_queues_[request.workerid].send(&item.first); + RAY_CHECK(send_queues_[request.workerid].send(&item.first), "error sending over IPC"); } else if (item.second == MemoryStatusType::NOT_READY || item.second == MemoryStatusType::NOT_PRESENT || item.second == MemoryStatusType::PRE_ALLOCED) { std::lock_guard lock(get_queue_lock_); get_queue_.push_back(std::make_pair(request.workerid, request.objref)); @@ -237,7 +237,7 @@ void ObjStoreService::process_requests() { // TODO(rkn): Should memory_lock_ be used in this method? ObjRequest request; while (true) { - recv_queue_.receive(&request); + RAY_CHECK(recv_queue_.receive(&request), "error receiving over IPC"); switch (request.type) { case ObjRequestType::ALLOC: { RAY_LOG(RAY_VERBOSE, "Request (worker " << request.workerid << " to objstore " << objstoreid_ << "): Allocate object with objref " << request.objref << " and size " << request.size); @@ -271,7 +271,7 @@ void ObjStoreService::process_gets_for_objref(ObjRef objref) { for (size_t i = 0; i < get_queue_.size(); ++i) { if (get_queue_[i].second == objref) { ObjHandle& elem = memory_[objref].first; - send_queues_[get_queue_[i].first].send(&item.first); + RAY_CHECK(send_queues_[get_queue_[i].first].send(&item.first), "error sending over IPC"); // Remove the get task from the queue std::swap(get_queue_[i], get_queue_[get_queue_.size() - 1]); get_queue_.pop_back(); diff --git a/src/worker.cc b/src/worker.cc index 7dffbea22..203ba0412 100644 --- a/src/worker.cc +++ b/src/worker.cc @@ -13,27 +13,27 @@ extern "C" { inline WorkerServiceImpl::WorkerServiceImpl(const std::string& worker_address) : worker_address_(worker_address) { - send_queue_.connect(worker_address_, false); + RAY_CHECK(send_queue_.connect(worker_address_, false), "error connecting send_queue_"); } Status WorkerServiceImpl::ExecuteTask(ServerContext* context, const ExecuteTaskRequest* request, ExecuteTaskReply* reply) { task_ = std::unique_ptr(new Task(request->task())); // Copy task RAY_LOG(RAY_INFO, "invoked task " << request->task().name()); WorkerMessage message = { &task_ }; - send_queue_.send(&message); + RAY_CHECK(send_queue_.send(&message), "error sending over IPC"); return Status::OK; } Status WorkerServiceImpl::Die(ServerContext* context, const DieRequest* request, DieReply* reply) { WorkerMessage message = { NULL }; - send_queue_.send(&message); + RAY_CHECK(send_queue_.send(&message), "error sending over IPC"); return Status::OK; } Worker::Worker(const std::string& worker_address, std::shared_ptr scheduler_channel, std::shared_ptr objstore_channel) : worker_address_(worker_address), scheduler_stub_(Scheduler::NewStub(scheduler_channel)) { - receive_queue_.connect(worker_address_, true); + RAY_CHECK(receive_queue_.connect(worker_address_, true), "error connecting receive_queue_"); connected_ = true; } @@ -72,9 +72,8 @@ void Worker::register_worker(const std::string& worker_address, const std::strin workerid_ = reply.workerid(); objstoreid_ = reply.objstoreid(); segmentpool_ = std::make_shared(objstoreid_, false); - request_obj_queue_.connect(std::string("queue:") + objstore_address + std::string(":obj"), false); - std::string queue_name = std::string("queue:") + objstore_address + std::string(":worker:") + std::to_string(workerid_) + std::string(":obj"); - receive_obj_queue_.connect(queue_name, true); + RAY_CHECK(request_obj_queue_.connect(std::string("queue:") + objstore_address + std::string(":obj"), false), "error connecting request_obj_queue_"); + RAY_CHECK(receive_obj_queue_.connect(std::string("queue:") + objstore_address + std::string(":worker:") + std::to_string(workerid_) + std::string(":obj"), true), "error connecting receive_obj_queue_"); return; } @@ -107,9 +106,9 @@ slice Worker::get_object(ObjRef objref) { request.workerid = workerid_; request.type = ObjRequestType::GET; request.objref = objref; - request_obj_queue_.send(&request); + RAY_CHECK(request_obj_queue_.send(&request), "error sending over IPC"); ObjHandle result; - receive_obj_queue_.receive(&result); + RAY_CHECK(receive_obj_queue_.receive(&result), "error receiving over IPC"); slice slice; slice.data = segmentpool_->get_address(result); slice.len = result.size(); @@ -128,13 +127,13 @@ void Worker::put_object(ObjRef objref, const Obj* obj, std::vector &cont request.type = ObjRequestType::ALLOC; request.objref = objref; request.size = data.size(); - request_obj_queue_.send(&request); + RAY_CHECK(request_obj_queue_.send(&request), "error sending over IPC"); if (contained_objrefs.size() > 0) { RAY_LOG(RAY_REFCOUNT, "In put_object, calling increment_reference_count for contained objrefs"); increment_reference_count(contained_objrefs); // Notify the scheduler that some object references are serialized in the objstore. } ObjHandle result; - receive_obj_queue_.receive(&result); + RAY_CHECK(receive_obj_queue_.receive(&result), "error receiving over IPC"); uint8_t* target = segmentpool_->get_address(result); std::memcpy(target, &data[0], data.size()); // We immediately unmap here; if the object is going to be accessed again, it will be mapped again; @@ -142,7 +141,7 @@ void Worker::put_object(ObjRef objref, const Obj* obj, std::vector &cont segmentpool_->unmap_segment(result.segmentid()); request.type = ObjRequestType::WORKER_DONE; request.metadata_offset = 0; - request_obj_queue_.send(&request); + RAY_CHECK(request_obj_queue_.send(&request), "error sending over IPC"); // Notify the scheduler about the objrefs that we are serializing in the objstore. AddContainedObjRefsRequest contained_objrefs_request; @@ -176,9 +175,9 @@ PyObject* Worker::put_arrow(ObjRef objref, PyObject* value) { request.type = ObjRequestType::ALLOC; request.objref = objref; request.size = size; - request_obj_queue_.send(&request); + RAY_CHECK(request_obj_queue_.send(&request), "error sending over IPC"); ObjHandle result; - receive_obj_queue_.receive(&result); + RAY_CHECK(receive_obj_queue_.receive(&result), "error receiving over IPC"); int64_t metadata_offset; uint8_t* address = segmentpool_->get_address(result); auto source = std::make_shared(address, size); @@ -188,7 +187,7 @@ PyObject* Worker::put_arrow(ObjRef objref, PyObject* value) { segmentpool_->unmap_segment(result.segmentid()); request.type = ObjRequestType::WORKER_DONE; request.metadata_offset = metadata_offset; - request_obj_queue_.send(&request); + RAY_CHECK(request_obj_queue_.send(&request), "error sending over IPC"); Py_RETURN_NONE; } @@ -200,9 +199,9 @@ PyObject* Worker::get_arrow(ObjRef objref, SegmentId& segmentid) { request.workerid = workerid_; request.type = ObjRequestType::GET; request.objref = objref; - request_obj_queue_.send(&request); + RAY_CHECK(request_obj_queue_.send(&request), "error sending over IPC"); ObjHandle result; - receive_obj_queue_.receive(&result); + RAY_CHECK(receive_obj_queue_.receive(&result), "error receiving over IPC"); uint8_t* address = segmentpool_->get_address(result); auto source = std::make_shared(address, result.size()); segmentid = result.segmentid(); @@ -219,7 +218,7 @@ bool Worker::is_arrow(ObjRef objref) { request.objref = objref; request_obj_queue_.send(&request); ObjHandle result; - receive_obj_queue_.receive(&result); + RAY_CHECK(receive_obj_queue_.receive(&result), "error receiving over IPC"); return result.metadata_offset() != 0; } @@ -288,7 +287,7 @@ void Worker::register_function(const std::string& name, size_t num_return_vals) std::unique_ptr Worker::receive_next_task() { WorkerMessage message; - receive_queue_.receive(&message); + RAY_CHECK(receive_queue_.receive(&message), "error receiving over IPC"); return message.task ? std::move(*message.task) : std::unique_ptr(); }