diff --git a/src/ray/common/client_connection.cc b/src/ray/common/client_connection.cc index 1ca93ba2a..dde14d680 100644 --- a/src/ray/common/client_connection.cc +++ b/src/ray/common/client_connection.cc @@ -186,9 +186,10 @@ void ServerConnection::DoAsyncWrites() { template std::shared_ptr> ClientConnection::Create( ClientHandler &client_handler, MessageHandler &message_handler, - boost::asio::basic_stream_socket &&socket, const std::string &debug_label) { - std::shared_ptr> self( - new ClientConnection(message_handler, std::move(socket), debug_label)); + boost::asio::basic_stream_socket &&socket, const std::string &debug_label, + int64_t error_message_type) { + std::shared_ptr> self(new ClientConnection( + message_handler, std::move(socket), debug_label, error_message_type)); // Let our manager process our new connection. client_handler(*self); return self; @@ -197,10 +198,12 @@ std::shared_ptr> ClientConnection::Create( template ClientConnection::ClientConnection(MessageHandler &message_handler, boost::asio::basic_stream_socket &&socket, - const std::string &debug_label) + const std::string &debug_label, + int64_t error_message_type) : ServerConnection(std::move(socket)), message_handler_(message_handler), - debug_label_(debug_label) {} + debug_label_(debug_label), + error_message_type_(error_message_type) {} template const ClientID &ClientConnection::GetClientId() { @@ -230,7 +233,7 @@ template void ClientConnection::ProcessMessageHeader(const boost::system::error_code &error) { if (error) { // If there was an error, disconnect the client. - read_type_ = static_cast(protocol::MessageType::DisconnectClient); + read_type_ = error_message_type_; read_length_ = 0; ProcessMessage(error); return; @@ -251,7 +254,7 @@ void ClientConnection::ProcessMessageHeader(const boost::system::error_code & template void ClientConnection::ProcessMessage(const boost::system::error_code &error) { if (error) { - read_type_ = static_cast(protocol::MessageType::DisconnectClient); + read_type_ = error_message_type_; } int64_t start_ms = current_time_ms(); diff --git a/src/ray/common/client_connection.h b/src/ray/common/client_connection.h index 7246c2b81..2aee25681 100644 --- a/src/ray/common/client_connection.h +++ b/src/ray/common/client_connection.h @@ -148,7 +148,8 @@ class ClientConnection : public ServerConnection { /// \return std::shared_ptr. static std::shared_ptr> Create( ClientHandler &new_client_handler, MessageHandler &message_handler, - boost::asio::basic_stream_socket &&socket, const std::string &debug_label); + boost::asio::basic_stream_socket &&socket, const std::string &debug_label, + int64_t error_message_type); std::shared_ptr> shared_ClientConnection_from_this() { return std::static_pointer_cast>(shared_from_this()); @@ -169,7 +170,7 @@ class ClientConnection : public ServerConnection { /// A private constructor for a node client connection. ClientConnection(MessageHandler &message_handler, boost::asio::basic_stream_socket &&socket, - const std::string &debug_label); + const std::string &debug_label, int64_t error_message_type); /// Process an error from the last operation, then process the message /// header from the client. void ProcessMessageHeader(const boost::system::error_code &error); @@ -183,6 +184,8 @@ class ClientConnection : public ServerConnection { MessageHandler message_handler_; /// A label used for debug messages. const std::string debug_label_; + /// The value for disconnect client message. + int64_t error_message_type_; /// Buffers for the current message being read from the client. int64_t read_version_; int64_t read_type_; diff --git a/src/ray/object_manager/format/object_manager.fbs b/src/ray/object_manager/format/object_manager.fbs index 4b684ff89..dbafa3744 100644 --- a/src/ray/object_manager/format/object_manager.fbs +++ b/src/ray/object_manager/format/object_manager.fbs @@ -27,6 +27,7 @@ table ObjectInfo { enum MessageType:int { ConnectClient = 1, + DisconnectClient, PushRequest, PullRequest, FreeRequest diff --git a/src/ray/object_manager/object_buffer_pool.cc b/src/ray/object_manager/object_buffer_pool.cc index 15e9eefac..f7d1f1651 100644 --- a/src/ray/object_manager/object_buffer_pool.cc +++ b/src/ray/object_manager/object_buffer_pool.cc @@ -150,7 +150,6 @@ void ObjectBufferPool::SealChunk(const ObjectID &object_id, const uint64_t chunk CreateChunkState::REFERENCED); create_buffer_state_[object_id].chunk_state[chunk_index] = CreateChunkState::SEALED; create_buffer_state_[object_id].num_seals_remaining--; - RAY_CHECK(create_buffer_state_[object_id].num_seals_remaining >= 0); RAY_LOG(DEBUG) << "SealChunk" << object_id << " " << create_buffer_state_[object_id].num_seals_remaining; if (create_buffer_state_[object_id].num_seals_remaining == 0) { diff --git a/src/ray/object_manager/object_manager.cc b/src/ray/object_manager/object_manager.cc index 6afc5f3dc..517d01c4a 100644 --- a/src/ray/object_manager/object_manager.cc +++ b/src/ray/object_manager/object_manager.cc @@ -707,25 +707,26 @@ void ObjectManager::ProcessNewClient(TcpClientConnection &conn) { void ObjectManager::ProcessClientMessage(std::shared_ptr &conn, int64_t message_type, const uint8_t *message) { - switch (message_type) { - case static_cast(object_manager_protocol::MessageType::PushRequest): { + auto message_type_value = + static_cast(message_type); + switch (message_type_value) { + case object_manager_protocol::MessageType::PushRequest: { ReceivePushRequest(conn, message); break; } - case static_cast(object_manager_protocol::MessageType::PullRequest): { + case object_manager_protocol::MessageType::PullRequest: { ReceivePullRequest(conn, message); break; } - case static_cast(object_manager_protocol::MessageType::ConnectClient): { + case object_manager_protocol::MessageType::ConnectClient: { ConnectClient(conn, message); break; } - case static_cast(object_manager_protocol::MessageType::FreeRequest): { + case object_manager_protocol::MessageType::FreeRequest: { ReceiveFreeRequest(conn, message); break; } - case static_cast(protocol::MessageType::DisconnectClient): { - // TODO(hme): Disconnect without depending on the node manager protocol. + case object_manager_protocol::MessageType::DisconnectClient: { DisconnectClient(conn, message); break; } diff --git a/src/ray/object_manager/test/object_manager_stress_test.cc b/src/ray/object_manager/test/object_manager_stress_test.cc index 84e27e5ed..b8b35ad62 100644 --- a/src/ray/object_manager/test/object_manager_stress_test.cc +++ b/src/ray/object_manager/test/object_manager_stress_test.cc @@ -74,9 +74,10 @@ class MockServer { object_manager_.ProcessClientMessage(client, message_type, message); }; // Accept a new local client and dispatch it to the node manager. - auto new_connection = - TcpClientConnection::Create(client_handler, message_handler, - std::move(object_manager_socket_), "object manager"); + auto new_connection = TcpClientConnection::Create( + client_handler, message_handler, std::move(object_manager_socket_), + "object manager", + static_cast(object_manager::protocol::MessageType::DisconnectClient)); DoAcceptObjectManager(); } diff --git a/src/ray/object_manager/test/object_manager_test.cc b/src/ray/object_manager/test/object_manager_test.cc index 4c108f2d3..12ad2c52c 100644 --- a/src/ray/object_manager/test/object_manager_test.cc +++ b/src/ray/object_manager/test/object_manager_test.cc @@ -65,9 +65,10 @@ class MockServer { object_manager_.ProcessClientMessage(client, message_type, message); }; // Accept a new local client and dispatch it to the node manager. - auto new_connection = - TcpClientConnection::Create(client_handler, message_handler, - std::move(object_manager_socket_), "object manager"); + auto new_connection = TcpClientConnection::Create( + client_handler, message_handler, std::move(object_manager_socket_), + "object manager", + static_cast(object_manager::protocol::MessageType::DisconnectClient)); DoAcceptObjectManager(); } diff --git a/src/ray/raylet/client_connection_test.cc b/src/ray/raylet/client_connection_test.cc index a68a6535c..5623beb90 100644 --- a/src/ray/raylet/client_connection_test.cc +++ b/src/ray/raylet/client_connection_test.cc @@ -13,7 +13,8 @@ namespace raylet { class ClientConnectionTest : public ::testing::Test { public: - ClientConnectionTest() : io_service_(), in_(io_service_), out_(io_service_) { + ClientConnectionTest() + : io_service_(), in_(io_service_), out_(io_service_), error_message_type_(1) { boost::asio::local::connect_pair(in_, out_); } @@ -21,6 +22,7 @@ class ClientConnectionTest : public ::testing::Test { boost::asio::io_service io_service_; boost::asio::local::stream_protocol::socket in_; boost::asio::local::stream_protocol::socket out_; + int64_t error_message_type_; }; TEST_F(ClientConnectionTest, SimpleSyncWrite) { @@ -37,11 +39,11 @@ TEST_F(ClientConnectionTest, SimpleSyncWrite) { num_messages += 1; }; - auto conn1 = LocalClientConnection::Create(client_handler, message_handler, - std::move(in_), "conn1"); + auto conn1 = LocalClientConnection::Create( + client_handler, message_handler, std::move(in_), "conn1", error_message_type_); - auto conn2 = LocalClientConnection::Create(client_handler, message_handler, - std::move(out_), "conn2"); + auto conn2 = LocalClientConnection::Create( + client_handler, message_handler, std::move(out_), "conn2", error_message_type_); RAY_CHECK_OK(conn1->WriteMessage(0, 5, arr)); RAY_CHECK_OK(conn2->WriteMessage(0, 5, arr)); @@ -83,11 +85,11 @@ TEST_F(ClientConnectionTest, SimpleAsyncWrite) { } }; - auto writer = LocalClientConnection::Create(client_handler, noop_handler, - std::move(in_), "writer"); + auto writer = LocalClientConnection::Create( + client_handler, noop_handler, std::move(in_), "writer", error_message_type_); reader = LocalClientConnection::Create(client_handler, message_handler, std::move(out_), - "reader"); + "reader", error_message_type_); std::function callback = [](const ray::Status &status) { RAY_CHECK_OK(status); @@ -111,8 +113,8 @@ TEST_F(ClientConnectionTest, SimpleAsyncError) { std::shared_ptr client, int64_t message_type, const uint8_t *message) {}; - auto writer = LocalClientConnection::Create(client_handler, noop_handler, - std::move(in_), "writer"); + auto writer = LocalClientConnection::Create( + client_handler, noop_handler, std::move(in_), "writer", error_message_type_); std::function callback = [](const ray::Status &status) { ASSERT_TRUE(!status.ok()); @@ -133,8 +135,8 @@ TEST_F(ClientConnectionTest, CallbackWithSharedRefDoesNotLeakConnection) { std::shared_ptr client, int64_t message_type, const uint8_t *message) {}; - auto writer = LocalClientConnection::Create(client_handler, noop_handler, - std::move(in_), "writer"); + auto writer = LocalClientConnection::Create( + client_handler, noop_handler, std::move(in_), "writer", error_message_type_); std::function callback = [writer](const ray::Status &status) { diff --git a/src/ray/raylet/raylet.cc b/src/ray/raylet/raylet.cc index 39028ce7f..23f0dba6c 100644 --- a/src/ray/raylet/raylet.cc +++ b/src/ray/raylet/raylet.cc @@ -101,7 +101,8 @@ void Raylet::HandleAcceptNodeManager(const boost::system::error_code &error) { }; // Accept a new TCP client and dispatch it to the node manager. auto new_connection = TcpClientConnection::Create( - client_handler, message_handler, std::move(node_manager_socket_), "node manager"); + client_handler, message_handler, std::move(node_manager_socket_), "node manager", + static_cast(protocol::MessageType::DisconnectClient)); } // We're ready to accept another client. DoAcceptNodeManager(); @@ -122,9 +123,10 @@ void Raylet::HandleAcceptObjectManager(const boost::system::error_code &error) { object_manager_.ProcessClientMessage(client, message_type, message); }; // Accept a new TCP client and dispatch it to the node manager. - auto new_connection = - TcpClientConnection::Create(client_handler, message_handler, - std::move(object_manager_socket_), "object manager"); + auto new_connection = TcpClientConnection::Create( + client_handler, message_handler, std::move(object_manager_socket_), + "object manager", + static_cast(object_manager::protocol::MessageType::DisconnectClient)); DoAcceptObjectManager(); } @@ -144,8 +146,9 @@ void Raylet::HandleAccept(const boost::system::error_code &error) { node_manager_.ProcessClientMessage(client, message_type, message); }; // Accept a new local client and dispatch it to the node manager. - auto new_connection = LocalClientConnection::Create(client_handler, message_handler, - std::move(socket_), "worker"); + auto new_connection = LocalClientConnection::Create( + client_handler, message_handler, std::move(socket_), "worker", + static_cast(protocol::MessageType::DisconnectClient)); } // We're ready to accept another client. DoAccept(); diff --git a/src/ray/raylet/worker_pool_test.cc b/src/ray/raylet/worker_pool_test.cc index 9b228457b..d235d0244 100644 --- a/src/ray/raylet/worker_pool_test.cc +++ b/src/ray/raylet/worker_pool_test.cc @@ -34,7 +34,7 @@ class WorkerPoolMock : public WorkerPool { class WorkerPoolTest : public ::testing::Test { public: - WorkerPoolTest() : worker_pool_(), io_service_() {} + WorkerPoolTest() : worker_pool_(), io_service_(), error_message_type_(1) {} std::shared_ptr CreateWorker(pid_t pid, const Language &language = Language::PYTHON) { @@ -46,14 +46,16 @@ class WorkerPoolTest : public ::testing::Test { HandleMessage(client, message_type, message); }; boost::asio::local::stream_protocol::socket socket(io_service_); - auto client = LocalClientConnection::Create(client_handler, message_handler, - std::move(socket), "worker"); + auto client = + LocalClientConnection::Create(client_handler, message_handler, std::move(socket), + "worker", error_message_type_); return std::shared_ptr(new Worker(pid, language, client)); } protected: WorkerPoolMock worker_pool_; boost::asio::io_service io_service_; + int64_t error_message_type_; private: void HandleNewClient(LocalClientConnection &){};