From 7218ff2091cf0ab8ec964398164300298fe15e55 Mon Sep 17 00:00:00 2001 From: Sebastian-Keith Date: Wed, 1 Apr 2026 17:25:29 +0800 Subject: [PATCH] [Fix]: Strengthen error handling logics in cm / ds / client modules - Make ds exits gracefully when disconnection with CM - Fix potential r/w race issue in client of ds connection map - Clean up stale ds connection entry when calling update_all_route table(sync with CM) - Fix no valid return code issue of delete sync api - Some minor bugfixs - Add more new UT cases on changed codes --- CMakeLists.txt | 14 +- src/client/CMakeLists.txt | 9 +- src/client/clnt_flags.cc | 1 + src/client/clnt_messenger.cc | 403 +++++++++++------- src/client/clnt_messenger.h | 60 ++- src/cluster_manager/cm_hb_monitor.cc | 43 +- src/cluster_manager/cm_hb_monitor.h | 17 +- src/cluster_manager/cm_node_manager.h | 24 +- src/cluster_manager/cm_rpc_handler.cc | 139 ++++-- src/cluster_manager/cm_service.cc | 58 +-- src/cluster_manager/cm_service.h | 10 +- src/cluster_manager/cm_shard_manager.cc | 100 +++-- src/cluster_manager/cm_shard_manager.h | 13 +- src/data_server/ds_flags.cc | 7 +- src/data_server/kv_cache_pool.h | 30 +- src/data_server/kv_rpc_service.cc | 151 ++++--- src/data_server/kv_rpc_service.h | 32 +- tests/client/test_clnt_messenger.cc | 342 ++++++++++++++- tests/cluster_manager/test_cm_hb_monitor.cc | 151 ++++--- tests/cluster_manager/test_cm_rebalance.cc | 226 +++++++++- tests/cluster_manager/test_cm_service.cc | 324 ++++++++++++-- .../cluster_manager/test_cm_shard_manager.cc | 92 +++- tests/data_server/test_ds_kv_service.cc | 110 ++++- 23 files changed, 1811 insertions(+), 545 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index 60912f4..b79dc0f 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -126,7 +126,18 @@ include_directories(${PROJECT_SOURCE_DIR}/src ${PROJECT_SOURCE_DIR}/third_party/sict/include/ ) -set(CMAKE_PREFIX_PATH "${CMAKE_PREFIX_PATH};${PROJECT_SOURCE_DIR}/third_party/fast_float/include/") +set(SIMM_FASTFLOAT_ROOT "/tmp/simm-fast_float") +set(SIMM_FASTFLOAT_INCLUDE_DIR "${SIMM_FASTFLOAT_ROOT}/include") +file(MAKE_DIRECTORY "${SIMM_FASTFLOAT_ROOT}") +if(NOT EXISTS "${SIMM_FASTFLOAT_INCLUDE_DIR}") + execute_process( + COMMAND ${CMAKE_COMMAND} -E create_symlink + "${PROJECT_SOURCE_DIR}/third_party/fast_float/include" + "${SIMM_FASTFLOAT_INCLUDE_DIR}" + ) +endif() +set(FASTFLOAT_INCLUDE_DIR "${SIMM_FASTFLOAT_INCLUDE_DIR}" CACHE PATH "fast_float include path" FORCE) +set(CMAKE_INCLUDE_PATH "${CMAKE_INCLUDE_PATH};${SIMM_FASTFLOAT_INCLUDE_DIR}") # set GNU SOURCE to fix folly dependency add_compile_definitions(_GNU_SOURCE) @@ -150,6 +161,7 @@ add_subdirectory(tools) if(ENABLE_TESTS) message(STATUS "Building tests ...") add_compile_definitions(SIMM_UNIT_TEST) + include_directories(${PROJECT_SOURCE_DIR}/third_party/gtest/googletest/include) add_subdirectory(tests) endif() diff --git a/src/client/CMakeLists.txt b/src/client/CMakeLists.txt index aa14cd6..dad338c 100644 --- a/src/client/CMakeLists.txt +++ b/src/client/CMakeLists.txt @@ -13,7 +13,14 @@ file(GLOB MODULE_SRC ) add_library(${MODULE_NAME}_static STATIC ${MODULE_SRC}) -# PRIVATE means that the libraries are only used by this target and +if(ENABLE_TESTS) + target_compile_definitions(${MODULE_NAME}_static PRIVATE SIMM_UNIT_TEST) + target_include_directories(${MODULE_NAME}_static PRIVATE + ${PROJECT_SOURCE_DIR}/third_party/gtest/googletest/include + ) +endif() + +# PRIVATE means that the libraries are only used by this target and # not propagated to targets that link against this one. # # **Caution** about sict, it maybe loaded by shared library at runtime, so diff --git a/src/client/clnt_flags.cc b/src/client/clnt_flags.cc index c2d46a5..dcba4da 100644 --- a/src/client/clnt_flags.cc +++ b/src/client/clnt_flags.cc @@ -4,6 +4,7 @@ DEFINE_bool(clnt_use_k8s, true, "simm client in K8S env"); DEFINE_uint32(clnt_thread_pool_size, 10, "simm client thread pool size"); DEFINE_uint32(clnt_cm_addr_check_interval_inSecs, 10, "simm client backgroud thread(check cm address update) trigger interval in seconds, default is 10s"); +DEFINE_uint32(clnt_syncreq_retry_count, 2, "simm client sync requests retry count"); DEFINE_int32(clnt_sync_req_timeout_ms, 1000, "simm client sync request timeout in milliseconds, default is 1s"); DEFINE_int32(clnt_async_req_timeout_ms, 3000, "simm client sync request timeout in milliseconds, default is 3s"); DEFINE_string(clnt_log_file, "/var/log/simm/simm_clnt.log", "simm client log file path & name"); diff --git a/src/client/clnt_messenger.cc b/src/client/clnt_messenger.cc index 098be45..80fb7d3 100644 --- a/src/client/clnt_messenger.cc +++ b/src/client/clnt_messenger.cc @@ -1,14 +1,14 @@ +#include #include +#include #include +#include #include #include #include #include #include -#include -#include -#include -#include +#include #include @@ -18,19 +18,17 @@ #include "proto/cm_clnt_rpcs.pb.h" #include "proto/ds_clnt_rpcs.pb.h" +#include "clnt_messenger.h" +#include "cluster_manager/cm_rpc_handler.h" #include "common/base/assert.h" +#include "common/context/context.h" #include "common/errcode/errcode_def.h" #include "common/hashkit/hashkit.h" #include "common/logging/logging.h" -#include "common/utils/k8s_util.h" -#include "common/context/context.h" #include "common/trace/trace.h" -#include "cluster_manager/cm_rpc_handler.h" +#include "common/utils/k8s_util.h" #include "data_server/kv_rpc_handler.h" -#include "clnt_messenger.h" -DECLARE_uint32(shard_total_num); -DECLARE_uint32(clnt_thread_pool_size); DECLARE_string(cm_namespace); DECLARE_string(cm_svc_name); DECLARE_string(cm_port_name); @@ -38,13 +36,18 @@ DECLARE_string(cm_primary_node_ip); DECLARE_int32(cm_rpc_inter_port); DECLARE_int32(clnt_sync_req_timeout_ms); DECLARE_int32(clnt_async_req_timeout_ms); +DECLARE_uint32(shard_total_num); +DECLARE_uint32(clnt_thread_pool_size); DECLARE_uint32(clnt_cm_addr_check_interval_inSecs); +DECLARE_uint32(clnt_syncreq_retry_count); DECLARE_bool(simm_enable_trace); DECLARE_LOG_MODULE("simm_client"); template -concept HasPBKey = requires(RequestPB pb) { { pb.key() }; }; +concept HasPBKey = requires(RequestPB pb) { + {pb.key()}; +}; namespace simm { namespace clnt { @@ -56,7 +59,7 @@ ClientMessenger::ClientMessenger() shard_table_.reserve(shard_num_); hashkit_ = &simm::hashkit::HashkitBase::Instance(); // create rpc client and rdma mempool - sicl::rpc::SiRPC::newInstance(rpc_client_, false/*is_server*/); + sicl::rpc::SiRPC::newInstance(rpc_client_, false /*is_server*/); // set client request timeout sync_req_timeout_ms_ = convert_timeout_setting_to_timer_tick(FLAGS_clnt_sync_req_timeout_ms); async_req_timeout_ms_ = convert_timeout_setting_to_timer_tick(FLAGS_clnt_async_req_timeout_ms); @@ -94,19 +97,17 @@ error_code_t ClientMessenger::ReInit() { return CommonErr::InvalidState; } - auto [ret, servers] = update_all_route_table(cm_addr_); + auto [ret, routing] = update_all_route_table(cm_addr_); if (ret != CommonErr::OK) { MLOG_ERROR("Update routing table from cluster manager failed : {}", ret); cm_addr_ = ""; return CommonErr::InvalidState; } - - // build connections with all data servers - for (const auto &server_addr : servers) { - ret = build_connection(server_addr); - if (ret != CommonErr::OK) { - MLOG_ERROR("Build connection with {} failed, error: {}", server_addr, ret); - } + ret = ApplyRouteTableDiff(*routing); + if (ret != CommonErr::OK) { + MLOG_ERROR("Apply routing table diff failed : {}", ret); + cm_addr_ = ""; + return CommonErr::InvalidState; } #ifdef SIMM_ENABLE_TRACE @@ -131,7 +132,7 @@ error_code_t ClientMessenger::Init() { std::unique_lock lock(failover_mutex_); // NOTE: We might miss messages if the conditon variable is notified while we rebuild connections. // So check on regular intervals to workaround that. - failover_condv_.wait_for(lock, std::chrono::seconds(10)); + failover_condv_.wait_for(lock, std::chrono::seconds(FLAGS_clnt_cm_addr_check_interval_inSecs)); } if (!failover_flag_.load()) { // avoid one meaningless cm address query below, for curl_easy_perform() in get_cm_address() @@ -169,6 +170,11 @@ error_code_t ClientMessenger::Init() { } std::string ClientMessenger::get_cm_address() { +#if defined(SIMM_UNIT_TEST) + if (test_get_cm_address_hook_) { + return test_get_cm_address_hook_(); + } +#endif const std::string default_cm_addr = FLAGS_cm_primary_node_ip + ":" + std::to_string(FLAGS_cm_rpc_inter_port); // get cluster manager pod info from K8S api, get vector of [pod_name, pod_ip] // TODO: change to get cluster manager info from etcd @@ -188,6 +194,11 @@ std::string ClientMessenger::get_cm_address() { } error_code_t ClientMessenger::build_connection(const std::string &addr) { +#if defined(SIMM_UNIT_TEST) + if (test_build_connection_hook_) { + return test_build_connection_hook_(addr); + } +#endif auto node_addr = simm::common::NodeAddress::ParseFromString(addr); if (!node_addr) { MLOG_ERROR("Invalid server address format: {}", addr); @@ -198,14 +209,60 @@ error_code_t ClientMessenger::build_connection(const std::string &addr) { MLOG_ERROR("Build connection with {} failed", addr); return ClntErr::BuildConnectionFailed; } + auto ds_ctx = GetOrCreateConnectionContext(addr); + ds_ctx->StoreConnection(connection); + ds_ctx->gen_num.fetch_add(1); + ds_ctx->active.store(true); + return CommonErr::OK; +} + +void ClientMessenger::ReleaseConnectionContext(const std::shared_ptr &ds_ctx) { + if (ds_ctx == nullptr) { + return; + } + ds_ctx->active.store(false); + ds_ctx->gen_num.fetch_add(1); + ds_ctx->StoreConnection(nullptr); +} + +std::shared_ptr ClientMessenger::GetOrCreateConnectionContext( + const std::string &addr) { auto ds_ctx = std::make_shared(addr); if (auto [existing, inserted] = ds_conn_ctxs_.emplace(addr, ds_ctx); !inserted) { ds_ctx = existing->second; } - ds_ctx->connection = connection; - ds_ctx->gen_num.fetch_add(1); - ds_ctx->active.store(true); - return CommonErr::OK; + return ds_ctx; +} + +void ClientMessenger::PruneStaleConnectionContexts(const std::unordered_set &live_servers) { + std::vector stale_servers; + for (const auto &[addr, ds_ctx] : ds_conn_ctxs_) { + if (!live_servers.contains(addr)) { + stale_servers.push_back(addr); + } + } + + if (stale_servers.empty()) { + return; + } + + std::vector stale_shards; + for (const auto &[shard_id, ds_ctx] : shard_table_) { + if (ds_ctx != nullptr && !live_servers.contains(ds_ctx->ip_port)) { + stale_shards.push_back(shard_id); + } + } + for (auto shard_id : stale_shards) { + shard_table_.erase(shard_id); + } + + for (const auto &addr : stale_servers) { + auto it = ds_conn_ctxs_.find(addr); + if (it != ds_conn_ctxs_.end()) { + ReleaseConnectionContext(it->second); + ds_conn_ctxs_.erase(addr); + } + } } template @@ -221,42 +278,48 @@ error_code_t ClientMessenger::call_sync(uint16_t shard_id, auto rpc_ctx = ctx->get_rpc_ctx(); auto retry_delay = std::chrono::milliseconds(100); - auto retry_count = 3; - for (auto i = 0; i <= retry_count; ++i) { + for (auto i = 0; i <= FLAGS_clnt_syncreq_retry_count; ++i) { auto ds_ctx = shard_table_[shard_id]; if (ds_ctx->active.load()) { auto tag = ds_ctx->gen_num.load(); + auto connection = ds_ctx->LoadConnection(); + if (connection == nullptr) { + MLOG_ERROR("Transport connection is null for shard id {}, data server is {}", shard_id, ds_ctx->ip_port); + ReconnectByErrors(rpc_ctx, ds_ctx, shard_id, tag); + } else { #ifdef SIMM_APIPERF - auto t1 = std::chrono::steady_clock::now(); + auto t1 = std::chrono::steady_clock::now(); #endif - SIMM_TRACE_POINT(*ctx, simm::trace::TracePointType::CLIENT_CALLSYNC_BEFORE_RPC); + SIMM_TRACE_POINT(*ctx, simm::trace::TracePointType::CLIENT_CALLSYNC_BEFORE_RPC); - rpc_client_->SendRequest(ds_ctx->connection, req_type, req, resp.get(), rpc_ctx); + rpc_client_->SendRequest(connection, req_type, req, resp.get(), rpc_ctx); - SIMM_TRACE_POINT(*ctx, simm::trace::TracePointType::CLIENT_CALLSYNC_AFTER_RPC); + SIMM_TRACE_POINT(*ctx, simm::trace::TracePointType::CLIENT_CALLSYNC_AFTER_RPC); #ifdef SIMM_APIPERF - auto t2 = std::chrono::steady_clock::now(); - if constexpr (HasPBKey) { - MLOG_INFO("Perf-callsync-sendreq key:{} Lat:{} us", req.key(), - std::chrono::duration_cast(t2 - t1).count());; - } + auto t2 = std::chrono::steady_clock::now(); + if constexpr (HasPBKey) { + MLOG_INFO("Perf-callsync-sendreq key:{} Lat:{} us", + req.key(), + std::chrono::duration_cast(t2 - t1).count()); + } #endif - if (!rpc_ctx->Failed()) { - MLOG_DEBUG("call_sync rpc succeed"); - return CommonErr::OK; + if (!rpc_ctx->Failed()) { + MLOG_DEBUG("call_sync rpc succeed"); + return CommonErr::OK; + } + ReconnectByErrors(rpc_ctx, ds_ctx, shard_id, tag); } - ReconnectByErrors(rpc_ctx, ds_ctx, shard_id, tag); } else { MLOG_WARN("Transport connection is inactive for shard id {}, data server is {}", shard_id, ds_ctx->ip_port); } - if (i < retry_count) { + if (i < FLAGS_clnt_syncreq_retry_count) { std::this_thread::sleep_for(retry_delay); retry_delay *= 2; } } - MLOG_ERROR("Failed to send request after {} retries (sync call)", retry_count); + MLOG_ERROR("Failed to send request after {} retries (sync call)", FLAGS_clnt_syncreq_retry_count); return ClntErr::ClntSendRPCFailed; } @@ -283,12 +346,7 @@ error_code_t ClientMessenger::execute(const std::string &addr, std::shared_ptr resp, std::shared_ptr ctx, Callback callback) { - auto ds_ctx = std::make_shared(addr); - // If oe data server was not connected yet or failed to connect in init process, it will be added - // into client connection contexts. - if (auto [existing, inserted] = ds_conn_ctxs_.emplace(addr, ds_ctx); !inserted) { - ds_ctx = existing->second; - } + auto ds_ctx = GetOrCreateConnectionContext(addr); if (!ds_ctx->active.load()) { // If one data server is new added and can't be connected, background failover thread will try // reconnect action by reinit @@ -299,22 +357,28 @@ error_code_t ClientMessenger::execute(const std::string &addr, } MLOG_DEBUG("Connect with {} succeed", addr); } - auto connection = ds_ctx->connection; + auto connection = ds_ctx->LoadConnection(); + if (connection == nullptr) { + MLOG_ERROR("Connection with {} is null", addr); + return ClntErr::BuildConnectionFailed; + } auto rpc_ctx = ctx->get_rpc_ctx(); if (callback) { - #ifdef SIMM_APIPERF +#ifdef SIMM_APIPERF auto t1 = std::chrono::steady_clock::now(); - #endif +#endif rpc_client_->SendRequest(connection, req_type, req, resp.get(), rpc_ctx, std::move(callback)); - #ifdef SIMM_APIPERF +#ifdef SIMM_APIPERF auto t2 = std::chrono::steady_clock::now(); if constexpr (HasPBKey) { - MLOG_INFO("Perf-exec-sendreq key:{} Lat:{} us", req.key(), - std::chrono::duration_cast(t2 - t1).count());; + MLOG_INFO("Perf-exec-sendreq key:{} Lat:{} us", + req.key(), + std::chrono::duration_cast(t2 - t1).count()); + ; } - #endif +#endif MLOG_DEBUG("Call async rpc succeed"); } else { MLOG_DEBUG("Call sync rpc to {}", addr); @@ -328,75 +392,105 @@ error_code_t ClientMessenger::execute(const std::string &addr, return CommonErr::OK; } -std::pair> ClientMessenger::update_all_route_table(const std::string &ip_port) { +std::pair> ClientMessenger::update_all_route_table( + const std::string &ip_port) { auto addr = simm::common::NodeAddress::ParseFromString(ip_port); if (!addr) { MLOG_ERROR("Invalid ip_port string when update all route table: {}", ip_port); - return {ClntErr::GetRoutingTableFailed, {}}; + return {ClntErr::GetRoutingTableFailed, nullptr}; + } + +#if defined(SIMM_UNIT_TEST) + if (test_route_query_hook_) { + auto [ret, resp] = test_route_query_hook_(ip_port); + if (ret != CommonErr::OK || resp == nullptr) { + return {ret, nullptr}; + } + if (resp->ret_code() != 0) { + return {ClntErr::GetRoutingTableFailed, nullptr}; + } + return {CommonErr::OK, resp}; } +#endif // Query shard routing table from Cluster Manager QueryShardRoutingTableAllRequestPB req; - auto resp = std::make_shared(); // NOTE: The cluster maanager may return an empty list for unknown reasons. // Retry if that happens. auto query_retry_delay = std::chrono::milliseconds(100); auto query_retry_count = 3; for (auto i = 0; i <= query_retry_count; ++i) { + auto connection = rpc_client_->connect(addr->node_ip_, addr->node_port_); + if (connection == nullptr) { + MLOG_ERROR("Build connection with cluster manager {} failed", ip_port); + return {ClntErr::BuildConnectionFailed, nullptr}; + } + + auto resp = std::make_shared(); auto ctx = std::make_shared(); sicl::rpc::RpcContext *ctx_p; sicl::rpc::RpcContext::newInstance(ctx_p); auto rpc_ctx = std::shared_ptr(ctx_p); ctx->set_rpc_ctx(rpc_ctx); rpc_ctx->set_timeout(sync_req_timeout_ms_); - auto ret = execute( - ip_port, - static_cast(cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_ALL), - req, - resp, - ctx, - nullptr); - if (ret != CommonErr::OK) { - MLOG_ERROR("RPC to cluster manager({}) failed: {}", ip_port, ret); - return {ret, {}}; + rpc_client_->SendRequest(connection, + static_cast(cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_ALL), + req, + resp.get(), + rpc_ctx); + if (rpc_ctx->Failed()) { + MLOG_ERROR("RPC to cluster manager({}) failed: {}({})", ip_port, rpc_ctx->ErrorCode(), rpc_ctx->ErrorText()); + return {ClntErr::ClntSendRPCFailed, nullptr}; } else if (resp->ret_code() != 0) { MLOG_ERROR("Cluster manager({}) respond with error: {}", ip_port, resp->ret_code()); - return {ClntErr::GetRoutingTableFailed, {}}; + return {ClntErr::GetRoutingTableFailed, nullptr}; } if (resp->shard_info().size() != 0) { - break; + MLOG_INFO("QueryShardRoutingTableAll from Cluster manager({}) succeed, total shards num: {}", + ip_port, + resp->shard_info().size()); + return {CommonErr::OK, resp}; } if (i == query_retry_count) { MLOG_ERROR( "QueryShardRoutingTableAll from Cluster manager({}) failed after {} retries", ip_port, query_retry_count); - return {ClntErr::GetRoutingTableTimeout, {}}; + return {ClntErr::GetRoutingTableTimeout, nullptr}; } else { std::this_thread::sleep_for(query_retry_delay); query_retry_delay *= 2; } } - MLOG_INFO("QueryShardRoutingTableAll from Cluster manager({}) succeed, total shards num: {}", - ip_port, - resp->shard_info().size()); + return {ClntErr::GetRoutingTableTimeout, nullptr}; +} - std::vector servers; - auto route_entries = resp->shard_info(); - for (auto entry : route_entries) { +error_code_t ClientMessenger::ApplyRouteTableDiff(const QueryShardRoutingTableAllResponsePB &routing) { + std::unordered_set live_servers; + std::vector servers_to_connect; + + shard_table_.clear(); + for (const auto &entry : routing.shard_info()) { std::string ip = entry.data_server_address().ip(); uint16_t port = static_cast(entry.data_server_address().port()); const auto address = ip + ":" + std::to_string(port); - servers.push_back(address); - auto ds_ctx = std::make_shared(address); - if (auto [existing, inserted] = ds_conn_ctxs_.emplace(address, ds_ctx); !inserted) { - ds_ctx = existing->second; + auto ds_ctx = GetOrCreateConnectionContext(address); + if (live_servers.insert(address).second && (!ds_ctx->active.load() || ds_ctx->LoadConnection() == nullptr)) { + servers_to_connect.push_back(address); } for (auto shard_id : entry.shard_ids()) { shard_table_.insert_or_assign(shard_id, ds_ctx); } } + PruneStaleConnectionContexts(live_servers); - return {CommonErr::OK, servers}; + for (const auto &addr : servers_to_connect) { + auto ret = build_connection(addr); + if (ret != CommonErr::OK) { + MLOG_ERROR("Build connection with {} failed during route diff apply, error: {}", addr, ret); + } + } + + return CommonErr::OK; } void ClientMessenger::ReconnectByErrors(std::shared_ptr rpc_ctx, @@ -428,7 +522,16 @@ void ClientMessenger::ReconnectByErrors(std::shared_ptr r } } -error_code_t ClientMessenger::Put(const std::string &key, std::shared_ptr memp, std::shared_ptr ctx) { +void ClientMessenger::HandleAsyncRequestFailure(std::shared_ptr rpc_ctx, + std::shared_ptr request_ds_ctx, + uint16_t shard_id, + size_t old_conn_gen_num) { + ReconnectByErrors(std::move(rpc_ctx), std::move(request_ds_ctx), shard_id, old_conn_gen_num); +} + +error_code_t ClientMessenger::Put(const std::string &key, + std::shared_ptr memp, + std::shared_ptr ctx) { uint16_t shard_id = hashkit_->generate_16bit_hash_value(key.c_str(), key.length()) % shard_num_; // memblock must have descr ptr @@ -463,12 +566,12 @@ error_code_t ClientMessenger::Put(const std::string &key, std::shared_ptr( shard_id, static_cast(simm::ds::KVServerRpcType::RPC_CLIENT_KV_PUT), req, resp, ctx); if (ret != CommonErr::OK) { - MLOG_ERROR("Put object {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + MLOG_ERROR("Put key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); return ret; } else { MLOG_DEBUG("{}: key({}) ret_code({})", resp->GetTypeName(), key, resp->ret_code()); if (resp->ret_code() != 0) { - MLOG_ERROR("KVPut {} failed: {}", key, resp->ret_code()); + MLOG_ERROR("KVPut key {} failed: {}", key, resp->ret_code()); // FIXME(szzhao): remove workaround after it's implemented in data server if (resp->ret_code() == DsErr::DataAlreadyExists) { @@ -485,7 +588,7 @@ error_code_t ClientMessenger::Put(const std::string &key, std::shared_ptr memp, - std::function callback, + std::function callback, std::shared_ptr ctx) { uint16_t shard_id = hashkit_->generate_16bit_hash_value(key.c_str(), key.length()) % shard_num_; @@ -516,25 +619,25 @@ error_code_t ClientMessenger::AsyncPut(const std::string &key, auto resp = std::make_shared(); auto ds_ctx_before_req = shard_table_.at(shard_id); auto gen_num_before_req = ds_ctx_before_req->gen_num.load(); - auto done = [this, resp, key, gen_num_before_req, shard_id, cb = std::move(callback)]( + auto done = [this, resp, key, ds_ctx_before_req, gen_num_before_req, shard_id, cb = std::move(callback)]( const google::protobuf::Message *rsp, const std::shared_ptr rpc_ctx) { if (rpc_ctx->Failed()) { - MLOG_ERROR("Async put kv {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); - ReconnectByErrors(rpc_ctx, this->shard_table_.at(shard_id), shard_id, gen_num_before_req); + MLOG_ERROR("Async put key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + HandleAsyncRequestFailure(rpc_ctx, ds_ctx_before_req, shard_id, gen_num_before_req); cb(rpc_ctx->ErrorCode()); } else { auto new_resp = static_cast(rsp); MLOG_DEBUG("{}: key({}) ret_code({})", new_resp->GetTypeName(), key, new_resp->ret_code()); if (new_resp->ret_code() != 0) { - MLOG_ERROR("AsyncKVPut {} failed: {}", key, new_resp->ret_code()); + MLOG_ERROR("AsyncKVPut key {} failed: {}", key, new_resp->ret_code()); } cb(new_resp->ret_code()); } }; - #ifdef SIMM_APIPERF - auto t1 = std::chrono::steady_clock::now(); - #endif +#ifdef SIMM_APIPERF + auto t1 = std::chrono::steady_clock::now(); +#endif call_async( shard_id, static_cast(simm::ds::KVServerRpcType::RPC_CLIENT_KV_PUT), @@ -542,16 +645,19 @@ error_code_t ClientMessenger::AsyncPut(const std::string &key, resp, ctx, std::move(done)); - #ifdef SIMM_APIPERF - auto t2 = std::chrono::steady_clock::now(); - MLOG_INFO("Perf-aput-callasync key:{} Lat:{} us", key, - std::chrono::duration_cast(t2 - t1).count()); - #endif +#ifdef SIMM_APIPERF + auto t2 = std::chrono::steady_clock::now(); + MLOG_INFO("Perf-aput-callasync key:{} Lat:{} us", + key, + std::chrono::duration_cast(t2 - t1).count()); +#endif return CommonErr::OK; } -int32_t ClientMessenger::Get(const std::string &key, std::shared_ptr memp, std::shared_ptr ctx) { +int32_t ClientMessenger::Get(const std::string &key, + std::shared_ptr memp, + std::shared_ptr ctx) { uint16_t shard_id = hashkit_->generate_16bit_hash_value(key.c_str(), key.length()) % shard_num_; sicl::transport::MemDesc *mem_desc; @@ -585,16 +691,17 @@ int32_t ClientMessenger::Get(const std::string &key, std::shared_ptr(simm::ds::KVServerRpcType::RPC_CLIENT_KV_GET), req, resp, ctx); #ifdef SIMM_APIPERF auto t2 = std::chrono::steady_clock::now(); - MLOG_INFO("Perf-get-callsync key:{} Lat:{} us", key, - std::chrono::duration_cast(t2 - t1).count()); + MLOG_INFO("Perf-get-callsync key:{} Lat:{} us", + key, + std::chrono::duration_cast(t2 - t1).count()); #endif if (ret != CommonErr::OK) { - MLOG_ERROR("Get kv {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + MLOG_ERROR("Get key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); return ret; } else { MLOG_DEBUG("{}: key({}) ret_code({}) val_len({})", resp->GetTypeName(), key, resp->ret_code(), resp->val_len()); if (resp->ret_code() != 0) { - MLOG_ERROR("KVGet {} failed: {}", key, resp->ret_code()); + MLOG_ERROR("KVGet key {} failed: {}", key, resp->ret_code()); return ClntErr::ClntGetObjectFailed; } } @@ -604,7 +711,7 @@ int32_t ClientMessenger::Get(const std::string &key, std::shared_ptr memp, - std::function callback, + std::function callback, std::shared_ptr ctx) { uint16_t shard_id = hashkit_->generate_16bit_hash_value(key.c_str(), key.length()) % shard_num_; @@ -635,11 +742,11 @@ error_code_t ClientMessenger::AsyncGet(const std::string &key, auto resp = std::make_shared(); auto ds_ctx_before_req = shard_table_.at(shard_id); auto gen_num_before_req = ds_ctx_before_req->gen_num.load(); - auto done = [this, resp, key, gen_num_before_req, shard_id, cb = std::move(callback)]( + auto done = [this, resp, key, ds_ctx_before_req, gen_num_before_req, shard_id, cb = std::move(callback)]( const google::protobuf::Message *rsp, const std::shared_ptr rpc_ctx) { if (rpc_ctx->Failed()) { - MLOG_ERROR("Async get object {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); - ReconnectByErrors(rpc_ctx, this->shard_table_.at(shard_id), shard_id, gen_num_before_req); + MLOG_ERROR("Async get key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + HandleAsyncRequestFailure(rpc_ctx, ds_ctx_before_req, shard_id, gen_num_before_req); cb(rpc_ctx->ErrorCode()); } else { auto new_resp = static_cast(rsp); @@ -649,7 +756,7 @@ error_code_t ClientMessenger::AsyncGet(const std::string &key, new_resp->ret_code(), new_resp->val_len()); if (new_resp->ret_code() != 0) { - MLOG_ERROR("AsyncKVGet {} failed: {}", key, new_resp->ret_code()); + MLOG_ERROR("AsyncKVGet key {} failed: {}", key, new_resp->ret_code()); cb(new_resp->ret_code()); } else { cb(new_resp->val_len()); @@ -657,9 +764,9 @@ error_code_t ClientMessenger::AsyncGet(const std::string &key, } }; - #ifdef SIMM_APIPERF - auto t1 = std::chrono::steady_clock::now(); - #endif +#ifdef SIMM_APIPERF + auto t1 = std::chrono::steady_clock::now(); +#endif call_async( shard_id, static_cast(simm::ds::KVServerRpcType::RPC_CLIENT_KV_GET), @@ -667,11 +774,12 @@ error_code_t ClientMessenger::AsyncGet(const std::string &key, resp, ctx, std::move(done)); - #ifdef SIMM_APIPERF - auto t2 = std::chrono::steady_clock::now(); - MLOG_INFO("Perf-aget-callasync key:{} Lat:{} us", key, - std::chrono::duration_cast(t2 - t1).count()); - #endif +#ifdef SIMM_APIPERF + auto t2 = std::chrono::steady_clock::now(); + MLOG_INFO("Perf-aget-callasync key:{} Lat:{} us", + key, + std::chrono::duration_cast(t2 - t1).count()); +#endif return CommonErr::OK; } @@ -692,18 +800,22 @@ error_code_t ClientMessenger::Delete(const std::string &key, std::shared_ptr( shard_id, static_cast(simm::ds::KVServerRpcType::RPC_CLIENT_KV_DEL), req, resp, ctx); if (ret != CommonErr::OK) { - MLOG_ERROR("Delete object {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + MLOG_ERROR("Delete key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + return ret; } else { MLOG_DEBUG("{}: key({}) ret_code({})", resp->GetTypeName(), key, resp->ret_code()); if (resp->ret_code() != 0) { - MLOG_ERROR("KVDelete {} failed: {}", key, resp->ret_code()); + MLOG_ERROR("KVDelete key {} failed: {}", key, resp->ret_code()); + return static_cast(resp->ret_code()); } } return CommonErr::OK; } -error_code_t ClientMessenger::AsyncDelete(const std::string &key, std::function callback, std::shared_ptr ctx) { +error_code_t ClientMessenger::AsyncDelete(const std::string &key, + std::function callback, + std::shared_ptr ctx) { uint16_t shard_id = hashkit_->generate_16bit_hash_value(key.c_str(), key.length()) % shard_num_; // build KVDel rpc @@ -718,17 +830,17 @@ error_code_t ClientMessenger::AsyncDelete(const std::string &key, std::function< auto resp = std::make_shared(); auto ds_ctx_before_req = shard_table_.at(shard_id); auto gen_num_before_req = ds_ctx_before_req->gen_num.load(); - auto done = [this, resp, key, gen_num_before_req, shard_id, cb = std::move(callback)]( + auto done = [this, resp, key, ds_ctx_before_req, gen_num_before_req, shard_id, cb = std::move(callback)]( const google::protobuf::Message *rsp, const std::shared_ptr rpc_ctx) { if (rpc_ctx->Failed()) { - MLOG_ERROR("Async delete kv {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); - ReconnectByErrors(rpc_ctx, this->shard_table_.at(shard_id), shard_id, gen_num_before_req); + MLOG_ERROR("AsyncDelete key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + HandleAsyncRequestFailure(rpc_ctx, ds_ctx_before_req, shard_id, gen_num_before_req); cb(rpc_ctx->ErrorCode()); } else { auto new_resp = static_cast(rsp); MLOG_DEBUG("{}: key({}) ret_code({})", new_resp->GetTypeName(), key, new_resp->ret_code()); if (new_resp->ret_code() != 0) { - MLOG_ERROR("AsyncKVDelete {} failed: {}", key, new_resp->ret_code()); + MLOG_ERROR("AsyncKVDelete key {} failed: {}", key, new_resp->ret_code()); } cb(new_resp->ret_code()); } @@ -760,12 +872,12 @@ error_code_t ClientMessenger::Exists(const std::string &key, std::shared_ptr( shard_id, static_cast(simm::ds::KVServerRpcType::RPC_CLIENT_KV_LOOKUP), req, resp, ctx); if (ret != CommonErr::OK) { - MLOG_ERROR("Lookup kv {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + MLOG_ERROR("Lookup key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); return ret; } else { MLOG_DEBUG("{}: key({}) ret_code({})", resp->GetTypeName(), key, resp->ret_code()); if (resp->ret_code() != 0) { - MLOG_ERROR("KVLookup {} failed: {}", key, resp->ret_code()); + MLOG_ERROR("KVLookup key {} failed: {}", key, resp->ret_code()); return resp->ret_code(); } } @@ -773,7 +885,9 @@ error_code_t ClientMessenger::Exists(const std::string &key, std::shared_ptr callback, std::shared_ptr ctx) { +error_code_t ClientMessenger::AsyncExists(const std::string &key, + std::function callback, + std::shared_ptr ctx) { uint16_t shard_id = hashkit_->generate_16bit_hash_value(key.c_str(), key.length()) % shard_num_; // build KVLookup rpc @@ -788,17 +902,17 @@ error_code_t ClientMessenger::AsyncExists(const std::string &key, std::function< auto resp = std::make_shared(); auto ds_ctx_before_req = shard_table_.at(shard_id); auto gen_num_before_req = ds_ctx_before_req->gen_num.load(); - auto done = [this, resp, key, gen_num_before_req, shard_id, cb = std::move(callback)]( + auto done = [this, resp, key, ds_ctx_before_req, gen_num_before_req, shard_id, cb = std::move(callback)]( const google::protobuf::Message *rsp, const std::shared_ptr rpc_ctx) { if (rpc_ctx->Failed()) { - MLOG_ERROR("Async lookup kv {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); - ReconnectByErrors(rpc_ctx, this->shard_table_.at(shard_id), shard_id, gen_num_before_req); + MLOG_ERROR("Async lookup key {} rpc failed: {}({})", key, rpc_ctx->ErrorText(), rpc_ctx->ErrorCode()); + HandleAsyncRequestFailure(rpc_ctx, ds_ctx_before_req, shard_id, gen_num_before_req); cb(rpc_ctx->ErrorCode()); } else { auto new_resp = static_cast(rsp); MLOG_DEBUG("{}: key({}) ret_code({})", new_resp->GetTypeName(), key, new_resp->ret_code()); if (new_resp->ret_code() != 0) { - MLOG_ERROR("AsyncKVLookup {} failed: {}", key, new_resp->ret_code()); + MLOG_ERROR("AsyncKVLookup key {} failed: {}", key, new_resp->ret_code()); } cb(new_resp->ret_code()); } @@ -820,17 +934,19 @@ std::vector ClientMessenger::MultiPut(const std::vector(keys.size(), CommonErr::InvalidArgument); } - auto multiGetResults = folly::coro::blockingWait([this, datas, keys]() -> folly::coro::Task> { - std::vector> tasks; - for (size_t i = 0; i < keys.size(); ++i) { - auto key = keys[i]; - auto memp = datas[i]; - auto ctx = std::make_shared(); - tasks.push_back(folly::via(&executor_, [this, key, memp, ctx]() -> error_code_t { return this->Put(key, memp, ctx); })); - } - auto results = co_await folly::coro::collectAllRange(std::move(tasks)); - co_return results; - }()); + auto multiGetResults = + folly::coro::blockingWait([this, datas, keys]() -> folly::coro::Task> { + std::vector> tasks; + for (size_t i = 0; i < keys.size(); ++i) { + auto key = keys[i]; + auto memp = datas[i]; + auto ctx = std::make_shared(); + tasks.push_back( + folly::via(&executor_, [this, key, memp, ctx]() -> error_code_t { return this->Put(key, memp, ctx); })); + } + auto results = co_await folly::coro::collectAllRange(std::move(tasks)); + co_return results; + }()); // Error handling and logging will happen in this->Put() and upper calls(like KVStore::MGet). // Now just return error code. @@ -849,7 +965,8 @@ std::vector ClientMessenger::MultiGet(const std::vector &k auto key = keys[i]; auto memp = datas[i]; auto ctx = std::make_shared(); - tasks.push_back(folly::via(&executor_, [this, key, memp, ctx]() -> int32_t { return this->Get(key, memp, ctx); })); + tasks.push_back( + folly::via(&executor_, [this, key, memp, ctx]() -> int32_t { return this->Get(key, memp, ctx); })); } auto results = co_await folly::coro::collectAllRange(std::move(tasks)); co_return results; @@ -893,8 +1010,6 @@ inline sicl::transport::TimerTick ClientMessenger::convert_timeout_setting_to_ti return sicl::transport::TimerTick::TIMER_1S; } else if (timeout_ms <= 3000) { return sicl::transport::TimerTick::TIMER_3S; - } else if (timeout_ms <= 10000) { - return sicl::transport::TimerTick::TIMER_10S; } else if (timeout_ms <= 5000) { return sicl::transport::TimerTick::TIMER_5S; } else if (timeout_ms <= 10000) { diff --git a/src/client/clnt_messenger.h b/src/client/clnt_messenger.h index 3fc8c2d..a560858 100644 --- a/src/client/clnt_messenger.h +++ b/src/client/clnt_messenger.h @@ -5,7 +5,9 @@ #include #include #include +#include #include +#include #include #include @@ -25,6 +27,7 @@ #include "common/base/memory.h" #include "common/errcode/errcode_def.h" #include "common/hashkit/hashkit.h" +#include "proto/cm_clnt_rpcs.pb.h" #include "common/trace/trace_server.h" @@ -34,6 +37,10 @@ namespace clnt { using Callback = std::function ctx)>; +#if defined(SIMM_UNIT_TEST) +class ClientMessengerTestPeer; +#endif + class ClientMessenger { public: explicit ClientMessenger(); @@ -64,13 +71,11 @@ class ClientMessenger { std::shared_ptr memp, std::function callback, std::shared_ptr ctx); - error_code_t Exists(const std::string &key, - std::shared_ptr ctx); + error_code_t Exists(const std::string &key, std::shared_ptr ctx); error_code_t AsyncExists(const std::string &key, std::function callback, std::shared_ptr ctx); - error_code_t Delete(const std::string &key, - std::shared_ptr ctx); + error_code_t Delete(const std::string &key, std::shared_ptr ctx); error_code_t AsyncDelete(const std::string &key, std::function callback, std::shared_ptr ctx); @@ -117,8 +122,9 @@ class ClientMessenger { std::shared_ptr ctx, Callback callback); - // Send rpc request to cluster_manager ip:port, get newest routing talbe and update entire local routing table - std::pair> update_all_route_table(const std::string &ip_port); + // Send rpc request to cluster_manager ip:port, get newest routing table. + std::pair> update_all_route_table( + const std::string &ip_port); struct ConnectionContext; // forward declaration // Trigger reconnect to data server upon specified rpc errors @@ -126,20 +132,45 @@ class ClientMessenger { std::shared_ptr ds_ctx, uint16_t shard_id, size_t old_conn_gen_num); + std::shared_ptr GetOrCreateConnectionContext(const std::string &addr); + void PruneStaleConnectionContexts(const std::unordered_set &live_servers); + void HandleAsyncRequestFailure(std::shared_ptr rpc_ctx, + std::shared_ptr request_ds_ctx, + uint16_t shard_id, + size_t old_conn_gen_num); // Convert timeout setting from user in milliseconds to sicl TimerTick type sicl::transport::TimerTick convert_timeout_setting_to_timer_tick(int32_t timeout_ms); + void ReleaseConnectionContext(const std::shared_ptr &ds_ctx); + error_code_t ApplyRouteTableDiff(const QueryShardRoutingTableAllResponsePB &routing); + +#if defined(SIMM_UNIT_TEST) + friend class ClientMessengerTestPeer; +#endif + private: struct ConnectionContext { std::string ip_port; - std::shared_ptr connection; std::atomic active{false}; /// monotonic to prevent ABA (race condition) std::atomic gen_num{0}; - ConnectionContext(const std::string &_ip_port) - : ip_port(_ip_port), connection{nullptr}, active(false), gen_num(0) {} + ConnectionContext(const std::string &_ip_port) : ip_port(_ip_port), active(false), gen_num(0) {} + + void StoreConnection(std::shared_ptr new_connection) { + std::unique_lock lock(connection_mutex_); + connection_ = std::move(new_connection); + } + + std::shared_ptr LoadConnection() const { + std::shared_lock lock(connection_mutex_); + return connection_; + } + + private: + mutable std::shared_mutex connection_mutex_; + std::shared_ptr connection_{nullptr}; }; // cluster manager address @@ -159,7 +190,7 @@ class ClientMessenger { // sicl objects sicl::rpc::SiRPC *rpc_client_{nullptr}; - sicl::rpc::SiRPC *admin_rpc_service_{nullptr}; + sicl::rpc::SiRPC *admin_rpc_service_{nullptr}; sicl::transport::IbvDeviceManager *ibv_mgr_{nullptr}; sicl::transport::Mempool *mempool_{nullptr}; @@ -177,6 +208,15 @@ class ClientMessenger { std::mutex failover_mutex_; std::condition_variable failover_condv_; std::unique_ptr trace_server_{nullptr}; + +#if defined(SIMM_UNIT_TEST) + private: + // ONLY for mock UT cases + std::function test_get_cm_address_hook_{nullptr}; + std::function>(const std::string &)> + test_route_query_hook_{nullptr}; + std::function test_build_connection_hook_{nullptr}; +#endif }; } // namespace clnt diff --git a/src/cluster_manager/cm_hb_monitor.cc b/src/cluster_manager/cm_hb_monitor.cc index 2178c74..3b6c954 100644 --- a/src/cluster_manager/cm_hb_monitor.cc +++ b/src/cluster_manager/cm_hb_monitor.cc @@ -3,10 +3,10 @@ #include +#include "cm_hb_monitor.h" #include "cm_node_manager.h" #include "common/base/assert.h" #include "common/logging/logging.h" -#include "cm_hb_monitor.h" DECLARE_LOG_MODULE("cluster_manager"); @@ -26,10 +26,14 @@ error_code_t ClusterManagerHBMonitor::Init() { } error_code_t ClusterManagerHBMonitor::Start() { + stop_flag_.store(false); + bg_scan_thread_baton_.reset(); + if (bg_scan_thread_.joinable()) { + MLOG_WARN("Background HB scan thread is still joinable, not start new thread"); + return CommonErr::OK; + } auto self = shared_from_this(); - bg_scan_thread_ = std::jthread([self]() { - self->BgHBScanLoop(); - }); + bg_scan_thread_ = std::jthread([self]() { self->BgHBScanLoop(); }); return CommonErr::OK; } @@ -50,7 +54,9 @@ error_code_t ClusterManagerHBMonitor::OnRecvNodeHeartbeat(const std::string &nod auto &entry = (*uomap_locked)[node_addr_str]; if (entry.size() >= FLAGS_cm_heartbeat_records_perserver) { MLOG_DEBUG("Heartbeat records(current:{}) for node({}) is full(limit:{}), removing oldest records", - entry.size(), node_addr_str, FLAGS_cm_heartbeat_records_perserver); + entry.size(), + node_addr_str, + FLAGS_cm_heartbeat_records_perserver); entry.pop_front(); } entry.push_back(hb_ts); @@ -61,10 +67,9 @@ error_code_t ClusterManagerHBMonitor::OnRecvNodeHeartbeat(const std::string &nod void ClusterManagerHBMonitor::BgHBScanLoop() { while (!stop_flag_.load()) { - { // code block for rlock - std::vector dead_dataservers{}; + std::vector dead_dataservers{}; + { // code block for rlock auto now = std::chrono::steady_clock::now(); - bool trigger_rebalance = false; auto uomap_locked = ds_hb_records_.rlock(); for (auto it = uomap_locked->begin(); it != uomap_locked->end();) { auto &hb_records = it->second; @@ -79,7 +84,8 @@ void ClusterManagerHBMonitor::BgHBScanLoop() { auto previous_hb_ts = hb_records.back(); if (now - previous_hb_ts.monotonic_tp_ > std::chrono::seconds(FLAGS_cm_heartbeat_timeout_inSecs)) { MLOG_ERROR("Dataserver node({}) heartbeat timeout({} secs), mark it as DEAD", - it->first, FLAGS_cm_heartbeat_timeout_inSecs); + it->first, + FLAGS_cm_heartbeat_timeout_inSecs); if (cm_node_manager_ptr_->QueryNodeStatus(it->first) == NodeStatus::DEAD) { ++it; continue; @@ -88,7 +94,6 @@ void ClusterManagerHBMonitor::BgHBScanLoop() { cm_node_manager_ptr_->UpdateNodeStatus(it->first, NodeStatus::DEAD); // record all dead servers in current scan round to trigger shard manager update in batch dead_dataservers.emplace_back(it->first); - trigger_rebalance = true; // FIXME(ytji): we still keep the heartbeat records for the node, // it = uomap_locked->erase(it); } @@ -96,12 +101,13 @@ void ClusterManagerHBMonitor::BgHBScanLoop() { ++it; } - // dead dataservers will trigger shard table refresh - if (!dead_dataservers.empty() && trigger_rebalance) { - HandleNodeFailure(dead_dataservers); - } } // release rlock + // dead dataservers will trigger shard table refresh + if (!dead_dataservers.empty()) { + HandleNodeFailure(dead_dataservers); + } + MLOG_DEBUG("ClusterManagerHBMonitor::BgHBScanLoop wait some seconds before next round..."); // sleep some seconds before next scan bg_scan_thread_baton_.timed_wait(std::chrono::milliseconds(FLAGS_cm_heartbeat_bg_scan_interval_inSecs * 1000)); @@ -109,14 +115,11 @@ void ClusterManagerHBMonitor::BgHBScanLoop() { } } -void ClusterManagerHBMonitor::HandleNodeFailure(const std::vector& dead_node_addresses) { +void ClusterManagerHBMonitor::HandleNodeFailure(const std::vector &dead_node_addresses) { MLOG_WARN("Handling failure of {} nodes", dead_node_addresses.size()); - + auto alive_servers = cm_node_manager_ptr_->GetAllNodeAddress(true /* alive only */); - if (alive_servers.empty()) { - return; - } - + // rebalance shards after node failure error_code_t ret = cm_shard_manager_ptr_->RebalanceShardsAfterNodeFailure(dead_node_addresses, alive_servers); if (ret != CommonErr::OK) { diff --git a/src/cluster_manager/cm_hb_monitor.h b/src/cluster_manager/cm_hb_monitor.h index 399789a..18c37ba 100644 --- a/src/cluster_manager/cm_hb_monitor.h +++ b/src/cluster_manager/cm_hb_monitor.h @@ -8,8 +8,12 @@ #include #include -#include "cm_shard_manager.h" +#if defined(SIMM_UNIT_TEST) +#include +#endif + #include "cm_node_manager.h" +#include "cm_shard_manager.h" #include "common/base/common_types.h" #include "common/errcode/errcode_def.h" @@ -21,13 +25,13 @@ class ClusterManagerHBMonitor : public std::enable_shared_from_this nm_ptr, std::shared_ptr sm_ptr) - : cm_node_manager_ptr_(nm_ptr), cm_shard_manager_ptr_(sm_ptr) {} + : cm_node_manager_ptr_(nm_ptr), cm_shard_manager_ptr_(sm_ptr) {} virtual ~ClusterManagerHBMonitor(); - ClusterManagerHBMonitor(const ClusterManagerHBMonitor&) = delete; - ClusterManagerHBMonitor& operator=(const ClusterManagerHBMonitor&) = delete; - ClusterManagerHBMonitor(ClusterManagerHBMonitor&&) = delete; - ClusterManagerHBMonitor& operator=(ClusterManagerHBMonitor&&) = delete; + ClusterManagerHBMonitor(const ClusterManagerHBMonitor &) = delete; + ClusterManagerHBMonitor &operator=(const ClusterManagerHBMonitor &) = delete; + ClusterManagerHBMonitor(ClusterManagerHBMonitor &&) = delete; + ClusterManagerHBMonitor &operator=(ClusterManagerHBMonitor &&) = delete; public: error_code_t Init(); @@ -57,6 +61,7 @@ class ClusterManagerHBMonitor : public std::enable_shared_from_this #include +#if defined(SIMM_UNIT_TEST) +#include +#endif + #include "rpc/connection.h" #include "rpc/rpc.h" #include "rpc/rpc_context.h" @@ -23,7 +27,7 @@ namespace cm { using NodeStatus = common::NodeStatus; -class ClusterManagerNodeManager : public std::enable_shared_from_this { +class ClusterManagerNodeManager : public std::enable_shared_from_this { public: /** * @brief ClusterManagerNodeManager controls all data server nodes. Get all @@ -34,9 +38,9 @@ class ClusterManagerNodeManager : public std::enable_shared_from_this> GetAllNodeAddress(bool alive = true); // add a new node in node manager, will send a rpc to get node resource info - error_code_t AddNode(const std::string & addr_str); + error_code_t AddNode(const std::string &addr_str); // delete node from node manager, remove all node info from map - error_code_t DelNode(const std::string & addr_str); + error_code_t DelNode(const std::string &addr_str); // update node status (RUNNING/DEAD) - error_code_t UpdateNodeStatus(const std::string & addr_str, NodeStatus status); + error_code_t UpdateNodeStatus(const std::string &addr_str, NodeStatus status); // query one data node status - NodeStatus QueryNodeStatus(const std::string & addr_str); + NodeStatus QueryNodeStatus(const std::string &addr_str); // query whether data node exists - bool QueryNodeExists(const std::string & addr_str); + bool QueryNodeExists(const std::string &addr_str); // get resource info of a single data node - std::shared_ptr GetNodeResource(const std::string & addr_str); + std::shared_ptr GetNodeResource(const std::string &addr_str); // get status of all data nodes // Returns an unordered_map of address string to node status @@ -74,7 +78,7 @@ class ClusterManagerNodeManager : public std::enable_shared_from_this> GetAllNodeResource(); private: - std::shared_ptr getNodeResource(const std::string & addr_str); + std::shared_ptr getNodeResource(const std::string &addr_str); void updateAllNodeResource(); private: diff --git a/src/cluster_manager/cm_rpc_handler.cc b/src/cluster_manager/cm_rpc_handler.cc index 74291d4..00a7c82 100644 --- a/src/cluster_manager/cm_rpc_handler.cc +++ b/src/cluster_manager/cm_rpc_handler.cc @@ -3,8 +3,8 @@ #include #include -#include #include +#include #include "cm_rpc_handler.h" #include "common/base/common_types.h" @@ -44,6 +44,9 @@ static inline void FillRespWithRoutingTableEntryHelper(const simm::cm::QueryResu using pbMsgMapType = std::unordered_map, std::vector>; pbMsgMapType pb_msg_map; for (const auto &entry : resmap) { + if (!entry.second) { + continue; + } pb_msg_map[entry.second].push_back(entry.first); } for (const auto &map_entry : pb_msg_map) { @@ -57,9 +60,19 @@ static inline void FillRespWithRoutingTableEntryHelper(const simm::cm::QueryResu } } +static inline bool HasIncompleteRoutingEntries(const simm::cm::QueryResultMap &resmap) { + for (const auto &[shard_id, node_addr] : resmap) { + if (!node_addr) { + MLOG_WARN("Shard {} has no assigned dataserver in routing table", shard_id); + return true; + } + } + return false; +} + void NewNodeHandshakeHandler::Work(const std::shared_ptr ctx, const std::shared_ptr conn, - [[maybe_unused]] const google::protobuf::Message *request) const { + [[maybe_unused]] const google::protobuf::Message *request) const { auto req_begin_ts = std::chrono::steady_clock::now(); auto req = dynamic_cast(request); auto resp = std::make_shared(); @@ -67,7 +80,7 @@ void NewNodeHandshakeHandler::Work(const std::shared_ptr error_code_t ret = CommonErr::OK; if (simm::common::ModuleServiceState::GetInstance().GracePeriodFinished()) { MLOG_INFO("Grace period is already finished, new dataserver({}) will be waited for joining the cluster", - node_addr.toString()); + node_addr.toString()); // already out of grace period, new dataserver nodes will be hold for // one timewindow, and be added in batch after current timewindow finishes } else { @@ -75,16 +88,19 @@ void NewNodeHandshakeHandler::Work(const std::shared_ptr // FIXME(ytji): needn't to use NodeAddress as intermediary struct MLOG_INFO("Still in Grace period new dataserver({}) will be added into cluster", node_addr.toString()); ret = node_manager_->AddNode(node_addr.toString()); - shard_manager_->BatchAssignRoutingTable(std::vector(req->shard_ids().begin(), req->shard_ids().end()), - std::make_shared(node_addr)); + shard_manager_->BatchAssignRoutingTable(std::vector(req->shard_ids().begin(), req->shard_ids().end()), + std::make_shared(node_addr)); } - + if (ret != CommonErr::OK) { MLOG_ERROR("Failed to register new node({}) into cluster, ret:{}", node_addr.toString(), ret); } resp->set_ret_code(ret); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("hand_shake", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("hand_shake", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("hand_shake"); if (ret != CommonErr::OK) { simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("hand_shake"); @@ -115,8 +131,11 @@ void NodeHeartBeatHandler::Work(const std::shared_ptr ctx } resp->set_ret_code(ret); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("heart_beat", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("heart_beat", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("heart_beat"); if (ret != CommonErr::OK) { simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("heart_beat"); @@ -156,8 +175,11 @@ void RoutingTableQuerySingleHandler::Work(const std::shared_ptrset_ret_code(ret); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("query_routing_table_single", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("query_routing_table_single", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("query_routing_table_single"); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_single"); SEND_RESPONSE(ctx, resp); @@ -168,13 +190,18 @@ void RoutingTableQuerySingleHandler::Work(const std::shared_ptrshard_id()); ret = CommonErr::CmTargetShardIdNotFound; + } else if (HasIncompleteRoutingEntries(query_res)) { + ret = CommonErr::CmRoutingInfoNotComplete; } else { FillRespWithRoutingTableEntryHelper(query_res, resp.get()); } resp->set_ret_code(ret); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("query_routing_table_single", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("query_routing_table_single", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("query_routing_table_single"); if (ret != CommonErr::OK) { simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_single"); @@ -192,8 +219,11 @@ void RoutingTableQueryBatchHandler::Work(const std::shared_ptrset_ret_code(CmErr::InitInGracePeriod); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("query_routing_table_batch", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("query_routing_table_batch", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("query_routing_table_batch"); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_batch"); SEND_RESPONSE(ctx, resp); @@ -206,9 +236,13 @@ void RoutingTableQueryBatchHandler::Work(const std::shared_ptrset_ret_code(CommonErr::CmTargetShardIdNotFound); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_batch"); + } else if (HasIncompleteRoutingEntries(query_res)) { + resp->set_ret_code(CommonErr::CmRoutingInfoNotComplete); + simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_batch"); } else if (query_res.size() != target_shards.size()) { MLOG_WARN("Some target shard ids not found in routing table, target_shards: {}, found_shards: {}", - target_shards.size(), query_res.size()); + target_shards.size(), + query_res.size()); // TODO(ytji): what error code should we return here? resp->set_ret_code(CommonErr::OK); } else { @@ -217,8 +251,11 @@ void RoutingTableQueryBatchHandler::Work(const std::shared_ptr( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("query_routing_table_batch", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("query_routing_table_batch"); SEND_RESPONSE(ctx, resp); } @@ -233,18 +270,26 @@ void RoutingTableQueryAllHandler::Work(const std::shared_ptrset_ret_code(CmErr::InitInGracePeriod); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("query_routing_table_all", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("query_routing_table_all", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("query_routing_table_all"); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_all"); SEND_RESPONSE(ctx, resp); - return;; + return; + ; } simm::cm::QueryResultMap query_res = shard_manager_->QueryAllShardRoutingInfos(); if (query_res.size() != FLAGS_shard_total_num) { MLOG_WARN("All routing table info query result is not complete, target_shards: {}, found_shards: {}", - FLAGS_shard_total_num, query_res.size()); + FLAGS_shard_total_num, + query_res.size()); + resp->set_ret_code(CommonErr::CmRoutingInfoNotComplete); + simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_all"); + } else if (HasIncompleteRoutingEntries(query_res)) { resp->set_ret_code(CommonErr::CmRoutingInfoNotComplete); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("query_routing_table_all"); } else { @@ -253,8 +298,11 @@ void RoutingTableQueryAllHandler::Work(const std::shared_ptr( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("query_routing_table_all", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("query_routing_table_all"); SEND_RESPONSE(ctx, resp); } @@ -268,8 +316,11 @@ void ListNodesHandler::Work(const std::shared_ptr ctx, if (FOLLY_UNLIKELY(!simm::common::ModuleServiceState::GetInstance().IsServiceReady())) { MLOG_WARN("Cluster Manager is still in grace period, service is not ready yet"); resp->set_ret_code(CmErr::InitInGracePeriod); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("list_nodes", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("list_nodes", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("list_nodes"); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("list_nodes"); SEND_RESPONSE(ctx, resp); @@ -315,15 +366,18 @@ void ListNodesHandler::Work(const std::shared_ptr ctx, } resp->set_ret_code(CommonErr::OK); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("list_nodes", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("list_nodes", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("list_nodes"); SEND_RESPONSE(ctx, resp); } void SetNodeStatusHandler::Work(const std::shared_ptr ctx, - const std::shared_ptr conn, - const google::protobuf::Message *request) const { + const std::shared_ptr conn, + const google::protobuf::Message *request) const { auto req_begin_ts = std::chrono::steady_clock::now(); auto req = dynamic_cast(request); auto resp = std::make_shared(); @@ -332,8 +386,11 @@ void SetNodeStatusHandler::Work(const std::shared_ptr ctx if (FOLLY_UNLIKELY(!simm::common::ModuleServiceState::GetInstance().IsServiceReady())) { MLOG_WARN("Cluster Manager is still in grace period, service is not ready yet"); resp->set_ret_code(CommonErr::TargetUnavailable); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("set_node_status", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("set_node_status", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("set_node_status"); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("set_node_status"); SEND_RESPONSE(ctx, resp); @@ -342,12 +399,15 @@ void SetNodeStatusHandler::Work(const std::shared_ptr ctx // Not exist should return error directly, use another method std::string addr_str = req->node().ip() + ":" + std::to_string(req->node().port()); - if(!node_manager_->QueryNodeExists(addr_str)) { + if (!node_manager_->QueryNodeExists(addr_str)) { MLOG_ERROR("Target node({}) does not exist in cluster, cannot set status", addr_str); ret = CommonErr::TargetNotFound; resp->set_ret_code(ret); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("set_node_status", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("set_node_status", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("set_node_status"); simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("set_node_status"); SEND_RESPONSE(ctx, resp); @@ -363,8 +423,11 @@ void SetNodeStatusHandler::Work(const std::shared_ptr ctx } resp->set_ret_code(ret); - simm::common::Metrics::Instance("cluster_manager").ObserveRequestDuration("set_node_status", static_cast( - std::chrono::duration_cast(std::chrono::steady_clock::now() - req_begin_ts).count())); + simm::common::Metrics::Instance("cluster_manager") + .ObserveRequestDuration("set_node_status", + static_cast(std::chrono::duration_cast( + std::chrono::steady_clock::now() - req_begin_ts) + .count())); simm::common::Metrics::Instance("cluster_manager").IncRequestsTotal("set_node_status"); if (ret != CommonErr::OK) { simm::common::Metrics::Instance("cluster_manager").IncErrorsTotal("set_node_status"); diff --git a/src/cluster_manager/cm_service.cc b/src/cluster_manager/cm_service.cc index 45f5218..4576cb6 100644 --- a/src/cluster_manager/cm_service.cc +++ b/src/cluster_manager/cm_service.cc @@ -1,16 +1,16 @@ +#include #include #include -#include #include +#include "cm_rpc_handler.h" +#include "cm_service.h" #include "common/logging/logging.h" #include "common/rpc_handlers/common_rpc_handlers.h" -#include "proto/common.pb.h" #include "proto/cm_clnt_rpcs.pb.h" +#include "proto/common.pb.h" #include "proto/ds_cm_rpcs.pb.h" -#include "cm_rpc_handler.h" -#include "cm_service.h" DECLARE_int32(cm_rpc_intra_port); DECLARE_int32(cm_rpc_inter_port); @@ -39,7 +39,7 @@ error_code_t ClusterManagerService::Start() { error_code_t ret = StartRPCServices(); if (ret != CommonErr::OK) { MLOG_ERROR("Failed to Start RPC services, ret:{}", ret); - is_running_.store(false); // revert state + is_running_.store(false); // revert state return ret; } @@ -47,13 +47,14 @@ error_code_t ClusterManagerService::Start() { std::this_thread::sleep_for(std::chrono::seconds(FLAGS_cm_cluster_init_grace_period_inSecs)); auto registered_dataservers_vec = node_manager_->GetAllNodeAddress(/* alive = true */); - for (const auto & ds_entry : registered_dataservers_vec) { + for (const auto &ds_entry : registered_dataservers_vec) { MLOG_INFO("DS({}:{}) joined into cluster", ds_entry->node_ip_, ds_entry->node_port_); } ret = shard_manager_->InitShardRoutingTable(registered_dataservers_vec); if (ret != CommonErr::OK) { MLOG_ERROR("Failed to init global shard routing table, ret:{}", ret); - is_running_.store(false); // revert state + StopRPCServices(); + is_running_.store(false); // revert state return ret; } @@ -66,7 +67,7 @@ error_code_t ClusterManagerService::Start() { // start dataservers resource query background thread // FIXME(ytji): for v0930 version, background thread(to query resource stats from ds) in // node manager module is not actived yet, so just comment init action - //node_manager_->Init(); + // node_manager_->Init(); MLOG_INFO("ClusterManager service starts successfully!"); } else { @@ -85,7 +86,7 @@ error_code_t ClusterManagerService::Stop() { error_code_t ret = StopRPCServices(); if (ret != CommonErr::OK) { MLOG_ERROR("Failed to Stop RPC services, ret:{}", ret); - is_running_.store(true); // revert state + is_running_.store(true); // revert state return ret; } @@ -97,7 +98,7 @@ error_code_t ClusterManagerService::Stop() { // stop node manager background thread // FIXME(ytji): for v0930 version, background thread(to query resource stats from ds) in // node manager module is not actived yet, so just comment stop action - //node_manager_->Stop(); + // node_manager_->Stop(); MLOG_INFO("Cluster Manager service stopped successfully..."); } else { @@ -132,17 +133,20 @@ error_code_t ClusterManagerService::StartRPCServices() { } if (auto res = create_and_start_rpc_fn(inter_rpc_service_, FLAGS_cm_rpc_inter_port, FLAGS_cm_rpc_inter_name); res != sicl::transport::SICL_SUCCESS) { + StopRPCServices(); return CmErr::InitInterRPCServiceFailed; } if (auto res = create_and_start_rpc_fn(admin_rpc_service_, FLAGS_cm_rpc_admin_port, FLAGS_cm_rpc_admin_name); res != sicl::transport::SICL_SUCCESS) { + StopRPCServices(); return CmErr::InitAdminRPCServiceFailed; } // Register handlers for in-cluster RPC requests inter_rpc_service_->RegisterHandler( static_cast(simm::cm::ClusterManagerRpcType::RPC_NEW_NODE_HANDSHAKE), - new NewNodeHandshakeHandler(inter_rpc_service_.get(), new NewNodeHandShakeRequestPB, node_manager_, shard_manager_)); + new NewNodeHandshakeHandler( + inter_rpc_service_.get(), new NewNodeHandShakeRequestPB, node_manager_, shard_manager_)); inter_rpc_service_->RegisterHandler( static_cast(simm::cm::ClusterManagerRpcType::RPC_NODE_HEARTBEAT), new NodeHeartBeatHandler(inter_rpc_service_.get(), new DataServerHeartBeatRequestPB, hb_monitor_)); @@ -164,31 +168,29 @@ error_code_t ClusterManagerService::StartRPCServices() { inter_rpc_service_.get(), new QueryShardRoutingTableBatchRequestPB, shard_manager_)); inter_rpc_service_->RegisterHandler( static_cast(simm::cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_ALL), - new RoutingTableQueryAllHandler(inter_rpc_service_.get(), new QueryShardRoutingTableAllRequestPB, shard_manager_)); + new RoutingTableQueryAllHandler( + inter_rpc_service_.get(), new QueryShardRoutingTableAllRequestPB, shard_manager_)); // Register handlers for admin RPC requests admin_rpc_service_->RegisterHandler( - static_cast(simm::common::CommonRpcType::RPC_GET_GFLAG_REQ), - new simm::common::GetGFlagHandler(admin_rpc_service_.get(), new proto::common::GetGFlagValueRequestPB)); + static_cast(simm::common::CommonRpcType::RPC_GET_GFLAG_REQ), + new simm::common::GetGFlagHandler(admin_rpc_service_.get(), new proto::common::GetGFlagValueRequestPB)); admin_rpc_service_->RegisterHandler( - static_cast(simm::common::CommonRpcType::RPC_SET_GFLAG_REQ), - new simm::common::SetGFlagHandler(admin_rpc_service_.get(), new proto::common::SetGFlagValueRequestPB)); + static_cast(simm::common::CommonRpcType::RPC_SET_GFLAG_REQ), + new simm::common::SetGFlagHandler(admin_rpc_service_.get(), new proto::common::SetGFlagValueRequestPB)); admin_rpc_service_->RegisterHandler( - static_cast(simm::common::CommonRpcType::RPC_LIST_GFLAGS_REQ), - new simm::common::ListGFlagsHandler(admin_rpc_service_.get(), new proto::common::ListAllGFlagsRequestPB)); + static_cast(simm::common::CommonRpcType::RPC_LIST_GFLAGS_REQ), + new simm::common::ListGFlagsHandler(admin_rpc_service_.get(), new proto::common::ListAllGFlagsRequestPB)); admin_rpc_service_->RegisterHandler( - static_cast(simm::common::CommonRpcType::RPC_LIST_SHARD_REQ), - new RoutingTableQueryAllHandler( - admin_rpc_service_.get(), new QueryShardRoutingTableAllRequestPB, shard_manager_) - ); + static_cast(simm::common::CommonRpcType::RPC_LIST_SHARD_REQ), + new RoutingTableQueryAllHandler( + admin_rpc_service_.get(), new QueryShardRoutingTableAllRequestPB, shard_manager_)); admin_rpc_service_->RegisterHandler( - static_cast(simm::common::CommonRpcType::RPC_LIST_NODE_REQ), - new ListNodesHandler(admin_rpc_service_.get(), new ListNodesRequestPB, node_manager_) - ); + static_cast(simm::common::CommonRpcType::RPC_LIST_NODE_REQ), + new ListNodesHandler(admin_rpc_service_.get(), new ListNodesRequestPB, node_manager_)); admin_rpc_service_->RegisterHandler( - static_cast(simm::common::CommonRpcType::RPC_SET_NODE_STATUS_REQ), - new SetNodeStatusHandler(admin_rpc_service_.get(), new SetNodeStatusRequestPB, node_manager_) - ); + static_cast(simm::common::CommonRpcType::RPC_SET_NODE_STATUS_REQ), + new SetNodeStatusHandler(admin_rpc_service_.get(), new SetNodeStatusRequestPB, node_manager_)); return CommonErr::OK; } diff --git a/src/cluster_manager/cm_service.h b/src/cluster_manager/cm_service.h index 815c4ed..59baaf1 100644 --- a/src/cluster_manager/cm_service.h +++ b/src/cluster_manager/cm_service.h @@ -7,10 +7,10 @@ #include "rpc/rpc_context.h" #include "transport/types.h" -#include "common/errcode/errcode_def.h" #include "cm_hb_monitor.h" #include "cm_node_manager.h" #include "cm_shard_manager.h" +#include "common/errcode/errcode_def.h" namespace simm { namespace cm { @@ -38,9 +38,9 @@ class ClusterManagerService { virtual ~ClusterManagerService(); ClusterManagerService(const ClusterManagerService &) = delete; - ClusterManagerService& operator=(const ClusterManagerService &) = delete; + ClusterManagerService &operator=(const ClusterManagerService &) = delete; ClusterManagerService(ClusterManagerService &&) = delete; - ClusterManagerService& operator=(ClusterManagerService &&) = delete; + ClusterManagerService &operator=(ClusterManagerService &&) = delete; public: error_code_t Init(); @@ -69,10 +69,12 @@ class ClusterManagerService { std::atomic is_running_{false}; // Node failure and rebalance - void HandleNodeFailure(const std::vector& dead_node_addresses); + void HandleNodeFailure(const std::vector &dead_node_addresses); #if defined(SIMM_UNIT_TEST) FRIEND_TEST(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs); + FRIEND_TEST(ClusterManagerServiceTest, TestQueryRoutingTableRejectsRequestsDuringGracePeriod); + FRIEND_TEST(ClusterManagerServiceTest, TestQueryRoutingTableReturnsIncompleteWhenShardUnavailable); FRIEND_TEST(ClusterManagerServiceTest, TestNodeRejoinRPC); #endif }; diff --git a/src/cluster_manager/cm_shard_manager.cc b/src/cluster_manager/cm_shard_manager.cc index 1f3ba0a..4e50fc7 100644 --- a/src/cluster_manager/cm_shard_manager.cc +++ b/src/cluster_manager/cm_shard_manager.cc @@ -2,11 +2,11 @@ #include #include +#include #include "cm_shard_manager.h" #include "common/base/assert.h" #include "common/logging/logging.h" #include "folly/Likely.h" -#include DECLARE_LOG_MODULE("cluster_manager"); @@ -134,25 +134,18 @@ error_code_t ClusterManagerShardManager::InitShardRoutingTable( } error_code_t ClusterManagerShardManager::RebalanceShardsAfterNodeFailure( - const std::vector& dead_node_addresses, - const std::vector>& alive_servers) { - MLOG_DEBUG("Starting shard rebalance after node failure, dead nodes: {}", boost::algorithm::join(dead_node_addresses, ", ")); - - if (FOLLY_UNLIKELY(alive_servers.empty())) { - return CmErr::NoAvailableDataservers; - } - - if (FOLLY_UNLIKELY(alive_servers.size() < FLAGS_dataserver_min_num)) { - return CmErr::InsufficientDataservers; - } + const std::vector &dead_node_addresses, + const std::vector> &alive_servers) { + MLOG_DEBUG("Starting shard rebalance after node failure, dead nodes: {}", + boost::algorithm::join(dead_node_addresses, ", ")); // Find all orphaned shards from dead nodes // TODO(zbhe): change data structure for faster lookup / backlink std::vector orphaned_shards; - for (const auto& entry : mShardRoutingTable) { + for (const auto &entry : mShardRoutingTable) { if (entry.second) { std::string node_addr = entry.second->node_ip_ + ":" + std::to_string(entry.second->node_port_); - for (const auto& dead_addr : dead_node_addresses) { + for (const auto &dead_addr : dead_node_addresses) { if (node_addr == dead_addr) { orphaned_shards.push_back(entry.first); break; @@ -160,59 +153,100 @@ error_code_t ClusterManagerShardManager::RebalanceShardsAfterNodeFailure( } } } - + MLOG_DEBUG("Found {} orphaned shards from dead nodes", orphaned_shards.size()); + if (orphaned_shards.empty()) { + return CommonErr::OK; + } + + if (FOLLY_UNLIKELY(alive_servers.empty())) { + (void)MarkShardsUnavailableForNodes(dead_node_addresses); + return CmErr::NoAvailableDataservers; + } + + if (FOLLY_UNLIKELY(alive_servers.size() < FLAGS_dataserver_min_num)) { + (void)MarkShardsUnavailableForNodes(dead_node_addresses); + return CmErr::InsufficientDataservers; + } + // Reassign return ReassignOrphanedShards(orphaned_shards, alive_servers); } // TODO(zbhe): reuse logic for init and rebalance error_code_t ClusterManagerShardManager::ReassignOrphanedShards( - const std::vector& orphaned_shards, - const std::vector>& alive_servers) { - + const std::vector &orphaned_shards, + const std::vector> &alive_servers) { if (orphaned_shards.empty() || alive_servers.empty()) { return CommonErr::OK; } - + // Current: Reassign based on average distribution size_t server_count = alive_servers.size(); size_t shards_per_server = orphaned_shards.size() / server_count; size_t remaining_shards = orphaned_shards.size() % server_count; - + std::vector target_shards; std::vector> target_servers; - + size_t shard_index = 0; for (size_t server_idx = 0; server_idx < server_count && shard_index < orphaned_shards.size(); ++server_idx) { size_t shards_for_this_server = shards_per_server + (server_idx < remaining_shards ? 1 : 0); - + for (size_t i = 0; i < shards_for_this_server && shard_index < orphaned_shards.size(); ++i) { target_shards.push_back(orphaned_shards[shard_index]); target_servers.push_back(alive_servers[server_idx]); - - MLOG_DEBUG("Planning to reassign shard {} to server {}:{}", - orphaned_shards[shard_index], - alive_servers[server_idx]->node_ip_, + + MLOG_DEBUG("Planning to reassign shard {} to server {}:{}", + orphaned_shards[shard_index], + alive_servers[server_idx]->node_ip_, alive_servers[server_idx]->node_port_); - + shard_index++; } } - - error_code_t ret = BatchModifyRoutingTable(target_shards, target_servers); // TODO: migrate with szzhao + + error_code_t ret = BatchModifyRoutingTable(target_shards, target_servers); // TODO: migrate with szzhao if (ret != CommonErr::OK) { MLOG_ERROR("Failed to batch modify routing table during rebalance, ret: {}", ret); return ret; } - - MLOG_DEBUG("Redistributed {} shards to {} alive servers", - target_shards.size(), server_count); - + + MLOG_DEBUG("Redistributed {} shards to {} alive servers", target_shards.size(), server_count); + return CommonErr::OK; } +error_code_t ClusterManagerShardManager::MarkShardsUnavailableForNodes( + const std::vector &target_node_addresses) { + if (target_node_addresses.empty()) { + return CommonErr::OK; + } + + std::vector target_shards; + std::vector> target_servers; + for (const auto &entry : mShardRoutingTable) { + if (!entry.second) { + continue; + } + const auto node_addr = entry.second->toString(); + for (const auto &target_node : target_node_addresses) { + if (node_addr == target_node) { + target_shards.push_back(entry.first); + target_servers.push_back(nullptr); + break; + } + } + } + + if (target_shards.empty()) { + return CommonErr::OK; + } + + return BatchModifyRoutingTable(target_shards, target_servers); +} + // TODO(ytji): update args type and implement this function error_code_t ClusterManagerShardManager::TriggerShardsMigration( const std::vector &target_shards, diff --git a/src/cluster_manager/cm_shard_manager.h b/src/cluster_manager/cm_shard_manager.h index 1c0c5b8..d8241ab 100644 --- a/src/cluster_manager/cm_shard_manager.h +++ b/src/cluster_manager/cm_shard_manager.h @@ -6,6 +6,10 @@ #include +#if defined(SIMM_UNIT_TEST) +#include +#endif + #include "common/base/common_types.h" #include "common/errcode/errcode_def.h" @@ -68,8 +72,15 @@ class ClusterManagerShardManager : public std::enable_shared_from_this &orphaned_shards, const std::vector> &alive_servers); + // Mark shards routed to target nodes as unavailable (nullptr). + error_code_t MarkShardsUnavailableForNodes(const std::vector &target_node_addresses); + // should we need this interface? - void CleanRoutingTable() { mShardRoutingTable.clear(); }; + void CleanRoutingTable() { + for (uint32_t i = 0; i < mShardNum; ++i) { + mShardRoutingTable.assign(i, nullptr); + } + }; // TODO(ytji): create policy factory to create different shard placement // policies diff --git a/src/data_server/ds_flags.cc b/src/data_server/ds_flags.cc index ee34a1f..76b5ae9 100644 --- a/src/data_server/ds_flags.cc +++ b/src/data_server/ds_flags.cc @@ -8,6 +8,9 @@ DEFINE_int64(ds_hash_seed, 0, "Hash seed for key string mapping in hash table"); DEFINE_int32(register_cooldown_sec, 10, "Duration seconds to register with Cluster Manager"); DEFINE_int32(heartbeat_cooldown_sec, 5, "Duration seconds to heartbeat with Cluster Manager"); DEFINE_uint32(cm_hb_tolerance_count, 5, "Count of allowed failed heartbeats before reconnecting to Cluster Manager"); +DEFINE_bool(ds_process_exit_cm_disconnection, + true, + "Exit data server gracefully after repeated CM heartbeat failures; if false, trigger CM re-register logic"); DEFINE_uint32(cm_connect_retry_interval_sec, 1, "Duration seconds between attempts to reconnect to Cluster Manager"); DEFINE_uint32(ds_free_memory_usable_ratio, 80, "Percentage of free memory can be used by data server"); DEFINE_int32(ds_initial_blocks, 3, "Initial number of cache blocks to pre-allocate"); @@ -16,7 +19,9 @@ DEFINE_uint32(ds_bg_evict_interval_ms, 500, "background chunk level eviction che DEFINE_uint32(ds_bg_evict_prefetch_factor, 1, "background chunk level eviction prefetch factor"); DEFINE_double(ds_bg_evict_trigger_threshold, 0.85, "background chunk level eviction trigger threshold"); DEFINE_uint32(ds_bg_evict_slab_class_num, 3, "background chunk level eviction for how much slab class number"); -DEFINE_uint64(ds_bg_evict_chunk_cooldown_ms, 3600000, "background chunk level eviction for empty chunk cooldown duration"); +DEFINE_uint64(ds_bg_evict_chunk_cooldown_ms, + 3600000, + "background chunk level eviction for empty chunk cooldown duration"); DEFINE_bool(ds_clean_stale_block, false, "clean stale block before initialization"); // Used in k8s scenarios where /proc/meminfo is inaccurate, see doc comments for GetMemoryFreeToUse DEFINE_uint64(memory_limit_bytes, 0, "Memory limit for data servers, 0 means to check /proc/meminfo"); diff --git a/src/data_server/kv_cache_pool.h b/src/data_server/kv_cache_pool.h index f795dfa..b751683 100644 --- a/src/data_server/kv_cache_pool.h +++ b/src/data_server/kv_cache_pool.h @@ -1,6 +1,7 @@ #pragma once #include +#include #include #include #include @@ -9,13 +10,17 @@ #include #include #include -#include + +#if defined(SIMM_UNIT_TEST) +#include +#endif + #include "common/base/common_types.h" #include "common/base/consts.h" #include "data_server/ds_common.h" #include "data_server/ds_memory_allocator.h" -#include "transport/mempool.h" #include "folly/executors/CPUThreadPoolExecutor.h" +#include "transport/mempool.h" namespace simm { namespace ds { @@ -122,9 +127,15 @@ constexpr size_t slab_index(SlabClass sc) { } // Slab sizes in bytes -constexpr size_t SLAB_CLASS_SIZES[] = {0, 1ULL << 12, 1ULL << 15, - 1ULL << 18, 1ULL << 20, 1ULL << 22, - 1ULL << 23, 1ULL << 24, (1ULL << 26) - (20ULL << 10)}; +constexpr size_t SLAB_CLASS_SIZES[] = {0, + 1ULL << 12, + 1ULL << 15, + 1ULL << 18, + 1ULL << 20, + 1ULL << 22, + 1ULL << 23, + 1ULL << 24, + (1ULL << 26) - (20ULL << 10)}; // Get size in bytes of a slab class constexpr size_t slab_size(SlabClass sc) { @@ -136,7 +147,8 @@ constexpr size_t max_slab_cnt(size_t slab_sz) { for (size_t n = 1;; n++) { size_t meta_aligned = align_up(n * META_SIZE, PAGE_SIZE); size_t total = HEADER_SIZE + meta_aligned + n * slab_sz; - if (total > CHUNK_SIZE) break; + if (total > CHUNK_SIZE) + break; max_n = n; } return max_n; @@ -240,7 +252,7 @@ struct KVMeta { uint64_t value_crc; // 8 bytes uint32_t ctime; // 4 bytes uint32_t ttl; // 4 bytes - char* key_ptr; // 8 bytes (tmp pointer to key string) + char *key_ptr; // 8 bytes (tmp pointer to key string) uint32_t shard_id; // 4 bytes uint8_t reserved[20]; // 20 bytes char key[448]; // 448 bytes (total 512 bytes) @@ -322,7 +334,7 @@ class KVCachePool { } private: - void clean_stale_block(const std::string& shm_path); + void clean_stale_block(const std::string &shm_path); bool init_block(size_t idx); @@ -362,7 +374,7 @@ class KVCachePool { FREETOUSE = 14, ALLOCATE = 15, }; - // each block uses 64 bits to represent status of all chunks(16) + // each block uses 64 bits to represent status of all chunks(16) std::deque> chunk_status_; // each chunk uses a counter to write down assigned number diff --git a/src/data_server/kv_rpc_service.cc b/src/data_server/kv_rpc_service.cc index 6a2fead..7b0b2ff 100644 --- a/src/data_server/kv_rpc_service.cc +++ b/src/data_server/kv_rpc_service.cc @@ -1,5 +1,6 @@ #include #include +#include #include #include @@ -10,11 +11,11 @@ #include "cluster_manager/cm_rpc_handler.h" #include "common/base/consts.h" +#include "common/context/context.h" #include "common/errcode/errcode_def.h" #include "common/logging/logging.h" #include "common/rpc_handlers/common_rpc_handlers.h" #include "common/trace/trace.h" -#include "common/context/context.h" #include "common/utils/ip_util.h" #include "common/utils/k8s_util.h" #include "common/utils/sys_util.h" @@ -36,6 +37,7 @@ DECLARE_string(cm_namespace); DECLARE_string(cm_svc_name); DECLARE_string(cm_port_name); DECLARE_uint32(cm_hb_tolerance_count); +DECLARE_bool(ds_process_exit_cm_disconnection); DECLARE_bool(simm_enable_trace); namespace simm { @@ -113,8 +115,8 @@ error_code_t KVRpcService::Init() { system_level_free_memory = static_cast(result); } auto cache_bound_bytes = system_level_free_memory * FLAGS_ds_free_memory_usable_ratio / 100; - int ret = cache_pool_->init(cache_bound_bytes, io_service->GetMempool(), - cache_evictor_.get(), FLAGS_ds_initial_blocks); + int ret = + cache_pool_->init(cache_bound_bytes, io_service->GetMempool(), cache_evictor_.get(), FLAGS_ds_initial_blocks); if (ret) { MLOG_ERROR("KVRpcService::Init new cache pool failed, ret:{}", ret); return DsErr::InitFailed; @@ -176,7 +178,8 @@ error_code_t KVRpcService::StartRPCServices() { } res = mgt_service_->Start(FLAGS_mgt_service_port); if (static_cast(res) != sicl::transport::Result::SICL_SUCCESS) { - MLOG_ERROR("Failed to start RPC management service on port:{}, res:{}", FLAGS_mgt_service_port, std::to_string(res)); + MLOG_ERROR( + "Failed to start RPC management service on port:{}, res:{}", FLAGS_mgt_service_port, std::to_string(res)); return DsErr::InitRPCServiceFailed; } res = admin_rpc_service_->Start(FLAGS_ds_rpc_admin_port); @@ -193,29 +196,29 @@ error_code_t KVRpcService::RegisterHandlers() { res = io_service_->RegisterHandler(static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_GET), new KVGetHandler(this, new KVGetRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_GET({}) failed", - static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_GET)); + MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_GET({}) failed", + static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_GET)); return DsErr::RegisterRPCHandlerFailed; } res = io_service_->RegisterHandler(static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_PUT), new KVPutHandler(this, new KVPutRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_PUT({}) failed", - static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_PUT)); + MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_PUT({}) failed", + static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_PUT)); return DsErr::RegisterRPCHandlerFailed; } res = io_service_->RegisterHandler(static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_DEL), new KVDelHandler(this, new KVDelRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_DEL({}) failed", - static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_DEL)); + MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_DEL({}) failed", + static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_DEL)); return DsErr::RegisterRPCHandlerFailed; } res = io_service_->RegisterHandler(static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_LOOKUP), new KVLookupHandler(this, new KVLookupRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_LOOKUP({}) failed", - static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_LOOKUP)); + MLOG_ERROR("RegisterHandler for RPC_CLIENT_KV_LOOKUP({}) failed", + static_cast(ds::KVServerRpcType::RPC_CLIENT_KV_LOOKUP)); return DsErr::RegisterRPCHandlerFailed; } // control plane @@ -223,8 +226,8 @@ error_code_t KVRpcService::RegisterHandlers() { static_cast(cm::ClusterManagerRpcType::RPC_DATASERVER_RESOURCE_QUERY), new MgtResourceHandler(this, new DataServerResourceRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_DATASERVER_RESOURCE_QUERY({}) failed", - static_cast(cm::ClusterManagerRpcType::RPC_DATASERVER_RESOURCE_QUERY)); + MLOG_ERROR("RegisterHandler for RPC_DATASERVER_RESOURCE_QUERY({}) failed", + static_cast(cm::ClusterManagerRpcType::RPC_DATASERVER_RESOURCE_QUERY)); return DsErr::RegisterRPCHandlerFailed; } // admin @@ -232,24 +235,24 @@ error_code_t KVRpcService::RegisterHandlers() { static_cast(simm::common::CommonRpcType::RPC_GET_GFLAG_REQ), new simm::common::GetGFlagHandler(admin_rpc_service_.get(), new proto::common::GetGFlagValueRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_GET_GFLAG_REQ({}) failed", - static_cast(simm::common::CommonRpcType::RPC_GET_GFLAG_REQ)); + MLOG_ERROR("RegisterHandler for RPC_GET_GFLAG_REQ({}) failed", + static_cast(simm::common::CommonRpcType::RPC_GET_GFLAG_REQ)); return DsErr::RegisterRPCHandlerFailed; } res = admin_rpc_service_->RegisterHandler( static_cast(simm::common::CommonRpcType::RPC_SET_GFLAG_REQ), new simm::common::SetGFlagHandler(admin_rpc_service_.get(), new proto::common::SetGFlagValueRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_SET_GFLAG_REQ({}) failed", - static_cast(simm::common::CommonRpcType::RPC_SET_GFLAG_REQ)); + MLOG_ERROR("RegisterHandler for RPC_SET_GFLAG_REQ({}) failed", + static_cast(simm::common::CommonRpcType::RPC_SET_GFLAG_REQ)); return DsErr::RegisterRPCHandlerFailed; } res = admin_rpc_service_->RegisterHandler( static_cast(simm::common::CommonRpcType::RPC_LIST_GFLAGS_REQ), new simm::common::ListGFlagsHandler(admin_rpc_service_.get(), new proto::common::ListAllGFlagsRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_LIST_GFLAGS_REQ({}) failed", - static_cast(simm::common::CommonRpcType::RPC_LIST_GFLAGS_REQ)); + MLOG_ERROR("RegisterHandler for RPC_LIST_GFLAGS_REQ({}) failed", + static_cast(simm::common::CommonRpcType::RPC_LIST_GFLAGS_REQ)); return DsErr::RegisterRPCHandlerFailed; } #ifdef SIMM_ENABLE_TRACE @@ -257,8 +260,8 @@ error_code_t KVRpcService::RegisterHandlers() { static_cast(simm::common::CommonRpcType::RPC_TRACE_TOGGLE_REQ), new simm::common::TraceToggleHandler(admin_rpc_service_.get(), new proto::common::TraceToggleRequestPB)); if (!res) { - MLOG_ERROR("RegisterHandler for RPC_TRACE_TOGGLE_REQ({}) failed", - static_cast(simm::common::CommonRpcType::RPC_TRACE_TOGGLE_REQ)); + MLOG_ERROR("RegisterHandler for RPC_TRACE_TOGGLE_REQ({}) failed", + static_cast(simm::common::CommonRpcType::RPC_TRACE_TOGGLE_REQ)); return DsErr::RegisterRPCHandlerFailed; } #endif @@ -299,7 +302,7 @@ void KVRpcService::KeepAlive() { while (!is_registered_) { RegisterOnCluster(); register_condv_.try_wait_for(std::chrono::seconds(FLAGS_register_cooldown_sec)); - register_condv_.reset(); // folly::Bation should be reset before reuse + register_condv_.reset(); // folly::Bation should be reset before reuse } while (!cm_ready_) { register_condv_.reset(); @@ -384,15 +387,18 @@ void KVRpcService::RegisterToRestartedManager() { auto response = dynamic_cast(rsp); auto ret_code = response->ret_code(); if (ret_code == CommonErr::OK) { - heartbeat_failure_count.store(0); + heartbeat_failure_count_.store(0); cm_ready_ = true; register_condv_.post(); } else { MLOG_ERROR("RegisterToRestartedManager return not OK: {}", ret_code); } } else { - MLOG_ERROR( - "Failed to RegisterToRestartedManager to Cluster Manager {}:{}: {} ({})", FLAGS_cm_primary_node_ip, FLAGS_cm_rpc_inter_port, ctx->ErrorText(), ctx->ErrorCode()); + MLOG_ERROR("Failed to RegisterToRestartedManager to Cluster Manager {}:{}: {} ({})", + FLAGS_cm_primary_node_ip, + FLAGS_cm_rpc_inter_port, + ctx->ErrorText(), + ctx->ErrorCode()); } }; @@ -421,20 +427,12 @@ void KVRpcService::HeartBeatToCluster() { const std::shared_ptr ctx) { if (!ctx->Failed()) { auto response = dynamic_cast(rsp); - if (response->ret_code() == CommonErr::OK) { - // nothing - } else { + OnHeartbeatResult(false, response->ret_code()); + if (response->ret_code() != CommonErr::OK) { MLOG_ERROR("HeartBeatToCluster return not OK"); } } else { - heartbeat_failure_count.fetch_add(1); - MLOG_ERROR("Failed to HeartbeatToCluster to Cluster Manager {}:{} (counter={})", - FLAGS_cm_primary_node_ip, - FLAGS_cm_rpc_inter_port, - heartbeat_failure_count.load()); - if (heartbeat_failure_count.load() >= FLAGS_cm_hb_tolerance_count) { - cm_ready_.store(false); - } + OnHeartbeatResult(true, CommonErr::InvalidState); } }; @@ -447,6 +445,39 @@ void KVRpcService::HeartBeatToCluster() { heartbeat_done); } +void KVRpcService::OnHeartbeatResult(bool rpc_failed, error_code_t ret_code) { + if (!rpc_failed && ret_code == CommonErr::OK) { + // reset failure counter after one HB request succeed + heartbeat_failure_count_.store(0); + return; + } + + if (!rpc_failed) { + return; + } + + const auto failure_count = heartbeat_failure_count_.fetch_add(1) + 1; + MLOG_ERROR("Failed to HeartbeatToCluster to Cluster Manager {}:{} (counter={})", + FLAGS_cm_primary_node_ip, + FLAGS_cm_rpc_inter_port, + failure_count); + if (failure_count >= FLAGS_cm_hb_tolerance_count) { + heartbeat_failure_count_.store(0); + if (FLAGS_ds_process_exit_cm_disconnection) { + MLOG_CRITICAL("Exit data server for disconnection(HB failed {} times) with cluster manager", + FLAGS_cm_hb_tolerance_count); + HandleClusterManagerDisconnect(); + } else { + MLOG_WARN("HB failed {} times with cluster manager, switch to CM re-register flow", FLAGS_cm_hb_tolerance_count); + cm_ready_.store(false); + } + } +} + +void KVRpcService::HandleClusterManagerDisconnect() { + cluster_disconnect_handler_(); +} + error_code_t KVRpcService::KVGet(std::shared_ptr ctx, const KVGetRequestPB *req, KVEntryIntrusivePtr &entry) { @@ -462,11 +493,11 @@ error_code_t KVRpcService::KVGet(std::shared_ptr ctx, return DsErr::KeyNotFound; } } - std::unique_ptr key_meta( - object_pool_->AcquireKey(), KVObjectPool::KVMetaDeleter{}); + std::unique_ptr key_meta(object_pool_->AcquireKey(), + KVObjectPool::KVMetaDeleter{}); key_meta->key_hash = key_hash; key_meta->key_len = key_str.length(); - key_meta->key_ptr= (char *)key_str.c_str(); + key_meta->key_ptr = (char *)key_str.c_str(); KVHashKey key(key_meta.get()); KVEntryIntrusivePtr value; bool found = table->Find(key, value); @@ -528,11 +559,11 @@ error_code_t KVRpcService::KVPut(std::shared_ptr ctx, all_tables_[shard_id] = table; } } - std::unique_ptr key_meta( - object_pool_->AcquireKey(), KVObjectPool::KVMetaDeleter{}); + std::unique_ptr key_meta(object_pool_->AcquireKey(), + KVObjectPool::KVMetaDeleter{}); key_meta->key_hash = key_hash; key_meta->key_len = key_str.length(); - key_meta->key_ptr= (char *)key_str.c_str(); + key_meta->key_ptr = (char *)key_str.c_str(); KVHashKey key(key_meta.get()); KVEntryIntrusivePtr value; bool found = table->Find(key, value); @@ -558,8 +589,8 @@ error_code_t KVRpcService::KVPut(std::shared_ptr ctx, value->ref_cnt_lock.ClearExclusive(); return DsErr::CachePoolAllocateFailed; } - shard_used_bytes_[shard_id].fetch_add( - SLAB_CLASS_SIZES[value->slab_info.slab_class] + META_SIZE, std::memory_order_relaxed); + shard_used_bytes_[shard_id].fetch_add(SLAB_CLASS_SIZES[value->slab_info.slab_class] + META_SIZE, + std::memory_order_relaxed); auto [meta, _] = KVCachePool::GetBufferPair(&value->slab_info); meta->key_len = key_str.length(); // meta->value_len = req->val_len(); since allocate success would write down value_len @@ -583,8 +614,8 @@ error_code_t KVRpcService::KVPut(std::shared_ptr ctx, value->status.store(static_cast(KVStatus::KV_CLEAR), std::memory_order_release); table->Remove(key); cache_pool_->free(&value->slab_info); - shard_used_bytes_[shard_id].fetch_sub( - SLAB_CLASS_SIZES[value->slab_info.slab_class] + META_SIZE, std::memory_order_relaxed); + shard_used_bytes_[shard_id].fetch_sub(SLAB_CLASS_SIZES[value->slab_info.slab_class] + META_SIZE, + std::memory_order_relaxed); value->ref_cnt_lock.ClearExclusive(); return DsErr::HashTableOperationError; } @@ -606,8 +637,8 @@ void KVRpcService::KVPutFailedRewind(uint32_t shard_id, KVEntryIntrusivePtr &ent KVHashKey key(meta); table->Remove(key); cache_pool_->free(&entry->slab_info); - shard_used_bytes_[shard_id].fetch_sub( - SLAB_CLASS_SIZES[entry->slab_info.slab_class] + META_SIZE, std::memory_order_relaxed); + shard_used_bytes_[shard_id].fetch_sub(SLAB_CLASS_SIZES[entry->slab_info.slab_class] + META_SIZE, + std::memory_order_relaxed); entry->ref_cnt_lock.ClearExclusive(); } @@ -625,8 +656,7 @@ void KVRpcService::KVPutSuccessHooks(KVEntryIntrusivePtr &entry) { #endif } -error_code_t KVRpcService::KVDel(std::shared_ptr ctx, - const KVDelRequestPB *req) { +error_code_t KVRpcService::KVDel(std::shared_ptr ctx, const KVDelRequestPB *req) { uint32_t shard_id = req->shard_id(); auto &key_str = req->key(); uint64_t key_hash = Hasher(key_str); @@ -639,11 +669,11 @@ error_code_t KVRpcService::KVDel(std::shared_ptr ctx, return DsErr::KeyNotFound; } } - std::unique_ptr key_meta( - object_pool_->AcquireKey(), KVObjectPool::KVMetaDeleter{}); + std::unique_ptr key_meta(object_pool_->AcquireKey(), + KVObjectPool::KVMetaDeleter{}); key_meta->key_hash = key_hash; key_meta->key_len = key_str.length(); - key_meta->key_ptr= (char *)key_str.c_str(); + key_meta->key_ptr = (char *)key_str.c_str(); KVHashKey key(key_meta.get()); KVEntryIntrusivePtr value; bool found = table->Find(key, value); @@ -708,8 +738,8 @@ error_code_t KVRpcService::KVDel(std::shared_ptr ctx, if (cache_pool_->is_active(&value->slab_info)) { cache_pool_->free(&value->slab_info); } - shard_used_bytes_[shard_id].fetch_sub( - SLAB_CLASS_SIZES[value->slab_info.slab_class] + META_SIZE, std::memory_order_relaxed); + shard_used_bytes_[shard_id].fetch_sub(SLAB_CLASS_SIZES[value->slab_info.slab_class] + META_SIZE, + std::memory_order_relaxed); value->ref_cnt_lock.ClearExclusive(); return CommonErr::OK; } else { @@ -721,8 +751,7 @@ error_code_t KVRpcService::KVDel(std::shared_ptr ctx, return DsErr::KVStatusNotValid; } -error_code_t KVRpcService::KVLookup(std::shared_ptr ctx, - const KVLookupRequestPB *req) { +error_code_t KVRpcService::KVLookup(std::shared_ptr ctx, const KVLookupRequestPB *req) { uint32_t shard_id = req->shard_id(); auto &key_str = req->key(); uint64_t key_hash = Hasher(key_str); @@ -735,11 +764,11 @@ error_code_t KVRpcService::KVLookup(std::shared_ptr c return DsErr::KeyNotFound; } } - std::unique_ptr key_meta( - object_pool_->AcquireKey(), KVObjectPool::KVMetaDeleter{}); + std::unique_ptr key_meta(object_pool_->AcquireKey(), + KVObjectPool::KVMetaDeleter{}); key_meta->key_hash = key_hash; key_meta->key_len = key_str.length(); - key_meta->key_ptr= (char *)key_str.c_str(); + key_meta->key_ptr = (char *)key_str.c_str(); KVHashKey key(key_meta.get()); KVEntryIntrusivePtr value; bool found = table->Find(key, value); diff --git a/src/data_server/kv_rpc_service.h b/src/data_server/kv_rpc_service.h index f999281..8593bb2 100644 --- a/src/data_server/kv_rpc_service.h +++ b/src/data_server/kv_rpc_service.h @@ -2,13 +2,20 @@ #include #include +#include +#include +#include #include #include +#if defined(SIMM_UNIT_TEST) +#include +#endif + #include "common/base/common_types.h" #include "common/base/consts.h" -#include "common/errcode/errcode_def.h" #include "common/context/context.h" +#include "common/errcode/errcode_def.h" #include "folly/synchronization/Baton.h" #include "rpc/connection.h" #include "rpc/rpc.h" @@ -81,12 +88,14 @@ class KVRpcService { void RegisterOnCluster(); void RegisterToRestartedManager(); void HeartBeatToCluster(); + void OnHeartbeatResult(bool rpc_failed, error_code_t ret_code); + void HandleClusterManagerDisconnect(); private: - std::unique_ptr io_service_{nullptr}; // for client io requests - std::unique_ptr mgmt_client_{nullptr}; // for service controls - std::unique_ptr mgt_service_{nullptr}; // for service controls - std::unique_ptr admin_rpc_service_{nullptr}; // for maintainence scenario + std::unique_ptr io_service_{nullptr}; // for client io requests + std::unique_ptr mgmt_client_{nullptr}; // for service controls + std::unique_ptr mgt_service_{nullptr}; // for service controls + std::unique_ptr admin_rpc_service_{nullptr}; // for maintainence scenario std::unique_ptr object_pool_{nullptr}; std::unique_ptr cache_pool_{nullptr}; @@ -102,8 +111,14 @@ class KVRpcService { folly::Baton<> register_condv_; std::mutex heartbeat_mutex_; std::condition_variable heartbeat_condv_; - std::atomic heartbeat_failure_count{0}; + std::atomic heartbeat_failure_count_{0}; std::unique_ptr keepalive_thread_{nullptr}; + std::function cluster_disconnect_handler_{[]() { + // raise SIGTERM to trigger handler to do data_server clean destruction + if (std::raise(SIGTERM) != 0) { + std::_Exit(EXIT_FAILURE); + } + }}; std::string local_ip_; std::deque> shard_used_bytes_; @@ -111,6 +126,11 @@ class KVRpcService { friend class KVCacheEvictor; #if defined(SIMM_UNIT_TEST) FRIEND_TEST(KVServiceTest, TestClientHandlers); + FRIEND_TEST(KVServiceLightTest, TestHeartbeatFailureCountResetOnSuccess); + FRIEND_TEST(KVServiceLightTest, TestClusterManagerDisconnectHandlerInvokedOnToleranceReached); + FRIEND_TEST(KVServiceLightTest, TestHeartbeatFailureToleranceTriggersReconnectWhenExitDisabled); + FRIEND_TEST(KVServiceLightTest, TestClusterManagerDisconnectSignalPathRaisesSigterm); + FRIEND_TEST(KVServiceLightTest, TestShmAllocatorDestructorReleasesSharedMemory); #endif }; diff --git a/tests/client/test_clnt_messenger.cc b/tests/client/test_clnt_messenger.cc index 32c57d1..8dc4869 100644 --- a/tests/client/test_clnt_messenger.cc +++ b/tests/client/test_clnt_messenger.cc @@ -1,37 +1,339 @@ +#include +#include + #include #include -#include - -#include "simm/simm_common.h" -#include "simm/simm_kv.h" #include "client/clnt_messenger.h" -#include "common/base/memory.h" - -DECLARE_string(cm_primary_node_ip); -DECLARE_int32(cm_inter_port); +#include "rpc/connection.h" +#include "rpc/rpc.h" +#include "rpc/rpc_context.h" namespace simm { namespace clnt { -constexpr size_t kOneMB = 1 << 20; +namespace { -using simm::common::MemBlock; +class FakeConnection : public sicl::rpc::Connection { + public: + explicit FakeConnection(std::string name) : name_(std::move(name)) {} -class ClientMessengerTest : public ::testing::Test { - protected: - void SetUp() override { - auto ret = ClientMessenger::Instance().Init(); - ASSERT_EQ(ret, CommonErr::OK); + sicl::transport::Result send(const google::protobuf::Message &, + sicl::transport::MsgTag, + sicl::SendCallback, + sicl::transport::Channel *) const override { + return sicl::transport::SICL_SUCCESS; + } + + sicl::transport::Result read(const void *, + const size_t, + const uint64_t, + const std::vector &, + sicl::ReadCallback, + sicl::RequestParam) const override { + return sicl::transport::SICL_SUCCESS; + } + + sicl::transport::Result write(const void *, + const size_t, + const uint64_t, + const std::vector &, + sicl::WriteCallback, + sicl::RequestParam) const override { + return sicl::transport::SICL_SUCCESS; + } + + sicl::transport::Result write_imm(const uint32_t, sicl::WriteCallback) override { + return sicl::transport::SICL_SUCCESS; } - void TearDown() override {} + bool addChannel(sicl::transport::Channel *) override { return true; } + sicl::transport::Channel *getAChannel() const override { return nullptr; } + size_t getChannelCount() const override { return 1; } + sicl::transport::Result remove(const sicl::transport::Channel *) override { return sicl::transport::SICL_SUCCESS; } + bool hasChannel(const sicl::transport::Channel *) const override { return false; } + std::string toString() const override { return name_; } + void SendResponse(const google::protobuf::Message &, + std::shared_ptr, + sicl::rpc::RpcResponseDoneFn) const override {} + uuids::uuid getGroupID() const override { return {}; } + std::pair getIPPort() const override { return {"127.0.0.1", 12345}; } + sicl::transport::Result recv_large(sicl::rpc::SiRPC &, + void *, + size_t, + std::function) const override { + return sicl::transport::SICL_SUCCESS; + } + + private: + std::string name_; }; -TEST_F(ClientMessengerTest, TestGetServerAddress) { - std::string key = "test"; - // auto addr = ClientMessenger::Instance().GetServerAddress(key); - // std::cout << addr << std::endl; +std::shared_ptr MakeFailedRpcContext(int error_code) { + sicl::rpc::RpcContext *ctx_raw = nullptr; + sicl::rpc::RpcContext::newInstance(ctx_raw); + auto ctx = std::shared_ptr(ctx_raw); + ctx->SetError(error_code, "mock failure"); + return ctx; +} + +std::shared_ptr BuildRoutingResponse( + const std::vector>> &routing_entries) { + auto resp = std::make_shared(); + resp->set_ret_code(CommonErr::OK); + for (const auto &[addr, shard_ids] : routing_entries) { + auto node_addr = simm::common::NodeAddress::ParseFromString(addr); + EXPECT_TRUE(node_addr.has_value()); + if (!node_addr.has_value()) { + continue; + } + auto *entry = resp->add_shard_info(); + entry->mutable_data_server_address()->set_ip(node_addr->node_ip_); + entry->mutable_data_server_address()->set_port(node_addr->node_port_); + for (auto shard_id : shard_ids) { + entry->add_shard_ids(shard_id); + } + } + return resp; +} + +} // namespace + +class ClientMessengerTestPeer { + public: + static void ResetState() { + auto &messenger = ClientMessenger::Instance(); + messenger.shard_table_.clear(); + messenger.ds_conn_ctxs_.clear(); + messenger.cm_addr_.clear(); + messenger.initialized_ = false; + messenger.test_get_cm_address_hook_ = nullptr; + messenger.test_route_query_hook_ = nullptr; + messenger.test_build_connection_hook_ = nullptr; + } + + static void InstallDsContext(const std::string &addr, + bool active, + size_t gen_num, + const std::vector &shard_ids, + std::shared_ptr connection = nullptr) { + auto &messenger = ClientMessenger::Instance(); + auto ds_ctx = messenger.GetOrCreateConnectionContext(addr); + ds_ctx->StoreConnection(std::move(connection)); + ds_ctx->active.store(active); + ds_ctx->gen_num.store(gen_num); + for (auto shard_id : shard_ids) { + messenger.shard_table_.insert_or_assign(shard_id, ds_ctx); + } + } + + static bool HasDsContext(const std::string &addr) { + auto &messenger = ClientMessenger::Instance(); + return messenger.ds_conn_ctxs_.find(addr) != messenger.ds_conn_ctxs_.end(); + } + + static bool IsDsActive(const std::string &addr) { + auto &messenger = ClientMessenger::Instance(); + auto it = messenger.ds_conn_ctxs_.find(addr); + return it != messenger.ds_conn_ctxs_.end() && it->second->active.load(); + } + + static std::shared_ptr GetConnection(const std::string &addr) { + auto &messenger = ClientMessenger::Instance(); + auto it = messenger.ds_conn_ctxs_.find(addr); + return it == messenger.ds_conn_ctxs_.end() ? nullptr : it->second->LoadConnection(); + } + + static void SetConnection(const std::string &addr, std::shared_ptr connection) { + ClientMessenger::Instance().GetOrCreateConnectionContext(addr)->StoreConnection(std::move(connection)); + } + + static void MarkConnectionActive(const std::string &addr, bool active, size_t gen_num) { + auto ds_ctx = ClientMessenger::Instance().GetOrCreateConnectionContext(addr); + ds_ctx->active.store(active); + ds_ctx->gen_num.store(gen_num); + } + + static std::string ShardOwner(uint16_t shard_id) { + auto &messenger = ClientMessenger::Instance(); + auto it = messenger.shard_table_.find(shard_id); + return it == messenger.shard_table_.end() ? "" : it->second->ip_port; + } + + static void PruneConnections(const std::vector &live_servers) { + ClientMessenger::Instance().PruneStaleConnectionContexts( + std::unordered_set(live_servers.begin(), live_servers.end())); + } + + static void HandleAsyncFailure(std::shared_ptr rpc_ctx, + const std::string &request_addr, + uint16_t shard_id, + size_t old_conn_gen_num) { + auto &messenger = ClientMessenger::Instance(); + auto it = messenger.ds_conn_ctxs_.find(request_addr); + if (it != messenger.ds_conn_ctxs_.end()) { + messenger.HandleAsyncRequestFailure(std::move(rpc_ctx), it->second, shard_id, old_conn_gen_num); + } + } + + static void SetGetCmAddressHook(std::function hook) { + ClientMessenger::Instance().test_get_cm_address_hook_ = std::move(hook); + } + + static void SetRouteQueryHook( + std::function>(const std::string &)> + hook) { + ClientMessenger::Instance().test_route_query_hook_ = std::move(hook); + } + + static void SetBuildConnectionHook(std::function hook) { + ClientMessenger::Instance().test_build_connection_hook_ = std::move(hook); + } + + static error_code_t ReInit() { return ClientMessenger::Instance().ReInit(); } + + static bool IsInitialized() { return ClientMessenger::Instance().initialized_; } + + static std::string CmAddr() { return ClientMessenger::Instance().cm_addr_; } +}; + +class ClientMessengerUnitTest : public ::testing::Test { + protected: + void SetUp() override { ClientMessengerTestPeer::ResetState(); } + + void TearDown() override { ClientMessengerTestPeer::ResetState(); } +}; + +TEST_F(ClientMessengerUnitTest, PruneStaleConnectionsRemovesDeadDataservers) { + auto conn_a = std::make_shared("a"); + auto conn_b = std::make_shared("b"); + auto conn_c = std::make_shared("c"); + + ClientMessengerTestPeer::InstallDsContext("10.0.0.1:1001", true, 1, {0}, conn_a); + ClientMessengerTestPeer::InstallDsContext("10.0.0.2:1002", true, 1, {1}, conn_b); + ClientMessengerTestPeer::InstallDsContext("10.0.0.3:1003", true, 1, {2}, conn_c); + + ClientMessengerTestPeer::PruneConnections({"10.0.0.2:1002", "10.0.0.3:1003"}); + + EXPECT_FALSE(ClientMessengerTestPeer::HasDsContext("10.0.0.1:1001")); + EXPECT_TRUE(ClientMessengerTestPeer::HasDsContext("10.0.0.2:1002")); + EXPECT_TRUE(ClientMessengerTestPeer::HasDsContext("10.0.0.3:1003")); +} + +TEST_F(ClientMessengerUnitTest, AsyncFailureUsesRequestTimeDataserverContext) { + auto old_conn = std::make_shared("old"); + auto new_conn = std::make_shared("new"); + + ClientMessengerTestPeer::InstallDsContext("10.0.0.1:1001", true, 7, {}, old_conn); + ClientMessengerTestPeer::InstallDsContext("10.0.0.2:1002", true, 7, {0}, new_conn); + + auto rpc_ctx = MakeFailedRpcContext(sicl::transport::SICL_ERR_INVALID_STATE); + ClientMessengerTestPeer::HandleAsyncFailure(rpc_ctx, "10.0.0.1:1001", 0, 7); + + EXPECT_FALSE(ClientMessengerTestPeer::IsDsActive("10.0.0.1:1001")); + EXPECT_TRUE(ClientMessengerTestPeer::IsDsActive("10.0.0.2:1002")); + EXPECT_EQ(ClientMessengerTestPeer::ShardOwner(0), "10.0.0.2:1002"); +} + +TEST_F(ClientMessengerUnitTest, ConnectionAccessorsStayConsistentDuringConcurrentSwap) { + auto conn_a = std::make_shared("a"); + auto conn_b = std::make_shared("b"); + ClientMessengerTestPeer::InstallDsContext("10.0.0.1:1001", true, 1, {0}, conn_a); + + std::atomic stop{false}; + std::thread writer([&]() { + for (int i = 0; i < 20000; ++i) { + ClientMessengerTestPeer::SetConnection("10.0.0.1:1001", (i % 2 == 0) ? conn_a : conn_b); + } + stop.store(true); + }); + + std::thread reader([&]() { + while (!stop.load()) { + auto current = ClientMessengerTestPeer::GetConnection("10.0.0.1:1001"); + ASSERT_TRUE(current == conn_a || current == conn_b); + } + }); + + writer.join(); + reader.join(); + + auto final_conn = ClientMessengerTestPeer::GetConnection("10.0.0.1:1001"); + EXPECT_TRUE(final_conn == conn_a || final_conn == conn_b); +} + +TEST_F(ClientMessengerUnitTest, ReInitRefreshesRoutingBuildsConnectionsAndPrunesStaleServers) { + auto healthy_conn = std::make_shared("healthy"); + auto inactive_conn = std::make_shared("inactive"); + auto stale_conn = std::make_shared("stale"); + + ClientMessengerTestPeer::InstallDsContext("10.0.0.1:1001", true, 10, {0}, healthy_conn); + ClientMessengerTestPeer::InstallDsContext("10.0.0.2:1002", false, 4, {1}, inactive_conn); + ClientMessengerTestPeer::InstallDsContext("10.0.0.9:1999", true, 2, {9}, stale_conn); + + std::vector built_servers_order; + std::unordered_set built_servers; + ClientMessengerTestPeer::SetGetCmAddressHook([]() { return std::string("10.0.0.100:9000"); }); + ClientMessengerTestPeer::SetRouteQueryHook([](const std::string &) { + return std::make_pair( + CommonErr::OK, BuildRoutingResponse({{"10.0.0.1:1001", {0}}, {"10.0.0.2:1002", {1}}, {"10.0.0.3:1003", {2}}})); + }); + ClientMessengerTestPeer::SetBuildConnectionHook([&](const std::string &addr) { + built_servers.insert(addr); + built_servers_order.push_back(addr); + auto conn = std::make_shared(addr); + ClientMessengerTestPeer::SetConnection(addr, conn); + ClientMessengerTestPeer::MarkConnectionActive(addr, true, 100 + built_servers_order.size()); + return CommonErr::OK; + }); + + EXPECT_EQ(ClientMessengerTestPeer::ReInit(), CommonErr::OK); + + EXPECT_TRUE(ClientMessengerTestPeer::IsInitialized()); + EXPECT_EQ(ClientMessengerTestPeer::CmAddr(), "10.0.0.100:9000"); + EXPECT_FALSE(ClientMessengerTestPeer::HasDsContext("10.0.0.9:1999")); + EXPECT_TRUE(ClientMessengerTestPeer::HasDsContext("10.0.0.1:1001")); + EXPECT_TRUE(ClientMessengerTestPeer::HasDsContext("10.0.0.2:1002")); + EXPECT_TRUE(ClientMessengerTestPeer::HasDsContext("10.0.0.3:1003")); + EXPECT_TRUE(ClientMessengerTestPeer::IsDsActive("10.0.0.1:1001")); + EXPECT_TRUE(ClientMessengerTestPeer::IsDsActive("10.0.0.2:1002")); + EXPECT_TRUE(ClientMessengerTestPeer::IsDsActive("10.0.0.3:1003")); + EXPECT_EQ(ClientMessengerTestPeer::GetConnection("10.0.0.1:1001"), healthy_conn); + EXPECT_EQ(ClientMessengerTestPeer::ShardOwner(0), "10.0.0.1:1001"); + EXPECT_EQ(ClientMessengerTestPeer::ShardOwner(1), "10.0.0.2:1002"); + EXPECT_EQ(ClientMessengerTestPeer::ShardOwner(2), "10.0.0.3:1003"); + EXPECT_EQ(built_servers.size(), 2); + EXPECT_TRUE(built_servers.contains("10.0.0.2:1002")); + EXPECT_TRUE(built_servers.contains("10.0.0.3:1003")); + EXPECT_FALSE(built_servers.contains("10.0.0.1:1001")); +} + +TEST_F(ClientMessengerUnitTest, ReInitFailsWhenRouteQueryFailsAndClearsCmAddress) { + ClientMessengerTestPeer::SetGetCmAddressHook([]() { return std::string("10.0.0.100:9000"); }); + ClientMessengerTestPeer::SetRouteQueryHook([](const std::string &) { + return std::make_pair(ClntErr::GetRoutingTableFailed, std::shared_ptr{}); + }); + + EXPECT_EQ(ClientMessengerTestPeer::ReInit(), CommonErr::InvalidState); + EXPECT_FALSE(ClientMessengerTestPeer::IsInitialized()); + EXPECT_EQ(ClientMessengerTestPeer::CmAddr(), ""); +} + +TEST_F(ClientMessengerUnitTest, ReInitRemovesStaleConnectionFromLocalContextMap) { + ClientMessengerTestPeer::InstallDsContext("10.0.0.1:1001", true, 1, {0}, std::make_shared("old")); + ClientMessengerTestPeer::InstallDsContext("10.0.0.9:1999", true, 2, {9}, std::make_shared("stale")); + + ClientMessengerTestPeer::SetGetCmAddressHook([]() { return std::string("10.0.0.100:9000"); }); + ClientMessengerTestPeer::SetRouteQueryHook([](const std::string &) { + return std::make_pair(CommonErr::OK, BuildRoutingResponse({{"10.0.0.1:1001", {0, 1}}})); + }); + ClientMessengerTestPeer::SetBuildConnectionHook([](const std::string &) { return CommonErr::OK; }); + + EXPECT_EQ(ClientMessengerTestPeer::ReInit(), CommonErr::OK); + EXPECT_TRUE(ClientMessengerTestPeer::HasDsContext("10.0.0.1:1001")); + EXPECT_FALSE(ClientMessengerTestPeer::HasDsContext("10.0.0.9:1999")); + EXPECT_EQ(ClientMessengerTestPeer::ShardOwner(0), "10.0.0.1:1001"); + EXPECT_EQ(ClientMessengerTestPeer::ShardOwner(1), "10.0.0.1:1001"); } } // namespace clnt diff --git a/tests/cluster_manager/test_cm_hb_monitor.cc b/tests/cluster_manager/test_cm_hb_monitor.cc index c0f26e7..42bf6f7 100644 --- a/tests/cluster_manager/test_cm_hb_monitor.cc +++ b/tests/cluster_manager/test_cm_hb_monitor.cc @@ -10,11 +10,11 @@ #include #include -#include "cluster_manager/cm_shard_manager.h" #include "cluster_manager/cm_hb_monitor.h" -#include "cluster_manager/cm_service.h" -#include "cluster_manager/cm_rpc_handler.h" #include "cluster_manager/cm_node_manager.h" +#include "cluster_manager/cm_rpc_handler.h" +#include "cluster_manager/cm_service.h" +#include "cluster_manager/cm_shard_manager.h" #include "common/base/common_types.h" #include "common/errcode/errcode_def.h" #include "common/logging/logging.h" @@ -39,22 +39,20 @@ using CmNMPtr = std::shared_ptr; class ClusterManagerHBMonitorTest : public ::testing::Test { protected: - enum class caseType : int { - OnRecvNodeHeartbeat = 1, - bgHBScanLoop = 2 - }; + enum class caseType : int { OnRecvNodeHeartbeat = 1, bgHBScanLoop = 2 }; void SetUp() override { - #ifdef NDEBUG - simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "INFO"}; - #else - simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "DEBUG"}; - #endif - simm::logging::LoggerManager::Instance().UpdateConfig("cluster_manager", cm_log_config); +#ifdef NDEBUG + simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "INFO"}; +#else + simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "DEBUG"}; +#endif + simm::logging::LoggerManager::Instance().UpdateConfig("cluster_manager", cm_log_config); cm_shard_manager_ptr_ = std::make_shared(); cm_node_manager_ptr_ = std::make_shared(); - cm_hb_monitor_ptr_ = std::make_shared(cm_node_manager_ptr_, cm_shard_manager_ptr_); + cm_hb_monitor_ptr_ = + std::make_shared(cm_node_manager_ptr_, cm_shard_manager_ptr_); // cm_hb_monitor_ptr_->Init(); // error_code_t ret = cm_hb_monitor_ptr_->Start(); // EXPECT_EQ(ret, CommonErr::OK); @@ -76,7 +74,7 @@ class ClusterManagerHBMonitorTest : public ::testing::Test { TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { std::atomic serverStarted{false}; std::atomic stopServer{false}; - const uint32_t kDataserverNum = 30; + const uint32_t kDataserverNum = 30; const uint32_t kHBNumPerServer = 8; std::string ip_addr_prefix{"192.168.0."}; int32_t server_port = 40000; @@ -85,7 +83,7 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { // (1) Test ClusterManagerHBMonitor::OnRecvNodeHeartbeat() interface // ------------------------------------------------------------------------ MLOG_INFO("Case to test Test ClusterManagerHBMonitor::OnRecvNodeHeartbeat()"); - folly::CPUThreadPoolExecutor executor1(kDataserverNum+2); + folly::CPUThreadPoolExecutor executor1(kDataserverNum + 2); // light-weight cv // post() - signaled state // reset() - set to unsignaled state again @@ -96,14 +94,15 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { uint32_t old_flag_val_1 = FLAGS_cm_cluster_init_grace_period_inSecs; uint32_t old_flag_val_2 = FLAGS_cm_heartbeat_records_perserver; - FLAGS_cm_cluster_init_grace_period_inSecs = 5; // shorten grace period in cm service start + FLAGS_cm_cluster_init_grace_period_inSecs = 5; // shorten grace period in cm service start FLAGS_cm_heartbeat_records_perserver = 3; auto server_thread = [&]() { // start RPC service and wait for requests simm::common::ModuleServiceState::GetInstance().Reset(FLAGS_cm_cluster_init_grace_period_inSecs); MLOG_INFO("Mark cm into grace period, duration is {} seconds", FLAGS_cm_cluster_init_grace_period_inSecs); - simm::cm::ClusterManagerService cm_service(cm_shard_manager_ptr_, this->cm_node_manager_ptr_, this->cm_hb_monitor_ptr_); + simm::cm::ClusterManagerService cm_service( + cm_shard_manager_ptr_, this->cm_node_manager_ptr_, this->cm_hb_monitor_ptr_); error_code_t ret = cm_service.Init(); EXPECT_EQ(ret, CommonErr::OK); ret = cm_service.Start(); @@ -128,7 +127,7 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { MLOG_INFO("Server thread exit!"); }; - auto client_thread = [&](folly::CPUThreadPoolExecutor & executor) { + auto client_thread = [&](folly::CPUThreadPoolExecutor &executor) { // wait for server enters into grace period state std::this_thread::sleep_for(std::chrono::milliseconds(500)); @@ -140,12 +139,12 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { std::string server_addr_str = ip_addr_prefix + std::to_string(thread_num); // create sirpc object as rpc client - sicl::rpc::SiRPC* sirpc_client; + sicl::rpc::SiRPC *sirpc_client; sicl::rpc::SiRPC::newInstance(sirpc_client, false); - sicl::rpc::RpcContext* ctx_p = nullptr; + sicl::rpc::RpcContext *ctx_p = nullptr; sicl::rpc::RpcContext::newInstance(ctx_p); std::shared_ptr ctx = std::shared_ptr(ctx_p); - //FIXME(ytji): URGENT! - 100 servers under 1s timeout will make rpc timeout issue occasionally + // FIXME(ytji): URGENT! - 100 servers under 1s timeout will make rpc timeout issue occasionally ctx->set_timeout(sicl::transport::TimerTick::TIMER_3S); // send handshake rpc to cm to join into the cluster @@ -154,10 +153,9 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { auto _field = hs_req.mutable_node(); _field->set_ip(server_addr_str); _field->set_port(server_port); - auto hs_done_cb_ok = [&](const google::protobuf::Message* rsp, - const std::shared_ptr ctx) { + auto hs_done_cb_ok = [&](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); - auto response = dynamic_cast(rsp); + auto response = dynamic_cast(rsp); EXPECT_EQ(response->ret_code(), CommonErr::OK); // clean up response object delete hs_resp; @@ -182,10 +180,10 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { auto node_field = hb_req.mutable_node(); node_field->set_ip(server_addr_str); node_field->set_port(server_port); - auto hb_done_cb_ok = [&](const google::protobuf::Message* rsp, - const std::shared_ptr ctx) { + auto hb_done_cb_ok = [&](const google::protobuf::Message *rsp, + const std::shared_ptr ctx) { EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); - auto response = dynamic_cast(rsp); + auto response = dynamic_cast(rsp); EXPECT_EQ(response->ret_code(), CommonErr::OK); // clean up response object delete hb_resp; @@ -210,7 +208,7 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { } } else if (ctype == caseType::bgHBScanLoop) { MLOG_INFO("Enter into caseType::bgHBScanLoop"); - if ((thread_num == 0 || thread_num == 2 || thread_num == 4 || thread_num == 6 || thread_num == 8) && + if ((thread_num == 0 || thread_num == 2 || thread_num == 4 || thread_num == 6 || thread_num == 8) && (n == 4 || n == 5 || n == 6 || n == 7)) { // mock some hb requests missing } else { @@ -231,9 +229,7 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { }; for (uint32_t i = 0; i < kDataserverNum; ++i) { - executor.add([i, &client_sub_thread]() { - client_sub_thread(i); - }); + executor.add([i, &client_sub_thread]() { client_sub_thread(i); }); MLOG_INFO("client_sub_thread-{} added!", i); } @@ -263,20 +259,20 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { { auto locked_map = cm_hb_monitor_ptr_->ds_hb_records_.rlock(); for (uint32_t i = 0; i < kDataserverNum; ++i) { - // FIXME(ytji) : change the port hard code with flag value later - std::string server_addr = ip_addr_prefix + std::to_string(i) + ":" + std::to_string(server_port); - EXPECT_EQ(locked_map->at(server_addr).size(), FLAGS_cm_heartbeat_records_perserver); - auto it = locked_map->at(server_addr).begin(); - auto it_nxt = it + 1; - if (it_nxt != locked_map->at(server_addr).end()) { - if (i == 3 || i == 5) { - // FIXME(ytji): how to check ts? - //EXPECT_EQ(it_nxt->monotonic_tp_ - it->monotonic_tp_, std::chrono::seconds(2)); - } else { - // FIXME(ytji): how to check ts? - //EXPECT_EQ(it_nxt->monotonic_tp_ - it->monotonic_tp_, std::chrono::seconds(1)); - } + // FIXME(ytji) : change the port hard code with flag value later + std::string server_addr = ip_addr_prefix + std::to_string(i) + ":" + std::to_string(server_port); + EXPECT_EQ(locked_map->at(server_addr).size(), FLAGS_cm_heartbeat_records_perserver); + auto it = locked_map->at(server_addr).begin(); + auto it_nxt = it + 1; + if (it_nxt != locked_map->at(server_addr).end()) { + if (i == 3 || i == 5) { + // FIXME(ytji): how to check ts? + // EXPECT_EQ(it_nxt->monotonic_tp_ - it->monotonic_tp_, std::chrono::seconds(2)); + } else { + // FIXME(ytji): how to check ts? + // EXPECT_EQ(it_nxt->monotonic_tp_ - it->monotonic_tp_, std::chrono::seconds(1)); } + } } } @@ -295,7 +291,7 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { serverStarted.store(false); stopServer.store(false); - folly::CPUThreadPoolExecutor executor2(kDataserverNum+2); + folly::CPUThreadPoolExecutor executor2(kDataserverNum + 2); serverDone.reset(); // mock 5 servers 4 HB requests respectively // server-1 : req1, req2, req3, req4; req5 ~ req8 missed @@ -308,9 +304,9 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { uint32_t old_flag_val_3 = FLAGS_cm_heartbeat_timeout_inSecs; uint32_t old_flag_val_4 = FLAGS_cm_heartbeat_bg_scan_interval_inSecs; - FLAGS_cm_heartbeat_records_perserver = 5; // max 5 HB record per dataserver - FLAGS_cm_heartbeat_timeout_inSecs = 2; // 2 secs timeout - FLAGS_cm_heartbeat_bg_scan_interval_inSecs = 1; // 1s bg scan interval + FLAGS_cm_heartbeat_records_perserver = 5; // max 5 HB record per dataserver + FLAGS_cm_heartbeat_timeout_inSecs = 2; // 2 secs timeout + FLAGS_cm_heartbeat_bg_scan_interval_inSecs = 1; // 1s bg scan interval // mark hb monitor statis is runing cm_hb_monitor_ptr_->stop_flag_.store(false); @@ -332,25 +328,25 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { // size() is approximate and may not reflect concurrent updates or recent clear() operations. EXPECT_EQ(locked_map->size(), kDataserverNum); for (uint32_t i = 0; i < kDataserverNum; ++i) { - // FIXME(ytji) : change the port hard code with flag value later - std::string server_addr = ip_addr_prefix + std::to_string(i) + ":" + std::to_string(server_port); - auto it = locked_map->find(server_addr); - EXPECT_TRUE(it != locked_map->end()); + // FIXME(ytji) : change the port hard code with flag value later + std::string server_addr = ip_addr_prefix + std::to_string(i) + ":" + std::to_string(server_port); + auto it = locked_map->find(server_addr); + EXPECT_TRUE(it != locked_map->end()); + if (i == 0 || i == 2 || i == 4 || i == 6 || i == 8) { + EXPECT_EQ(it->second.size(), 4); + } else { + EXPECT_EQ(it->second.size(), FLAGS_cm_heartbeat_records_perserver); + } + auto it_nm = cm_node_manager_ptr_->node_status_map_.find(server_addr); + // EXPECT_TRUE(it_nm != cm_node_manager_ptr_->node_status_map_.end()); + if (it_nm != cm_node_manager_ptr_->node_status_map_.end()) { if (i == 0 || i == 2 || i == 4 || i == 6 || i == 8) { - EXPECT_EQ(it->second.size(), 4); + EXPECT_EQ(it_nm->second, NodeStatus::DEAD); } else { - EXPECT_EQ(it->second.size(), FLAGS_cm_heartbeat_records_perserver); - } - auto it_nm = cm_node_manager_ptr_->node_status_map_.find(server_addr); - //EXPECT_TRUE(it_nm != cm_node_manager_ptr_->node_status_map_.end()); - if (it_nm != cm_node_manager_ptr_->node_status_map_.end()) { - if (i == 0 || i == 2 || i == 4 || i == 6 || i == 8) { - EXPECT_EQ(it_nm->second, NodeStatus::DEAD); - } else { - // FIXME(ytji): all dataservers are marked as DEAD - //EXPECT_EQ(it_nm->second, NodeStatus::RUNNING); - } + // FIXME(ytji): all dataservers are marked as DEAD + // EXPECT_EQ(it_nm->second, NodeStatus::RUNNING); } + } } } @@ -361,12 +357,35 @@ TEST_F(ClusterManagerHBMonitorTest, TestHBMonitor) { FLAGS_cm_heartbeat_bg_scan_interval_inSecs = old_flag_val_4; } +TEST_F(ClusterManagerHBMonitorTest, TestRestartAfterStopResumesHeartbeatScanning) { + auto old_timeout = FLAGS_cm_heartbeat_timeout_inSecs; + auto old_scan_interval = FLAGS_cm_heartbeat_bg_scan_interval_inSecs; + auto restore_flags = folly::makeGuard([&]() { + FLAGS_cm_heartbeat_timeout_inSecs = old_timeout; + FLAGS_cm_heartbeat_bg_scan_interval_inSecs = old_scan_interval; + }); + FLAGS_cm_heartbeat_timeout_inSecs = 1; + FLAGS_cm_heartbeat_bg_scan_interval_inSecs = 1; + + const std::string node_addr = "127.0.0.1:47001"; + ASSERT_EQ(cm_node_manager_ptr_->AddNode(node_addr), CommonErr::OK); + ASSERT_EQ(cm_hb_monitor_ptr_->Start(), CommonErr::OK); + ASSERT_EQ(cm_hb_monitor_ptr_->OnRecvNodeHeartbeat(node_addr), CommonErr::OK); + ASSERT_EQ(cm_hb_monitor_ptr_->Stop(), CommonErr::OK); + EXPECT_EQ(cm_node_manager_ptr_->QueryNodeStatus(node_addr), NodeStatus::RUNNING); + + ASSERT_EQ(cm_hb_monitor_ptr_->Start(), CommonErr::OK); + std::this_thread::sleep_for(std::chrono::milliseconds(2200)); + EXPECT_EQ(cm_node_manager_ptr_->QueryNodeStatus(node_addr), NodeStatus::DEAD); + EXPECT_EQ(cm_hb_monitor_ptr_->Stop(), CommonErr::OK); +} + } // namespace cm } // namespace simm // should use main function in below way to accept gflag options changes // from UT test binary command line -int main(int argc, char** argv) { +int main(int argc, char **argv) { ::testing::InitGoogleTest(&argc, argv); gflags::ParseCommandLineFlags(&argc, &argv, true); return RUN_ALL_TESTS(); diff --git a/tests/cluster_manager/test_cm_rebalance.cc b/tests/cluster_manager/test_cm_rebalance.cc index 376fb1f..ed2b497 100644 --- a/tests/cluster_manager/test_cm_rebalance.cc +++ b/tests/cluster_manager/test_cm_rebalance.cc @@ -37,6 +37,24 @@ DECLARE_string(cm_log_file); namespace simm { namespace cm { +namespace { + +template +bool WaitUntil(Predicate pred, + std::chrono::milliseconds timeout, + std::chrono::milliseconds poll_interval = std::chrono::milliseconds(50)) { + const auto deadline = std::chrono::steady_clock::now() + timeout; + while (std::chrono::steady_clock::now() < deadline) { + if (pred()) { + return true; + } + std::this_thread::sleep_for(poll_interval); + } + return pred(); +} + +} // namespace + class ClusterManagerRebalanceTest : public ::testing::Test { protected: void SetUp() override { @@ -70,9 +88,15 @@ class ClusterManagerRebalanceTest : public ::testing::Test { // TODO: move MockDS to a common test utils file? class MockDataServer { public: - MockDataServer(const std::string &ip, int port) : ip_(ip), port_(port), is_running_(false), should_stop_(false) { - node_address_ = std::make_shared(ip, port); - } + MockDataServer(const std::string &ip, + int port, + std::chrono::milliseconds heartbeat_interval = std::chrono::seconds(2)) + : ip_(ip), + port_(port), + node_address_(std::make_shared(ip, port)), + heartbeat_interval_(heartbeat_interval), + is_running_(false), + should_stop_(false) {} ~MockDataServer() { Stop(); } @@ -116,7 +140,7 @@ class MockDataServer { private: void SendHandshakeRequest() { try { - sicl::rpc::SiRPC *sirpc_client; + sicl::rpc::SiRPC *sirpc_client = nullptr; sicl::rpc::SiRPC::newInstance(sirpc_client, false); sicl::rpc::RpcContext *ctx_p = nullptr; @@ -131,13 +155,15 @@ class MockDataServer { node_field->set_ip(ip_); node_field->set_port(port_); - auto hs_done_cb = [](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { + auto hs_done_cb = [hs_resp](const google::protobuf::Message *rsp, + const std::shared_ptr ctx) { if (ctx->ErrorCode() == sicl::transport::Result::SICL_SUCCESS) { auto response = dynamic_cast(rsp); MLOG_INFO("Handshake successful, ret_code: {}", response->ret_code()); } else { MLOG_ERROR("Handshake failed, error: {}", ctx->ErrorCode()); } + delete hs_resp; }; sirpc_client->SendRequest( @@ -166,13 +192,13 @@ class MockDataServer { MLOG_ERROR("Exception in HeartbeatLoop: {}", e.what()); } - std::this_thread::sleep_for(std::chrono::seconds(2)); + std::this_thread::sleep_for(heartbeat_interval_); } } void SendHeartbeat() { try { - sicl::rpc::SiRPC *sirpc_client; + sicl::rpc::SiRPC *sirpc_client = nullptr; sicl::rpc::SiRPC::newInstance(sirpc_client, false); sicl::rpc::RpcContext *ctx_p = nullptr; @@ -186,12 +212,14 @@ class MockDataServer { node_field->set_ip(ip_); node_field->set_port(port_); - auto hb_done_cb = [](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { + auto hb_done_cb = [hb_resp](const google::protobuf::Message *rsp, + const std::shared_ptr ctx) { if (ctx->ErrorCode() == sicl::transport::Result::SICL_SUCCESS) { MLOG_DEBUG("Heartbeat successful"); } else { MLOG_ERROR("Heartbeat failed, error: {}", ctx->ErrorCode()); } + delete hb_resp; }; sirpc_client->SendRequest("127.0.0.1", @@ -211,6 +239,7 @@ class MockDataServer { std::string ip_; int port_; std::shared_ptr node_address_; + std::chrono::milliseconds heartbeat_interval_; std::atomic is_running_; std::atomic should_stop_; std::jthread heartbeat_thread_; @@ -218,7 +247,7 @@ class MockDataServer { TEST_F(ClusterManagerRebalanceTest, TestShardRebalanceAfterNodeFailure) { constexpr uint32_t kDataserverNum = 4; - const std::string ip_prefix = "192.168.0."; // TODO: remove this? + const std::string ip_prefix = "192.168.0."; // TODO: remove this? const int base_port = 40000; std::atomic cm_started{false}; @@ -245,7 +274,7 @@ TEST_F(ClusterManagerRebalanceTest, TestShardRebalanceAfterNodeFailure) { MLOG_INFO("Starting CM service..."); simm::common::ModuleServiceState::GetInstance().Reset(FLAGS_cm_cluster_init_grace_period_inSecs); - + // Should pre-init DS to avoid TooFew error std::vector> pre_init_servers; for (uint32_t i = 0; i < kDataserverNum; ++i) { @@ -398,8 +427,8 @@ TEST_F(ClusterManagerRebalanceTest, TestShardRebalanceAfterNodeFailure) { TEST_F(ClusterManagerRebalanceTest, TestMultipleNodeFailures) { constexpr uint32_t kDataserverNum = 6; - constexpr uint32_t kFailureNum = 2; // TODO: what if remain nodes less than min req? - const std::string ip_prefix = "192.168.1."; // TODO: remove this? + constexpr uint32_t kFailureNum = 2; // TODO: what if remain nodes less than min req? + const std::string ip_prefix = "192.168.1."; // TODO: remove this? const int base_port = 41000; std::atomic cm_started{false}; @@ -530,7 +559,7 @@ TEST_F(ClusterManagerRebalanceTest, TestMultipleNodeFailures) { // No shards assigned to failed nodes for (const auto &failed_node : failed_nodes) { - EXPECT_NE(node_addr, failed_node) + EXPECT_NE(node_addr, failed_node) << "Found shard " << entry.first << " still assigned to failed node " << failed_node; } } @@ -578,9 +607,171 @@ TEST_F(ClusterManagerRebalanceTest, TestMultipleNodeFailures) { executor.join(); } +TEST_F(ClusterManagerRebalanceTest, TestFastTwoNodeFailuresOnFiveNodeCluster) { + constexpr uint32_t kDataserverNum = 5; + constexpr uint32_t kFailureNum = 2; + const std::string ip_prefix = "192.168.10."; + const int base_port = 42000; + + auto old_grace_period = FLAGS_cm_cluster_init_grace_period_inSecs; + auto old_hb_timeout = FLAGS_cm_heartbeat_timeout_inSecs; + auto old_hb_scan_interval = FLAGS_cm_heartbeat_bg_scan_interval_inSecs; + FLAGS_cm_cluster_init_grace_period_inSecs = 2; + FLAGS_cm_heartbeat_timeout_inSecs = 1; + FLAGS_cm_heartbeat_bg_scan_interval_inSecs = 1; + auto restore_flags = folly::makeGuard([&]() { + FLAGS_cm_cluster_init_grace_period_inSecs = old_grace_period; + FLAGS_cm_heartbeat_timeout_inSecs = old_hb_timeout; + FLAGS_cm_heartbeat_bg_scan_interval_inSecs = old_hb_scan_interval; + }); + + std::atomic cm_started{false}; + std::atomic test_completed{false}; + folly::Baton<> cm_done; + folly::CPUThreadPoolExecutor executor(10); + + auto shard_manager_ptr = std::make_shared(); + auto node_manager_ptr = std::make_shared(); + auto hb_monitor_ptr = std::make_shared(node_manager_ptr, shard_manager_ptr); + auto cm_service_ptr = + std::make_unique(shard_manager_ptr, node_manager_ptr, hb_monitor_ptr); + + auto cm_thread = [&]() { + MLOG_INFO("Starting CM service for fast two-node failure test..."); + + simm::common::ModuleServiceState::GetInstance().Reset(FLAGS_cm_cluster_init_grace_period_inSecs); + + auto ret = cm_service_ptr->Init(); + EXPECT_EQ(ret, CommonErr::OK); + + ret = cm_service_ptr->Start(); + EXPECT_EQ(ret, CommonErr::OK); + EXPECT_TRUE(cm_service_ptr->IsRunning()); + + cm_started.store(true); + MLOG_INFO("CM service started successfully for fast two-node failure test"); + + while (!test_completed.load()) { + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + } + + MLOG_INFO("Stopping CM service for fast two-node failure test..."); + ret = cm_service_ptr->Stop(); + EXPECT_EQ(ret, CommonErr::OK); + cm_done.post(); + }; + + executor.add(cm_thread); + + std::this_thread::sleep_for(std::chrono::milliseconds(200)); + + std::vector> dataservers; + dataservers.reserve(kDataserverNum); + for (uint32_t i = 0; i < kDataserverNum; ++i) { + std::string ip = ip_prefix + std::to_string(i + 1); + int port = base_port + i; + auto ds = std::make_unique(ip, port, std::chrono::milliseconds(200)); + auto ret = ds->Start(); + EXPECT_EQ(ret, CommonErr::OK); + dataservers.push_back(std::move(ds)); + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + } + + ASSERT_TRUE(WaitUntil([&]() { return cm_started.load(); }, std::chrono::milliseconds(3000))) + << "CM did not finish startup in time"; + + ASSERT_TRUE(WaitUntil([&]() { return node_manager_ptr->GetAllNodeAddress(true).size() == kDataserverNum; }, + std::chrono::milliseconds(1000))) + << "CM did not observe all dataservers as RUNNING in time"; + + auto initial_routing = shard_manager_ptr->QueryAllShardRoutingInfos(); + ASSERT_EQ(initial_routing.size(), FLAGS_shard_total_num); + std::unordered_map initial_shard_count; + for (const auto &entry : initial_routing) { + ASSERT_NE(entry.second, nullptr); + initial_shard_count[entry.second->toString()]++; + } + EXPECT_EQ(initial_shard_count.size(), kDataserverNum); + + std::vector failed_nodes; + for (uint32_t i = 0; i < kFailureNum; ++i) { + failed_nodes.push_back(dataservers[i]->GetAddressString()); + auto ret = dataservers[i]->Stop(); + EXPECT_EQ(ret, CommonErr::OK); + } + + auto routing_rebalanced = [&]() { + if (node_manager_ptr->GetAllNodeAddress(true).size() != kDataserverNum - kFailureNum) { + return false; + } + auto routing = shard_manager_ptr->QueryAllShardRoutingInfos(); + if (routing.size() != FLAGS_shard_total_num) { + return false; + } + std::unordered_set target_nodes; + for (const auto &entry : routing) { + if (!entry.second) { + return false; + } + auto node_addr = entry.second->toString(); + for (const auto &failed_node : failed_nodes) { + if (node_addr == failed_node) { + return false; + } + } + target_nodes.insert(node_addr); + } + return target_nodes.size() == kDataserverNum - kFailureNum; + }; + + ASSERT_TRUE(WaitUntil(routing_rebalanced, std::chrono::milliseconds(2500))) + << "CM did not finish two-node failure rebalance in time"; + + auto alive_nodes = node_manager_ptr->GetAllNodeAddress(true); + EXPECT_EQ(alive_nodes.size(), kDataserverNum - kFailureNum); + + auto rebalanced_routing = shard_manager_ptr->QueryAllShardRoutingInfos(); + ASSERT_EQ(rebalanced_routing.size(), FLAGS_shard_total_num); + + std::unordered_map rebalanced_shard_count; + std::unordered_set nodes_with_shards; + for (const auto &entry : rebalanced_routing) { + ASSERT_NE(entry.second, nullptr); + std::string node_addr = entry.second->toString(); + for (const auto &failed_node : failed_nodes) { + EXPECT_NE(node_addr, failed_node) << "Shard " << entry.first << " still points to failed node " << failed_node; + } + rebalanced_shard_count[node_addr]++; + nodes_with_shards.insert(node_addr); + } + + EXPECT_EQ(nodes_with_shards.size(), kDataserverNum - kFailureNum); + + uint32_t total_shards = 0; + for (const auto &[node_addr, shard_num] : rebalanced_shard_count) { + total_shards += shard_num; + } + EXPECT_EQ(total_shards, FLAGS_shard_total_num); + + auto min_max = std::minmax_element(rebalanced_shard_count.begin(), + rebalanced_shard_count.end(), + [](const auto &a, const auto &b) { return a.second < b.second; }); + ASSERT_NE(min_max.first, rebalanced_shard_count.end()); + ASSERT_NE(min_max.second, rebalanced_shard_count.end()); + EXPECT_LE(min_max.second->second - min_max.first->second, 3U); + + for (size_t i = kFailureNum; i < dataservers.size(); ++i) { + dataservers[i]->Stop(); + } + + test_completed.store(true); + cm_done.wait(); + executor.join(); +} + TEST_F(ClusterManagerRebalanceTest, TestShardRebalanceOnlyOnce) { constexpr uint32_t kDataserverNum = 4; - const std::string ip_prefix = "192.168.0."; // TODO: remove this? + const std::string ip_prefix = "192.168.0."; // TODO: remove this? const int base_port = 40000; std::atomic cm_started{false}; @@ -607,7 +798,7 @@ TEST_F(ClusterManagerRebalanceTest, TestShardRebalanceOnlyOnce) { MLOG_INFO("Starting CM service..."); simm::common::ModuleServiceState::GetInstance().Reset(FLAGS_cm_cluster_init_grace_period_inSecs); - + // Should pre-init DS to avoid TooFew error std::vector> pre_init_servers; for (uint32_t i = 0; i < kDataserverNum; ++i) { @@ -747,7 +938,8 @@ TEST_F(ClusterManagerRebalanceTest, TestShardRebalanceOnlyOnce) { MLOG_INFO("Rebalance test completed successfully"); // Sleep 2xhb time - std::this_thread::sleep_for(std::chrono::seconds((FLAGS_cm_heartbeat_timeout_inSecs + FLAGS_cm_heartbeat_bg_scan_interval_inSecs) * 2)); + std::this_thread::sleep_for( + std::chrono::seconds((FLAGS_cm_heartbeat_timeout_inSecs + FLAGS_cm_heartbeat_bg_scan_interval_inSecs) * 2)); // Clean up for (size_t i = 1; i < dataservers.size(); ++i) { @@ -768,4 +960,4 @@ int main(int argc, char **argv) { ::testing::InitGoogleTest(&argc, argv); gflags::ParseCommandLineFlags(&argc, &argv, true); return RUN_ALL_TESTS(); -} \ No newline at end of file +} diff --git a/tests/cluster_manager/test_cm_service.cc b/tests/cluster_manager/test_cm_service.cc index 8e24391..6b60ac0 100644 --- a/tests/cluster_manager/test_cm_service.cc +++ b/tests/cluster_manager/test_cm_service.cc @@ -18,6 +18,7 @@ #include "common/base/common_types.h" #include "common/errcode/errcode_def.h" #include "common/logging/logging.h" +#include "common/rpc_handlers/common_rpc_handlers.h" #include "proto/cm_clnt_rpcs.pb.h" #include "proto/ds_cm_rpcs.pb.h" @@ -29,6 +30,7 @@ DECLARE_uint32(cm_cluster_init_grace_period_inSecs); DECLARE_string(cm_log_file); DECLARE_uint32(dataserver_min_num); DECLARE_string(cm_primary_node_ip); +DECLARE_int32(cm_rpc_admin_port); namespace simm { namespace cm { @@ -38,17 +40,19 @@ using CmSrvPtr = std::unique_ptr; class ClusterManagerServiceTest : public ::testing::Test { protected: void SetUp() override { - #ifdef NDEBUG - simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "INFO"}; - #else - simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "DEBUG"}; - #endif - simm::logging::LoggerManager::Instance().UpdateConfig("cluster_manager", cm_log_config); +#ifdef NDEBUG + simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "INFO"}; +#else + simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{FLAGS_cm_log_file, "DEBUG"}; +#endif + simm::logging::LoggerManager::Instance().UpdateConfig("cluster_manager", cm_log_config); cm_service_ptr_ = std::make_unique(); } - void TearDown() override { cm_service_ptr_.reset(); } + void TearDown() override { + cm_service_ptr_.reset(); + } CmSrvPtr cm_service_ptr_{nullptr}; }; @@ -65,8 +69,8 @@ TEST_F(ClusterManagerServiceTest, TestStartAndStop) { ret = cm_service_ptr_->Start(); EXPECT_EQ(ret, CmErr::InitDataserverNodesTooFew); EXPECT_TRUE(!cm_service_ptr_->IsRunning()); - //EXPECT_EQ(ret, CommonErr::OK); - //EXPECT_TRUE(cm_service_ptr_->IsRunning()); + // EXPECT_EQ(ret, CommonErr::OK); + // EXPECT_TRUE(cm_service_ptr_->IsRunning()); // FIXME(ytji) : for cm isn't started, so stop action also will fail ret = cm_service_ptr_->Stop(); @@ -76,6 +80,25 @@ TEST_F(ClusterManagerServiceTest, TestStartAndStop) { FLAGS_cm_cluster_init_grace_period_inSecs = old_flag_val; } +TEST_F(ClusterManagerServiceTest, TestStartFailureCleansUpRPCServices) { + auto old_grace = FLAGS_cm_cluster_init_grace_period_inSecs; + auto old_min_ds = FLAGS_dataserver_min_num; + FLAGS_cm_cluster_init_grace_period_inSecs = 1; + FLAGS_dataserver_min_num = 3; + + ASSERT_EQ(cm_service_ptr_->Init(), CommonErr::OK); + auto ret = cm_service_ptr_->Start(); + EXPECT_EQ(ret, CmErr::InitDataserverNodesTooFew); + EXPECT_FALSE(cm_service_ptr_->IsRunning()); + + ret = cm_service_ptr_->Start(); + EXPECT_EQ(ret, CmErr::InitDataserverNodesTooFew); + EXPECT_FALSE(cm_service_ptr_->IsRunning()); + + FLAGS_cm_cluster_init_grace_period_inSecs = old_grace; + FLAGS_dataserver_min_num = old_min_ds; +} + TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { std::atomic serverStarted{false}; std::atomic stopServer{false}; @@ -91,10 +114,10 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { auto server_thread = [&]() { auto ret = cm_service_ptr_->Init(); EXPECT_EQ(ret, CommonErr::OK); - //FIXME(ytji) : for no dataserver sends hand-shake RPC to cm, so it will start failed - //ret = cm_service_ptr_->Start(); - //EXPECT_EQ(ret, CommonErr::OK); - //EXPECT_TRUE(cm_service_ptr_->IsRunning()); + // FIXME(ytji) : for no dataserver sends hand-shake RPC to cm, so it will start failed + // ret = cm_service_ptr_->Start(); + // EXPECT_EQ(ret, CommonErr::OK); + // EXPECT_TRUE(cm_service_ptr_->IsRunning()); ret = cm_service_ptr_->StartRPCServices(); EXPECT_EQ(ret, CommonErr::OK); @@ -125,10 +148,10 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { std::this_thread::sleep_for(std::chrono::milliseconds(500)); } - //FIXME(ytji) : for no dataserver sends hand-shake RPC to cm, it doesn't start, so - // stop action will always fail - //ret = cm_service_ptr_->Stop(); - //EXPECT_EQ(ret, CommonErr::OK); + // FIXME(ytji) : for no dataserver sends hand-shake RPC to cm, it doesn't start, so + // stop action will always fail + // ret = cm_service_ptr_->Stop(); + // EXPECT_EQ(ret, CommonErr::OK); ret = cm_service_ptr_->StopRPCServices(); EXPECT_EQ(ret, CommonErr::OK); @@ -142,10 +165,10 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { } // create sirpc object as rpc client - sicl::rpc::SiRPC* sirpc_client; + sicl::rpc::SiRPC *sirpc_client; sicl::rpc::SiRPC::newInstance(sirpc_client, false); - sicl::rpc::RpcContext* ctx_p = nullptr; + sicl::rpc::RpcContext *ctx_p = nullptr; sicl::rpc::RpcContext::newInstance(ctx_p); std::shared_ptr ctx = std::shared_ptr(ctx_p); ctx->set_timeout(sicl::transport::TimerTick::TIMER_1S); @@ -154,10 +177,10 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { QueryShardRoutingTableSingleRequestPB single_query_req; auto single_query_res_found = new QueryShardRoutingTableSingleResponsePB(); single_query_req.set_shard_id(0); // query shard 0 - auto done_cb_ok = [&single_query_res_found, &done_latch](const google::protobuf::Message* rsp, + auto done_cb_ok = [&single_query_res_found, &done_latch](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { ASSERT_TRUE(ctx->ErrorCode() == sicl::transport::Result::SICL_SUCCESS); - auto response = dynamic_cast(rsp); + auto response = dynamic_cast(rsp); EXPECT_EQ(response->ret_code(), CommonErr::OK); EXPECT_EQ(response->shard_info_size(), 1); EXPECT_EQ(response->shard_info(0).shard_ids_size(), 1); @@ -182,10 +205,10 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { single_query_req.set_shard_id(FLAGS_shard_total_num); auto single_query_res_nfound = new QueryShardRoutingTableSingleResponsePB(); - auto done_cb_not_found = [&single_query_res_nfound, &done_latch](const google::protobuf::Message* rsp, + auto done_cb_not_found = [&single_query_res_nfound, &done_latch](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { ASSERT_TRUE(ctx->ErrorCode() == sicl::transport::Result::SICL_SUCCESS); - auto response = dynamic_cast(rsp); + auto response = dynamic_cast(rsp); EXPECT_EQ(response->ret_code(), CommonErr::CmTargetShardIdNotFound); EXPECT_TRUE(response->shard_info().empty()); // clean up response object @@ -216,10 +239,10 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { QueryShardRoutingTableBatchRequestPB batch_query_req; auto batch_query_res = new QueryShardRoutingTableBatchResponsePB(); batch_query_req.mutable_shard_ids()->Add(target_shards_vec.begin(), target_shards_vec.end()); - auto done_cb_batch = [&batch_query_res, &done_latch](const google::protobuf::Message* rsp, + auto done_cb_batch = [&batch_query_res, &done_latch](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { ASSERT_TRUE(ctx->ErrorCode() == sicl::transport::Result::SICL_SUCCESS); - auto response = dynamic_cast(rsp); + auto response = dynamic_cast(rsp); EXPECT_EQ(response->ret_code(), CommonErr::OK); EXPECT_EQ(response->shard_info_size(), 8); // all 8 dataservers foud for (uint32_t i = 0; i < kDataserverNum; ++i) { @@ -277,10 +300,10 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { QueryShardRoutingTableAllRequestPB all_query_req; auto all_query_res = new QueryShardRoutingTableAllResponsePB(); - auto done_cb_all = [&all_query_res, &done_latch](const google::protobuf::Message* rsp, + auto done_cb_all = [&all_query_res, &done_latch](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { EXPECT_TRUE(ctx->ErrorCode() == sicl::transport::Result::SICL_SUCCESS); - auto response = dynamic_cast(rsp); + auto response = dynamic_cast(rsp); EXPECT_EQ(response->ret_code(), CommonErr::OK); // FIXME(ytji) : fix below checks // EXPECT_EQ(response->shard_info_size(), 8); // all 8 dataservers foud @@ -323,6 +346,210 @@ TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableInfoRPCs) { executor.join(); } +TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableRejectsRequestsDuringGracePeriod) { + auto old_grace = FLAGS_cm_cluster_init_grace_period_inSecs; + auto restore = folly::makeGuard([&]() { + FLAGS_cm_cluster_init_grace_period_inSecs = old_grace; + simm::common::ModuleServiceState::GetInstance().MarkServiceReady(); + }); + FLAGS_cm_cluster_init_grace_period_inSecs = 5; + + ASSERT_EQ(cm_service_ptr_->Init(), CommonErr::OK); + simm::common::ModuleServiceState::GetInstance().Reset(FLAGS_cm_cluster_init_grace_period_inSecs); + ASSERT_EQ(cm_service_ptr_->StartRPCServices(), CommonErr::OK); + auto stop_rpc = folly::makeGuard([&]() { EXPECT_EQ(cm_service_ptr_->StopRPCServices(), CommonErr::OK); }); + + sicl::rpc::SiRPC *sirpc_client = nullptr; + ASSERT_EQ(sicl::rpc::SiRPC::newInstance(sirpc_client, false), sicl::transport::Result::SICL_SUCCESS); + auto client = std::unique_ptr(sirpc_client); + + std::latch done_latch(4); + + auto make_ctx = []() { + sicl::rpc::RpcContext *ctx_p = nullptr; + sicl::rpc::RpcContext::newInstance(ctx_p); + auto ctx = std::shared_ptr(ctx_p); + ctx->set_timeout(sicl::transport::TimerTick::TIMER_1S); + return ctx; + }; + + QueryShardRoutingTableSingleRequestPB single_req; + single_req.set_shard_id(0); + auto *single_rsp = new QueryShardRoutingTableSingleResponsePB(); + client->SendRequest("127.0.0.1", + FLAGS_cm_rpc_inter_port, + static_cast(simm::cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_SINGLE), + single_req, + single_rsp, + make_ctx(), + [&done_latch, single_rsp](const google::protobuf::Message *rsp, + const std::shared_ptr ctx) { + EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); + auto response = dynamic_cast(rsp); + ASSERT_NE(response, nullptr); + EXPECT_EQ(response->ret_code(), CmErr::InitInGracePeriod); + delete single_rsp; + done_latch.count_down(); + }); + + QueryShardRoutingTableBatchRequestPB batch_req; + batch_req.add_shard_ids(0); + batch_req.add_shard_ids(1); + auto *batch_rsp = new QueryShardRoutingTableBatchResponsePB(); + client->SendRequest( + "127.0.0.1", + FLAGS_cm_rpc_inter_port, + static_cast(simm::cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_BATCH), + batch_req, + batch_rsp, + make_ctx(), + [&done_latch, batch_rsp](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { + EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); + auto response = dynamic_cast(rsp); + ASSERT_NE(response, nullptr); + EXPECT_EQ(response->ret_code(), CmErr::InitInGracePeriod); + delete batch_rsp; + done_latch.count_down(); + }); + + QueryShardRoutingTableAllRequestPB all_req; + auto *all_rsp = new QueryShardRoutingTableAllResponsePB(); + client->SendRequest( + "127.0.0.1", + FLAGS_cm_rpc_inter_port, + static_cast(simm::cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_ALL), + all_req, + all_rsp, + make_ctx(), + [&done_latch, all_rsp](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { + EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); + auto response = dynamic_cast(rsp); + ASSERT_NE(response, nullptr); + EXPECT_EQ(response->ret_code(), CmErr::InitInGracePeriod); + delete all_rsp; + done_latch.count_down(); + }); + + ListNodesRequestPB list_req; + auto *list_rsp = new ListNodesResponsePB(); + client->SendRequest( + "127.0.0.1", + FLAGS_cm_rpc_admin_port, + static_cast(simm::common::CommonRpcType::RPC_LIST_NODE_REQ), + list_req, + list_rsp, + make_ctx(), + [&done_latch, list_rsp](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { + EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); + auto response = dynamic_cast(rsp); + ASSERT_NE(response, nullptr); + EXPECT_EQ(response->ret_code(), CmErr::InitInGracePeriod); + delete list_rsp; + done_latch.count_down(); + }); + + done_latch.wait(); +} + +TEST_F(ClusterManagerServiceTest, TestQueryRoutingTableReturnsIncompleteWhenShardUnavailable) { + auto old_min_ds = FLAGS_dataserver_min_num; + auto restore = folly::makeGuard([&]() { + FLAGS_dataserver_min_num = old_min_ds; + simm::common::ModuleServiceState::GetInstance().MarkServiceReady(); + }); + FLAGS_dataserver_min_num = 3; + + ASSERT_EQ(cm_service_ptr_->Init(), CommonErr::OK); + simm::common::ModuleServiceState::GetInstance().MarkServiceReady(); + ASSERT_EQ(cm_service_ptr_->StartRPCServices(), CommonErr::OK); + auto stop_rpc = folly::makeGuard([&]() { EXPECT_EQ(cm_service_ptr_->StopRPCServices(), CommonErr::OK); }); + + std::vector> all_servers; + for (uint32_t i = 0; i < 3; ++i) { + all_servers.push_back(std::make_shared("192.168.2." + std::to_string(i + 1), 43000 + i)); + } + ASSERT_EQ(cm_service_ptr_->shard_manager_->InitShardRoutingTable(all_servers), CommonErr::OK); + ASSERT_EQ(cm_service_ptr_->shard_manager_->RebalanceShardsAfterNodeFailure({all_servers[0]->toString()}, + {all_servers[1], all_servers[2]}), + CmErr::InsufficientDataservers); + + sicl::rpc::SiRPC *sirpc_client = nullptr; + ASSERT_EQ(sicl::rpc::SiRPC::newInstance(sirpc_client, false), sicl::transport::Result::SICL_SUCCESS); + auto client = std::unique_ptr(sirpc_client); + + std::latch done_latch(3); + + auto make_ctx = []() { + sicl::rpc::RpcContext *ctx_p = nullptr; + sicl::rpc::RpcContext::newInstance(ctx_p); + auto ctx = std::shared_ptr(ctx_p); + ctx->set_timeout(sicl::transport::TimerTick::TIMER_1S); + return ctx; + }; + + QueryShardRoutingTableSingleRequestPB single_req; + single_req.set_shard_id(0); + auto *single_rsp = new QueryShardRoutingTableSingleResponsePB(); + client->SendRequest("127.0.0.1", + FLAGS_cm_rpc_inter_port, + static_cast(simm::cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_SINGLE), + single_req, + single_rsp, + make_ctx(), + [&done_latch, single_rsp](const google::protobuf::Message *rsp, + const std::shared_ptr ctx) { + EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); + auto response = dynamic_cast(rsp); + ASSERT_NE(response, nullptr); + EXPECT_EQ(response->ret_code(), CommonErr::CmRoutingInfoNotComplete); + EXPECT_EQ(response->shard_info_size(), 0); + delete single_rsp; + done_latch.count_down(); + }); + + QueryShardRoutingTableBatchRequestPB batch_req; + batch_req.add_shard_ids(0); + batch_req.add_shard_ids(1); + auto *batch_rsp = new QueryShardRoutingTableBatchResponsePB(); + client->SendRequest( + "127.0.0.1", + FLAGS_cm_rpc_inter_port, + static_cast(simm::cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_BATCH), + batch_req, + batch_rsp, + make_ctx(), + [&done_latch, batch_rsp](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { + EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); + auto response = dynamic_cast(rsp); + ASSERT_NE(response, nullptr); + EXPECT_EQ(response->ret_code(), CommonErr::CmRoutingInfoNotComplete); + EXPECT_GE(response->shard_info_size(), 1); + delete batch_rsp; + done_latch.count_down(); + }); + + QueryShardRoutingTableAllRequestPB all_req; + auto *all_rsp = new QueryShardRoutingTableAllResponsePB(); + client->SendRequest( + "127.0.0.1", + FLAGS_cm_rpc_inter_port, + static_cast(simm::cm::ClusterManagerRpcType::RPC_ROUTING_TABLE_QUERY_ALL), + all_req, + all_rsp, + make_ctx(), + [&done_latch, all_rsp](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { + EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); + auto response = dynamic_cast(rsp); + ASSERT_NE(response, nullptr); + EXPECT_EQ(response->ret_code(), CommonErr::CmRoutingInfoNotComplete); + EXPECT_GT(response->shard_info_size(), 0); + delete all_rsp; + done_latch.count_down(); + }); + + done_latch.wait(); +} + TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { std::atomic serverStarted{false}; std::atomic stopServer{false}; @@ -340,8 +567,8 @@ TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { auto shard_manager_ptr = std::make_shared(); auto node_manager_ptr = std::make_shared(); auto hb_monitor_ptr = std::make_shared(node_manager_ptr, shard_manager_ptr); - auto cm_service_ptr = std::make_unique( - shard_manager_ptr, node_manager_ptr, hb_monitor_ptr); + auto cm_service_ptr = + std::make_unique(shard_manager_ptr, node_manager_ptr, hb_monitor_ptr); auto ret = cm_service_ptr->Init(); EXPECT_EQ(ret, CommonErr::OK); ret = cm_service_ptr->Start(); @@ -357,11 +584,11 @@ TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { // test checks shard_id_t steps = static_cast(FLAGS_shard_total_num / kDataserverNum); - shard_id_t rests = static_cast(FLAGS_shard_total_num - steps * (kDataserverNum - 1)); + shard_id_t remainder = static_cast(FLAGS_shard_total_num % kDataserverNum); auto all_servers_vec = node_manager_ptr->GetAllNodeAddress(); EXPECT_EQ(all_servers_vec.size(), kDataserverNum); std::unordered_set ds_set; - for (const auto & ptr : all_servers_vec) { + for (const auto &ptr : all_servers_vec) { ds_set.insert(ptr->toString()); } EXPECT_EQ(ds_set.size(), kDataserverNum); @@ -376,16 +603,24 @@ TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { EXPECT_EQ(st, NodeStatus::RUNNING); } auto query_res = shard_manager_ptr->QueryAllShardRoutingInfos(); - for (const auto & entry : query_res) { + for (const auto &entry : query_res) { std::string ds_addr_str = entry.second->toString(); auto it = ds_set.find(ds_addr_str); EXPECT_TRUE(it != ds_set.end()); converted_routing_map[ds_addr_str]++; } - for (const auto & entry : converted_routing_map) { - EXPECT_TRUE(entry.second == steps || entry.second == rests); + EXPECT_EQ(converted_routing_map.size(), kDataserverNum); + for (const auto &entry : converted_routing_map) { + EXPECT_TRUE(entry.second == steps || entry.second == steps + 1); MLOG_INFO("DS({}), holds ({}) shards", entry.first, entry.second); } + uint32_t nodes_with_extra_shard = 0; + for (const auto &entry : converted_routing_map) { + if (entry.second == steps + 1) { + ++nodes_with_extra_shard; + } + } + EXPECT_EQ(nodes_with_extra_shard, remainder); MLOG_INFO("Server thread start to stop!"); ret = cm_service_ptr->Stop(); @@ -401,8 +636,8 @@ TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { }; auto client_thread = [&]() { - //while (!serverStarted.load()) { - // just hang 1 second to wait for RPC service ready of cm service + // while (!serverStarted.load()) { + // just hang 1 second to wait for RPC service ready of cm service std::this_thread::sleep_for(std::chrono::milliseconds(1000)); //} @@ -413,12 +648,12 @@ TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { std::string server_addr_str = ip_addr_prefix + std::to_string(thread_num); // create sirpc object as rpc client - sicl::rpc::SiRPC* sirpc_client; + sicl::rpc::SiRPC *sirpc_client; sicl::rpc::SiRPC::newInstance(sirpc_client, false); - sicl::rpc::RpcContext* ctx_p = nullptr; + sicl::rpc::RpcContext *ctx_p = nullptr; sicl::rpc::RpcContext::newInstance(ctx_p); std::shared_ptr ctx = std::shared_ptr(ctx_p); - //FIXME(ytji): URGENT! - 100 servers under 1s timeout will make rpc timeout issue occasionally + // FIXME(ytji): URGENT! - 100 servers under 1s timeout will make rpc timeout issue occasionally ctx->set_timeout(sicl::transport::TimerTick::TIMER_3S); // send handshake rpc to cm to join into the cluster @@ -427,10 +662,9 @@ TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { auto node_field = hs_req.mutable_node(); node_field->set_ip(server_addr_str); node_field->set_port(server_port); - auto hs_done_cb_ok = [&](const google::protobuf::Message* rsp, - const std::shared_ptr ctx) { + auto hs_done_cb_ok = [&](const google::protobuf::Message *rsp, const std::shared_ptr ctx) { EXPECT_EQ(ctx->ErrorCode(), sicl::transport::Result::SICL_SUCCESS); - auto response = dynamic_cast(rsp); + auto response = dynamic_cast(rsp); EXPECT_EQ(response->ret_code(), CommonErr::OK); // clean up response object // delete hs_resp; @@ -451,9 +685,7 @@ TEST_F(ClusterManagerServiceTest, TestNewNodeHandShakeRPC) { }; for (uint32_t i = 0; i < kDataserverNum; ++i) { - executor.add([i, &client_sub_thread]() { - client_sub_thread(i); - }); + executor.add([i, &client_sub_thread]() { client_sub_thread(i); }); MLOG_INFO("client_sub_thread-{} added!", i); } @@ -550,7 +782,7 @@ TEST_F(ClusterManagerServiceTest, TestNodeRejoinRPC) { } // namespace cm } // namespace simm -int main(int argc, char** argv) { +int main(int argc, char **argv) { ::testing::InitGoogleTest(&argc, argv); gflags::ParseCommandLineFlags(&argc, &argv, true); return RUN_ALL_TESTS(); diff --git a/tests/cluster_manager/test_cm_shard_manager.cc b/tests/cluster_manager/test_cm_shard_manager.cc index b2633a4..0eb8a6d 100644 --- a/tests/cluster_manager/test_cm_shard_manager.cc +++ b/tests/cluster_manager/test_cm_shard_manager.cc @@ -7,6 +7,7 @@ #include #include +#include #include #include #include @@ -21,6 +22,7 @@ DECLARE_LOG_MODULE("cluster_manager"); DECLARE_uint32(shard_total_num); +DECLARE_uint32(dataserver_min_num); namespace simm { namespace cm { @@ -30,17 +32,19 @@ using CmSmPtr = std::unique_ptr; class ClusterManagerShardManagerTest : public ::testing::Test { protected: void SetUp() override { - #ifdef NDEBUG - simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{"/tmp/simm_cm.log", "INFO"}; - #else - simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{"/tmp/simm_cm.log", "DEBUG"}; - #endif - simm::logging::LoggerManager::Instance().UpdateConfig("cluster_manager", cm_log_config); +#ifdef NDEBUG + simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{"/tmp/simm_cm.log", "INFO"}; +#else + simm::logging::LogConfig cm_log_config = simm::logging::LogConfig{"/tmp/simm_cm.log", "DEBUG"}; +#endif + simm::logging::LoggerManager::Instance().UpdateConfig("cluster_manager", cm_log_config); cm_shard_manager_ptr_ = std::make_unique(); } - void TearDown() override { cm_shard_manager_ptr_.reset(); } + void TearDown() override { + cm_shard_manager_ptr_.reset(); + } CmSmPtr cm_shard_manager_ptr_{nullptr}; }; @@ -51,7 +55,7 @@ TEST_F(ClusterManagerShardManagerTest, TestGetTotalShardNum) { TEST_F(ClusterManagerShardManagerTest, TestTriggerShardsMigration) { std::vector target_shards; - std::vector target_servers; + std::vector target_servers; error_code_t ret = cm_shard_manager_ptr_->TriggerShardsMigration(target_shards, target_servers); EXPECT_EQ(ret, CommonErr::NotImplemented); @@ -72,8 +76,9 @@ TEST_F(ClusterManagerShardManagerTest, TestInitShardRoutingTable) { std::array kDataserverNumArray = {2, 8, 10, 50, 200, 500, 1000, 5000}; // set different sub step length to check rounting table entry results std::array kCheckSubStepArray = {13, 71, 100, 255}; - for (auto& kCheckSubStep : kCheckSubStepArray) { - for (auto& kDataserverNum : kDataserverNumArray) { + for (auto &kCheckSubStep : kCheckSubStepArray) { + for (auto &kDataserverNum : kDataserverNumArray) { + cm_shard_manager_ptr_->CleanRoutingTable(); uint32_t totalShardNum = cm_shard_manager_ptr_->GetTotalShardNum(); EXPECT_EQ(totalShardNum, FLAGS_shard_total_num); uint32_t step = totalShardNum / kDataserverNum; @@ -110,7 +115,7 @@ TEST_F(ClusterManagerShardManagerTest, TestInitShardRoutingTable) { } next_loop: - for (auto& server : all_servers) { + for (auto &server : all_servers) { server.reset(); } all_servers.clear(); @@ -121,7 +126,7 @@ TEST_F(ClusterManagerShardManagerTest, TestInitShardRoutingTable) { cm_shard_manager_ptr_->CleanRoutingTable(); } -folly::coro::Task concurrentReadTask(const CmSmPtr& testMgrPtr) { +folly::coro::Task concurrentReadTask(const CmSmPtr &testMgrPtr) { constexpr uint32_t kRandQueryNum = 1000; for (uint32_t i = 0; i < kRandQueryNum; ++i) { uint32_t rand_shard_id = folly::Random::rand32(0, 16384); // [min, max) @@ -136,13 +141,13 @@ folly::coro::Task concurrentReadTask(const CmSmPtr& testMgrPtr) { co_return; } -folly::coro::Task concurrentWriteTask(const CmSmPtr& testMgrPtr) { +folly::coro::Task concurrentWriteTask(const CmSmPtr &testMgrPtr) { constexpr uint32_t kRandModifyNum = 5000; for (uint32_t i = 0; i < kRandModifyNum; ++i) { uint32_t rand_shard_id = folly::Random::rand32(0, 16384); // [min, max) - auto new_server = std::make_shared( - "192.168.1." + std::to_string(folly::Random::rand32(0, 255)), - static_cast(folly::Random::rand32(30000, 40000))); + auto new_server = + std::make_shared("192.168.1." + std::to_string(folly::Random::rand32(0, 255)), + static_cast(folly::Random::rand32(30000, 40000))); auto ret = testMgrPtr->ModifyRoutingTable(rand_shard_id, new_server); EXPECT_EQ(ret, CommonErr::OK); // TODO(ytji) : maybe add some check for routing table after modify @@ -258,7 +263,7 @@ TEST_F(ClusterManagerShardManagerTest, TestModifyShardRoutingTable) { // test single shard modify shard_id_t target_shard = 0; simm::cm::QueryResultMap resmap = cm_shard_manager_ptr_->QueryShardRoutingInfo(target_shard); - EXPECT_EQ(resmap.size(), 0); + EXPECT_EQ(resmap.size(), 1); EXPECT_EQ(resmap.at(target_shard)->node_ip_, "192.168.1.1"); EXPECT_EQ(resmap.at(target_shard)->node_port_, 30000); auto new_server = std::make_shared("192.168.1.3", 30002); @@ -319,6 +324,57 @@ TEST_F(ClusterManagerShardManagerTest, TestModifyShardRoutingTable) { cm_shard_manager_ptr_->CleanRoutingTable(); } +TEST_F(ClusterManagerShardManagerTest, TestFailureBelowMinMarksDeadNodeShardsUnavailable) { + auto old_min_ds = FLAGS_dataserver_min_num; + auto restore_flag = folly::makeGuard([&]() { FLAGS_dataserver_min_num = old_min_ds; }); + FLAGS_dataserver_min_num = 3; + + std::vector> all_servers; + for (uint32_t i = 0; i < 3; ++i) { + all_servers.push_back(std::make_shared("192.168.2." + std::to_string(i + 1), 32000 + i)); + } + + ASSERT_EQ(cm_shard_manager_ptr_->InitShardRoutingTable(all_servers), CommonErr::OK); + const auto failed_node = all_servers[0]->toString(); + + std::vector> alive_servers{all_servers[1], all_servers[2]}; + EXPECT_EQ(cm_shard_manager_ptr_->RebalanceShardsAfterNodeFailure({failed_node}, alive_servers), + CmErr::InsufficientDataservers); + + auto routing_info = cm_shard_manager_ptr_->QueryAllShardRoutingInfos(); + uint32_t unavailable_shards = 0; + for (const auto &[shard_id, node_addr] : routing_info) { + if (!node_addr) { + ++unavailable_shards; + continue; + } + EXPECT_NE(node_addr->toString(), failed_node); + } + EXPECT_GT(unavailable_shards, 0U); +} + +TEST_F(ClusterManagerShardManagerTest, TestFailureWithNoAliveServersMarksAllFailedNodeShardsUnavailable) { + std::vector> all_servers; + for (uint32_t i = 0; i < 3; ++i) { + all_servers.push_back(std::make_shared("192.168.3." + std::to_string(i + 1), 33000 + i)); + } + + ASSERT_EQ(cm_shard_manager_ptr_->InitShardRoutingTable(all_servers), CommonErr::OK); + + std::vector failed_nodes; + for (const auto &server : all_servers) { + failed_nodes.push_back(server->toString()); + } + + EXPECT_EQ(cm_shard_manager_ptr_->RebalanceShardsAfterNodeFailure(failed_nodes, {}), CmErr::NoAvailableDataservers); + + auto routing_info = cm_shard_manager_ptr_->QueryAllShardRoutingInfos(); + ASSERT_EQ(routing_info.size(), FLAGS_shard_total_num); + for (const auto &[shard_id, node_addr] : routing_info) { + EXPECT_EQ(node_addr, nullptr) << "Shard " << shard_id << " should be unavailable when no dataservers are alive"; + } +} + // copy from folly example test case TEST_F(ClusterManagerShardManagerTest, TestVectorOfTaskWithExecutorUsage) { folly::CPUThreadPoolExecutor threadPool{4, std::make_shared("TestThreadPool")}; @@ -341,7 +397,7 @@ TEST_F(ClusterManagerShardManagerTest, TestVectorOfTaskWithExecutorUsage) { } // namespace cm } // namespace simm -int main(int argc, char** argv) { +int main(int argc, char **argv) { ::testing::InitGoogleTest(&argc, argv); gflags::ParseCommandLineFlags(&argc, &argv, true); return RUN_ALL_TESTS(); diff --git a/tests/data_server/test_ds_kv_service.cc b/tests/data_server/test_ds_kv_service.cc index 2ad6d48..2eb0271 100644 --- a/tests/data_server/test_ds_kv_service.cc +++ b/tests/data_server/test_ds_kv_service.cc @@ -1,7 +1,11 @@ +#include +#include +#include +#include + #include #include - #include "folly/Random.h" #include "folly/executors/IOThreadPoolExecutor.h" #include "folly/futures/Future.h" @@ -9,21 +13,42 @@ #include "transport/types.h" -#include "common/logging/logging.h" #include "common/base/hash.h" #include "common/hashkit/murmurhash.h" -#include "proto/ds_clnt_rpcs.pb.h" -#include "rpc/rpc_context.h" +#include "common/logging/logging.h" +#include "data_server/ds_memory_allocator.h" #include "data_server/kv_rpc_handler.h" #include "data_server/kv_rpc_service.h" +#include "proto/ds_clnt_rpcs.pb.h" +#include "rpc/rpc_context.h" DECLARE_LOG_MODULE("data_server"); DECLARE_uint64(memory_limit_bytes); DECLARE_uint32(ds_free_memory_usable_ratio); +DECLARE_uint32(cm_hb_tolerance_count); +DECLARE_bool(ds_process_exit_cm_disconnection); namespace simm { namespace ds { +namespace { + +std::atomic g_sigterm_received{false}; + +void TestSigTermHandler(int signal) { + if (signal == SIGTERM) { + g_sigterm_received.store(true); + } +} + +std::string MakeShmPath(const std::string &shm_name) { + std::ostringstream oss; + oss << "/dev/shm/" << shm_name; + return oss.str(); +} + +} // namespace + size_t get_random_size(size_t min = 1, size_t max = 1UL << 22) { return folly::Random::rand32(min, max + 1); } @@ -57,6 +82,14 @@ class KVServiceTest : public ::testing::Test { std::unique_ptr rpcServicePtr; }; +class KVServiceLightTest : public ::testing::Test { + protected: + void SetUp() override { rpcServicePtr = std::make_unique(); } + void TearDown() override { rpcServicePtr.reset(); } + + std::unique_ptr rpcServicePtr; +}; + TEST_F(KVServiceTest, TestClientHandlers) { size_t server_thread_num = 1; size_t client_thread_num = 3; @@ -214,10 +247,77 @@ TEST_F(KVServiceTest, TestClientHandlers) { serverPool->join(); clientPool->join(); } + +TEST_F(KVServiceLightTest, TestHeartbeatFailureCountResetOnSuccess) { + rpcServicePtr->heartbeat_failure_count_.store(4); + rpcServicePtr->cm_ready_.store(true); + + rpcServicePtr->OnHeartbeatResult(false, CommonErr::OK); + EXPECT_EQ(rpcServicePtr->heartbeat_failure_count_.load(), 0); + EXPECT_TRUE(rpcServicePtr->cm_ready_.load()); + + rpcServicePtr->OnHeartbeatResult(true, CommonErr::InvalidArgument); + EXPECT_EQ(rpcServicePtr->heartbeat_failure_count_.load(), 1); +} + +TEST_F(KVServiceLightTest, TestClusterManagerDisconnectHandlerInvokedOnToleranceReached) { + auto old_flag = FLAGS_ds_process_exit_cm_disconnection; + auto restore_flag = folly::makeGuard([&]() { FLAGS_ds_process_exit_cm_disconnection = old_flag; }); + FLAGS_ds_process_exit_cm_disconnection = true; + + bool disconnect_handler_invoked = false; + rpcServicePtr->cluster_disconnect_handler_ = [&]() { disconnect_handler_invoked = true; }; + rpcServicePtr->heartbeat_failure_count_.store(FLAGS_cm_hb_tolerance_count - 1); + + rpcServicePtr->OnHeartbeatResult(true, CommonErr::InvalidArgument); + + EXPECT_TRUE(disconnect_handler_invoked); + EXPECT_EQ(rpcServicePtr->heartbeat_failure_count_.load(), 0); +} + +TEST_F(KVServiceLightTest, TestHeartbeatFailureToleranceTriggersReconnectWhenExitDisabled) { + auto old_flag = FLAGS_ds_process_exit_cm_disconnection; + auto restore_flag = folly::makeGuard([&]() { FLAGS_ds_process_exit_cm_disconnection = old_flag; }); + FLAGS_ds_process_exit_cm_disconnection = false; + + bool disconnect_handler_invoked = false; + rpcServicePtr->cluster_disconnect_handler_ = [&]() { disconnect_handler_invoked = true; }; + rpcServicePtr->heartbeat_failure_count_.store(FLAGS_cm_hb_tolerance_count - 1); + rpcServicePtr->cm_ready_.store(true); + + rpcServicePtr->OnHeartbeatResult(true, CommonErr::InvalidArgument); + + EXPECT_FALSE(disconnect_handler_invoked); + EXPECT_FALSE(rpcServicePtr->cm_ready_.load()); + EXPECT_EQ(rpcServicePtr->heartbeat_failure_count_.load(), 0); +} + +TEST_F(KVServiceLightTest, TestClusterManagerDisconnectSignalPathRaisesSigterm) { + g_sigterm_received.store(false); + auto prev_handler = std::signal(SIGTERM, TestSigTermHandler); + auto restore_handler = folly::makeGuard([&]() { std::signal(SIGTERM, prev_handler); }); + + rpcServicePtr->cluster_disconnect_handler_(); + EXPECT_TRUE(g_sigterm_received.load()); +} + +TEST_F(KVServiceLightTest, TestShmAllocatorDestructorReleasesSharedMemory) { + auto shm_name = "simm_ds_test_shm_" + std::to_string(getpid()) + "_" + std::to_string(folly::Random::rand32()); + auto shm_path = MakeShmPath(shm_name); + + { + ShmAllocator allocator; + MemBlock block(4096, shm_name.c_str()); + ASSERT_EQ(allocator.allocate(&block), 0); + ASSERT_EQ(access(shm_path.c_str(), F_OK), 0); + } + + EXPECT_NE(access(shm_path.c_str(), F_OK), 0) << "shared memory should be unlinked after allocator destruction"; +} } // namespace ds } // namespace simm -int main(int argc, char** argv) { +int main(int argc, char **argv) { ::testing::InitGoogleTest(&argc, argv); gflags::ParseCommandLineFlags(&argc, &argv, true); return RUN_ALL_TESTS();