From 7e1326c0f6e0d198d126f2e555aff94f4536b1a2 Mon Sep 17 00:00:00 2001 From: "Siyuan (Ryans) Zhuang" Date: Thu, 9 Jul 2020 08:59:37 -0700 Subject: [PATCH] [Core] Enhance common client connection (#9367) * enhance client connection * add write buffer async * read message * add test --- BUILD.bazel | 4 +- src/ray/common/client_connection.cc | 81 ++++++++++++++----- src/ray/common/client_connection.h | 37 ++++++++- .../test}/client_connection_test.cc | 76 ++++++++++++----- src/ray/raylet/raylet.cc | 4 +- src/ray/raylet/worker_pool_test.cc | 9 ++- 6 files changed, 163 insertions(+), 48 deletions(-) rename src/ray/{raylet => common/test}/client_connection_test.cc (73%) diff --git a/BUILD.bazel b/BUILD.bazel index 152039c93..037cff77b 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -925,10 +925,10 @@ cc_test( cc_test( name = "client_connection_test", - srcs = ["src/ray/raylet/client_connection_test.cc"], + srcs = ["src/ray/common/test/client_connection_test.cc"], copts = COPTS, deps = [ - ":raylet_lib", + ":ray_common", "@com_google_googletest//:gtest_main", ], ) diff --git a/src/ray/common/client_connection.cc b/src/ray/common/client_connection.cc index 3abc1908e..9a5210fb3 100644 --- a/src/ray/common/client_connection.cc +++ b/src/ray/common/client_connection.cc @@ -29,13 +29,12 @@ namespace ray { -std::shared_ptr ServerConnection::Create( - boost::asio::generic::stream_protocol::socket &&socket) { +std::shared_ptr ServerConnection::Create(local_stream_socket &&socket) { std::shared_ptr self(new ServerConnection(std::move(socket))); return self; } -ServerConnection::ServerConnection(boost::asio::generic::stream_protocol::socket &&socket) +ServerConnection::ServerConnection(local_stream_socket &&socket) : socket_(std::move(socket)), async_write_max_messages_(1), async_write_queue_(), @@ -73,6 +72,17 @@ Status ServerConnection::WriteBuffer( return ray::Status::OK(); } +void ServerConnection::WriteBufferAsync( + const std::vector &buffer, + const std::function &handler) { + // Wait for the message to be written. + boost::asio::async_write( + socket_, buffer, + [handler](const boost::system::error_code &ec, size_t bytes_transferred) { + handler(boost_to_ray_status(ec)); + }); +} + Status ServerConnection::ReadBuffer( const std::vector &buffer) { boost::system::error_code error; @@ -95,18 +105,54 @@ Status ServerConnection::ReadBuffer( return Status::OK(); } +void ServerConnection::ReadBufferAsync( + const std::vector &buffer, + const std::function &handler) { + // Wait for the message to be read. + boost::asio::async_read( + socket_, buffer, + [handler](const boost::system::error_code &ec, size_t bytes_transferred) { + handler(boost_to_ray_status(ec)); + }); +} + ray::Status ServerConnection::WriteMessage(int64_t type, int64_t length, const uint8_t *message) { sync_writes_ += 1; bytes_written_ += length; - std::vector message_buffers; auto write_cookie = RayConfig::instance().ray_cookie(); - message_buffers.push_back(boost::asio::buffer(&write_cookie, sizeof(write_cookie))); - message_buffers.push_back(boost::asio::buffer(&type, sizeof(type))); - message_buffers.push_back(boost::asio::buffer(&length, sizeof(length))); - message_buffers.push_back(boost::asio::buffer(message, length)); - return WriteBuffer(message_buffers); + return WriteBuffer({ + boost::asio::buffer(&write_cookie, sizeof(write_cookie)), + boost::asio::buffer(&type, sizeof(type)), + boost::asio::buffer(&length, sizeof(length)), + boost::asio::buffer(message, length), + }); +} + +Status ServerConnection::ReadMessage(int64_t type, std::vector *message) { + int64_t read_cookie, read_type, read_length; + // Wait for a message header from the client. The message header includes the + // protocol version, the message type, and the length of the message. + RAY_RETURN_NOT_OK(ReadBuffer({ + boost::asio::buffer(&read_cookie, sizeof(read_cookie)), + boost::asio::buffer(&read_type, sizeof(read_type)), + boost::asio::buffer(&read_length, sizeof(read_length)), + })); + if (read_cookie != RayConfig::instance().ray_cookie()) { + std::ostringstream ss; + ss << "Ray cookie mismatch for received message. " + << "Received cookie: " << read_cookie; + return Status::IOError(ss.str()); + } + if (type != read_type) { + std::ostringstream ss; + ss << "Connection corrupted. Expected message type: " << type + << ", receviced message type: " << read_type; + return Status::IOError(ss.str()); + } + message->resize(read_length); + return ReadBuffer({boost::asio::buffer(*message)}); } void ServerConnection::WriteMessageAsync( @@ -203,8 +249,7 @@ void ServerConnection::DoAsyncWrites() { std::shared_ptr ClientConnection::Create( ClientHandler &client_handler, MessageHandler &message_handler, - boost::asio::generic::stream_protocol::socket &&socket, - const std::string &debug_label, + local_stream_socket &&socket, const std::string &debug_label, const std::vector &message_type_enum_names, int64_t error_message_type) { std::shared_ptr self( new ClientConnection(message_handler, std::move(socket), debug_label, @@ -215,8 +260,7 @@ std::shared_ptr ClientConnection::Create( } ClientConnection::ClientConnection( - MessageHandler &message_handler, - boost::asio::generic::stream_protocol::socket &&socket, + MessageHandler &message_handler, local_stream_socket &&socket, const std::string &debug_label, const std::vector &message_type_enum_names, int64_t error_message_type) : ServerConnection(std::move(socket)), @@ -234,10 +278,11 @@ void ClientConnection::Register() { void ClientConnection::ProcessMessages() { // Wait for a message header from the client. The message header includes the // protocol version, the message type, and the length of the message. - std::vector header; - header.push_back(boost::asio::buffer(&read_cookie_, sizeof(read_cookie_))); - header.push_back(boost::asio::buffer(&read_type_, sizeof(read_type_))); - header.push_back(boost::asio::buffer(&read_length_, sizeof(read_length_))); + std::vector header{ + boost::asio::buffer(&read_cookie_, sizeof(read_cookie_)), + boost::asio::buffer(&read_type_, sizeof(read_type_)), + boost::asio::buffer(&read_length_, sizeof(read_length_)), + }; boost::asio::async_read( ServerConnection::socket_, header, boost::bind(&ClientConnection::ProcessMessageHeader, @@ -307,7 +352,7 @@ void ClientConnection::ProcessMessage(const boost::system::error_code &error) { } int64_t start_ms = current_time_ms(); - message_handler_(shared_ClientConnection_from_this(), read_type_, read_message_.data()); + message_handler_(shared_ClientConnection_from_this(), read_type_, read_message_); int64_t interval = current_time_ms() - start_ms; if (interval > RayConfig::instance().handler_warning_timeout_ms()) { std::string message_type; diff --git a/src/ray/common/client_connection.h b/src/ray/common/client_connection.h index 8f57325fa..fab43bfae 100644 --- a/src/ray/common/client_connection.h +++ b/src/ray/common/client_connection.h @@ -63,24 +63,57 @@ class ServerConnection : public std::enable_shared_from_this { void WriteMessageAsync(int64_t type, int64_t length, const uint8_t *message, const std::function &handler); + /// Read a message from the client. + /// + /// \param type The message type (e.g., a flatbuffer enum). + /// \param message A pointer to the message buffer. + /// \return Status. + Status ReadMessage(int64_t type, std::vector *message); + /// Write a buffer to this connection. /// /// \param buffer The buffer. /// \return Status. Status WriteBuffer(const std::vector &buffer); + /// Write a buffer to this connection asynchronously. + /// + /// \param buffer The buffer. + /// \param handler A callback to run on write completion. + /// \return Status. + void WriteBufferAsync(const std::vector &buffer, + const std::function &handler); + /// Read a buffer from this connection. /// /// \param buffer The buffer. /// \return Status. Status ReadBuffer(const std::vector &buffer); + /// Read a buffer from this connection asynchronously. + /// + /// \param buffer The buffer. + /// \param handler A callback to run on read completion. + /// \return Status. + void ReadBufferAsync(const std::vector &buffer, + const std::function &handler); + /// Shuts down socket for this connection. void Close() { boost::system::error_code ec; socket_.close(ec); } + /// Get the native handle of the socket. + int GetNativeHandle() { return socket_.native_handle(); } + + /// Set the blocking flag of the underlying socket. + Status SetNonBlocking(bool nonblocking) { + boost::system::error_code ec; + socket_.native_non_blocking(nonblocking, ec); + return boost_to_ray_status(ec); + } + std::string DebugString() const; protected: @@ -133,8 +166,8 @@ class ServerConnection : public std::enable_shared_from_this { class ClientConnection; using ClientHandler = std::function; -using MessageHandler = - std::function, int64_t, const uint8_t *)>; +using MessageHandler = std::function, int64_t, + const std::vector &)>; /// \typename ClientConnection /// diff --git a/src/ray/raylet/client_connection_test.cc b/src/ray/common/test/client_connection_test.cc similarity index 73% rename from src/ray/raylet/client_connection_test.cc rename to src/ray/common/test/client_connection_test.cc index 083f024bb..f1a22234b 100644 --- a/src/ray/raylet/client_connection_test.cc +++ b/src/ray/common/test/client_connection_test.cc @@ -12,11 +12,11 @@ // See the License for the specific language governing permissions and // limitations under the License. +#include +#include #include #include -#include -#include #include "gmock/gmock.h" #include "gtest/gtest.h" @@ -72,12 +72,12 @@ TEST_F(ClientConnectionTest, SimpleSyncWrite) { ClientHandler client_handler = [](ClientConnection &client) {}; - MessageHandler message_handler = [&arr, &num_messages]( - std::shared_ptr client, - int64_t message_type, const uint8_t *message) { - ASSERT_TRUE(!std::memcmp(arr, message, 5)); - num_messages += 1; - }; + MessageHandler message_handler = + [&arr, &num_messages](std::shared_ptr client, + int64_t message_type, const std::vector &message) { + ASSERT_TRUE(!std::memcmp(arr, message.data(), 5)); + num_messages += 1; + }; auto conn1 = ClientConnection::Create(client_handler, message_handler, std::move(in_), "conn1", {}, error_message_type_); @@ -102,19 +102,21 @@ TEST_F(ClientConnectionTest, SimpleAsyncWrite) { ClientHandler client_handler = [](ClientConnection &client) {}; MessageHandler noop_handler = [](std::shared_ptr client, - int64_t message_type, const uint8_t *message) {}; + int64_t message_type, + const std::vector &message) {}; std::shared_ptr reader = NULL; MessageHandler message_handler = [&msg1, &msg2, &msg3, &num_messages, &reader]( std::shared_ptr client, - int64_t message_type, const uint8_t *message) { + int64_t message_type, + const std::vector &message) { if (num_messages == 0) { - ASSERT_TRUE(!std::memcmp(msg1, message, 5)); + ASSERT_TRUE(!std::memcmp(msg1, message.data(), 5)); } else if (num_messages == 1) { - ASSERT_TRUE(!std::memcmp(msg2, message, 5)); + ASSERT_TRUE(!std::memcmp(msg2, message.data(), 5)); } else { - ASSERT_TRUE(!std::memcmp(msg3, message, 5)); + ASSERT_TRUE(!std::memcmp(msg3, message.data(), 5)); } num_messages += 1; if (num_messages < 3) { @@ -140,13 +142,44 @@ TEST_F(ClientConnectionTest, SimpleAsyncWrite) { ASSERT_EQ(num_messages, 3); } +TEST_F(ClientConnectionTest, SimpleSyncReadWriteMessage) { + auto writer = ServerConnection::Create(std::move(in_)); + auto reader = ServerConnection::Create(std::move(out_)); + + const std::vector write_buffer = {1, 2, 3, 4, 5}; + std::vector read_buffer; + + RAY_CHECK_OK(writer->WriteMessage(42, write_buffer.size(), write_buffer.data())); + RAY_CHECK_OK(reader->ReadMessage(42, &read_buffer)); + RAY_CHECK(write_buffer == read_buffer); +} + +TEST_F(ClientConnectionTest, SimpleAsyncReadWriteBuffers) { + auto writer = ServerConnection::Create(std::move(in_)); + auto reader = ServerConnection::Create(std::move(out_)); + + const std::vector write_buffer = {1, 2, 3, 4, 5}; + std::vector read_buffer = {0, 0, 0, 0, 0}; + + writer->WriteBufferAsync({boost::asio::buffer(write_buffer)}, + [](const ray::Status &status) { RAY_CHECK_OK(status); }); + + reader->ReadBufferAsync({boost::asio::buffer(read_buffer)}, + [&write_buffer, &read_buffer](const ray::Status &status) { + RAY_CHECK_OK(status); + RAY_CHECK(write_buffer == read_buffer); + }); + io_service_.run(); +} + TEST_F(ClientConnectionTest, SimpleAsyncError) { const uint8_t msg1[5] = {1, 2, 3, 4, 5}; ClientHandler client_handler = [](ClientConnection &client) {}; MessageHandler noop_handler = [](std::shared_ptr client, - int64_t message_type, const uint8_t *message) {}; + int64_t message_type, + const std::vector &message) {}; auto writer = ClientConnection::Create(client_handler, noop_handler, std::move(in_), "writer", {}, error_message_type_); @@ -166,7 +199,8 @@ TEST_F(ClientConnectionTest, CallbackWithSharedRefDoesNotLeakConnection) { ClientHandler client_handler = [](ClientConnection &client) {}; MessageHandler noop_handler = [](std::shared_ptr client, - int64_t message_type, const uint8_t *message) {}; + int64_t message_type, + const std::vector &message) {}; auto writer = ClientConnection::Create(client_handler, noop_handler, std::move(in_), "writer", {}, error_message_type_); @@ -186,12 +220,12 @@ TEST_F(ClientConnectionTest, ProcessBadMessage) { ClientHandler client_handler = [](ClientConnection &client) {}; - MessageHandler message_handler = [&arr, &num_messages]( - std::shared_ptr client, - int64_t message_type, const uint8_t *message) { - ASSERT_TRUE(!std::memcmp(arr, message, 5)); - num_messages += 1; - }; + MessageHandler message_handler = + [&arr, &num_messages](std::shared_ptr client, + int64_t message_type, const std::vector &message) { + ASSERT_TRUE(!std::memcmp(arr, message.data(), 5)); + num_messages += 1; + }; auto writer = ClientConnection::Create(client_handler, message_handler, std::move(in_), "writer", {}, error_message_type_); diff --git a/src/ray/raylet/raylet.cc b/src/ray/raylet/raylet.cc index 4e28b9c3f..4db107806 100644 --- a/src/ray/raylet/raylet.cc +++ b/src/ray/raylet/raylet.cc @@ -134,8 +134,8 @@ void Raylet::HandleAccept(const boost::system::error_code &error) { }; MessageHandler message_handler = [this](std::shared_ptr client, int64_t message_type, - const uint8_t *message) { - node_manager_.ProcessClientMessage(client, message_type, message); + const std::vector &message) { + node_manager_.ProcessClientMessage(client, message_type, message.data()); }; // Accept a new local client and dispatch it to the node manager. auto new_connection = ClientConnection::Create( diff --git a/src/ray/raylet/worker_pool_test.cc b/src/ray/raylet/worker_pool_test.cc index 4cb35b1d4..ad90ef80b 100644 --- a/src/ray/raylet/worker_pool_test.cc +++ b/src/ray/raylet/worker_pool_test.cc @@ -104,9 +104,11 @@ class WorkerPoolTest : public ::testing::Test { const Language &language = Language::PYTHON) { std::function client_handler = [this](ClientConnection &client) { HandleNewClient(client); }; - std::function, int64_t, const uint8_t *)> + std::function, int64_t, + const std::vector &)> message_handler = [this](std::shared_ptr client, - int64_t message_type, const uint8_t *message) { + int64_t message_type, + const std::vector &message) { HandleMessage(client, message_type, message); }; local_stream_socket socket(io_service_); @@ -162,7 +164,8 @@ class WorkerPoolTest : public ::testing::Test { private: void HandleNewClient(ClientConnection &){}; - void HandleMessage(std::shared_ptr, int64_t, const uint8_t *){}; + void HandleMessage(std::shared_ptr, int64_t, + const std::vector &){}; }; static inline TaskSpecification ExampleTaskSpec(