diff --git a/src/client/clnt_flags.cc b/src/client/clnt_flags.cc index dcba4da..25ea741 100644 --- a/src/client/clnt_flags.cc +++ b/src/client/clnt_flags.cc @@ -1,10 +1,12 @@ #include 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_bool(clnt_syncreq_enable_retry, true, "enable simm client sync requests retry or not"); DEFINE_uint32(clnt_syncreq_retry_count, 2, "simm client sync requests retry count"); +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_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 80fb7d3..3f8ea1d 100644 --- a/src/client/clnt_messenger.cc +++ b/src/client/clnt_messenger.cc @@ -40,6 +40,7 @@ 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(clnt_syncreq_enable_retry); DECLARE_bool(simm_enable_trace); DECLARE_LOG_MODULE("simm_client"); @@ -193,27 +194,64 @@ std::string ClientMessenger::get_cm_address() { return cm_ips[0].second; } -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); +error_code_t ClientMessenger::build_connection(const std::string &addr, BuildConnWaitMode wait_mode) { + auto ds_ctx = GetOrCreateConnectionContext(addr); + if (ds_ctx->active.load() && ds_ctx->LoadConnection() != nullptr) { + return CommonErr::OK; } -#endif + + bool expected = false; + if (!ds_ctx->connecting_.compare_exchange_strong(expected, true, std::memory_order_acq_rel)) { + if (wait_mode == BuildConnWaitMode::kNoWait) { + return ClntErr::BuildConnectionFailed; + } + std::unique_lock lock(ds_ctx->connect_wait_mutex_); + ds_ctx->connect_cv_.wait(lock, [&]() { return !ds_ctx->connecting_.load(std::memory_order_acquire); }); + if (ds_ctx->active.load() && ds_ctx->LoadConnection() != nullptr) { + return CommonErr::OK; + } + return ClntErr::BuildConnectionFailed; + } + + error_code_t ret = CommonErr::OK; + std::shared_ptr connection; auto node_addr = simm::common::NodeAddress::ParseFromString(addr); if (!node_addr) { MLOG_ERROR("Invalid server address format: {}", addr); - return CommonErr::InvalidArgument; + ret = CommonErr::InvalidArgument; + } else { +#if defined(SIMM_UNIT_TEST) + if (test_build_connection_hook_) { + ret = test_build_connection_hook_(addr); + } else +#endif + { + connection = rpc_client_->connect(node_addr->node_ip_, node_addr->node_port_); + if (connection == nullptr) { + MLOG_ERROR("Build connection with {} failed", addr); + ret = ClntErr::BuildConnectionFailed; + } + } } - std::shared_ptr connection = rpc_client_->connect(node_addr->node_ip_, node_addr->node_port_); - if (connection == nullptr) { - MLOG_ERROR("Build connection with {} failed", addr); - return ClntErr::BuildConnectionFailed; + + if (ret == CommonErr::OK) { +#if !defined(SIMM_UNIT_TEST) + ds_ctx->StoreConnection(connection); + ds_ctx->gen_num.fetch_add(1); + ds_ctx->active.store(true); +#else + if (test_build_connection_hook_ == nullptr) { + ds_ctx->StoreConnection(connection); + ds_ctx->gen_num.fetch_add(1); + ds_ctx->active.store(true); + } else if (!(ds_ctx->active.load() && ds_ctx->LoadConnection() != nullptr)) { + ret = ClntErr::BuildConnectionFailed; + } +#endif } - auto ds_ctx = GetOrCreateConnectionContext(addr); - ds_ctx->StoreConnection(connection); - ds_ctx->gen_num.fetch_add(1); - ds_ctx->active.store(true); - return CommonErr::OK; + ds_ctx->connecting_.store(false, std::memory_order_release); + ds_ctx->connect_cv_.notify_all(); + return ret; } void ClientMessenger::ReleaseConnectionContext(const std::shared_ptr &ds_ctx) { @@ -285,7 +323,6 @@ error_code_t ClientMessenger::call_sync(uint16_t shard_id, 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(); @@ -307,19 +344,25 @@ error_code_t ClientMessenger::call_sync(uint16_t shard_id, MLOG_DEBUG("call_sync rpc succeed"); return CommonErr::OK; } - ReconnectByErrors(rpc_ctx, ds_ctx, shard_id, tag); } + // trigger background failover thread to do global routing table re-fetch and + // QP connections re-build by specified errors + 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 < FLAGS_clnt_syncreq_retry_count) { + if (!FLAGS_clnt_syncreq_enable_retry) { + MLOG_WARN("Sync requests retry mechanism ({} retries) disabled", FLAGS_clnt_syncreq_retry_count); + break; + } else 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)", FLAGS_clnt_syncreq_retry_count); + MLOG_ERROR("Failed to send request after {} retries (sync call)", + FLAGS_clnt_syncreq_enable_retry ? FLAGS_clnt_syncreq_retry_count : 0); return ClntErr::ClntSendRPCFailed; } @@ -350,7 +393,7 @@ error_code_t ClientMessenger::execute(const std::string &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 - auto ret = build_connection(addr); + auto ret = build_connection(addr, BuildConnWaitMode::kNoWait); if (ret != CommonErr::OK) { MLOG_ERROR("Connect with {} failed :{}", addr, ret); return ret; diff --git a/src/client/clnt_messenger.h b/src/client/clnt_messenger.h index a560858..6ad1b79 100644 --- a/src/client/clnt_messenger.h +++ b/src/client/clnt_messenger.h @@ -98,7 +98,12 @@ class ClientMessenger { error_code_t RegisterHandlers(); // build sicl connection - error_code_t build_connection(const std::string &addr); + enum class BuildConnWaitMode { + kWaitForInflight, + kNoWait, + }; + error_code_t build_connection(const std::string &addr, + BuildConnWaitMode wait_mode = BuildConnWaitMode::kWaitForInflight); // RPC options in threadpool template @@ -169,8 +174,13 @@ class ClientMessenger { } private: + friend class ClientMessenger; + mutable std::shared_mutex connection_mutex_; std::shared_ptr connection_{nullptr}; + std::mutex connect_wait_mutex_; + std::condition_variable connect_cv_; + std::atomic connecting_{false}; }; // cluster manager address diff --git a/tests/client/test_clnt_messenger.cc b/tests/client/test_clnt_messenger.cc index 8dc4869..a69d5c8 100644 --- a/tests/client/test_clnt_messenger.cc +++ b/tests/client/test_clnt_messenger.cc @@ -5,10 +5,16 @@ #include #include "client/clnt_messenger.h" +#include "data_server/kv_rpc_handler.h" +#include "proto/ds_clnt_rpcs.pb.h" #include "rpc/connection.h" #include "rpc/rpc.h" #include "rpc/rpc_context.h" +DECLARE_bool(clnt_syncreq_enable_retry); +DECLARE_uint32(clnt_syncreq_retry_count); +DECLARE_uint32(shard_total_num); + namespace simm { namespace clnt { @@ -69,6 +75,101 @@ class FakeConnection : public sicl::rpc::Connection { std::string name_; }; +class FakeSiRPC : public sicl::rpc::SiRPC { + public: + std::shared_ptr connect(const std::string &, const int) override { + connect_calls_.fetch_add(1); + if (connect_delay_.count() > 0) { + std::this_thread::sleep_for(connect_delay_); + } + if (connect_handler_) { + return connect_handler_(); + } + return std::make_shared("fake-connect"); + } + + std::shared_ptr connect(const std::string &, const int, const int, const int) override { + return nullptr; + } + + size_t addConnect(std::shared_ptr, size_t count) override { return count; } + + void SendRequest(const std::shared_ptr, + const sicl::rpc::ReqType, + const google::protobuf::Message &, + google::protobuf::Message *rsp, + const std::shared_ptr ctx, + const sicl::rpc::RpcDoneFn done) override { + async_send_calls_.fetch_add(1); + if (async_send_handler_) { + async_send_handler_(ctx); + } else { + ctx->SetError(sicl::transport::SICL_SUCCESS, ""); + } + done(rsp, ctx); + } + + void SendRequest(const std::shared_ptr, + const sicl::rpc::ReqType, + const google::protobuf::Message &, + google::protobuf::Message *, + const std::shared_ptr ctx) override { + sync_send_calls_.fetch_add(1); + if (sync_send_handler_) { + sync_send_handler_(ctx); + } else { + ctx->SetError(sicl::transport::SICL_SUCCESS, ""); + } + } + + void SendRequest(const std::string &, + const int, + const sicl::rpc::ReqType, + const google::protobuf::Message &, + google::protobuf::Message *, + const std::shared_ptr, + const sicl::rpc::RpcDoneFn) override {} + + int Stop() override { return 0; } + int Start(const int) override { return 0; } + void RunUntilAskedToQuit() override {} + bool RegisterHandler(const sicl::rpc::ReqType, const sicl::rpc::HandlerBase *) override { return true; } + void HandleRequest(const std::shared_ptr, + const std::shared_ptr, + const sicl::rpc::ReqType, + const void *, + const size_t) override {} + std::vector GetAllDevices() const override { return {}; } + sicl::transport::Mempool *GetMempool() const override { return nullptr; } + + void SetConnectHandler(std::function()> handler) { + connect_handler_ = std::move(handler); + } + + void SetSyncSendHandler(std::function &)> handler) { + sync_send_handler_ = std::move(handler); + } + + void SetAsyncSendHandler(std::function &)> handler) { + async_send_handler_ = std::move(handler); + } + + void SetConnectDelay(std::chrono::milliseconds delay) { connect_delay_ = delay; } + + int ConnectCalls() const { return connect_calls_.load(); } + int SyncSendCalls() const { return sync_send_calls_.load(); } + int AsyncSendCalls() const { return async_send_calls_.load(); } + + private: + std::atomic connect_calls_{0}; + std::atomic sync_send_calls_{0}; + std::atomic async_send_calls_{0}; + std::chrono::milliseconds connect_delay_{0}; + std::function()> connect_handler_{}; + std::function &)> sync_send_handler_{}; + std::function &)> async_send_handler_{}; +}; + std::shared_ptr MakeFailedRpcContext(int error_code) { sicl::rpc::RpcContext *ctx_raw = nullptr; sicl::rpc::RpcContext::newInstance(ctx_raw); @@ -195,13 +296,69 @@ class ClientMessengerTestPeer { static bool IsInitialized() { return ClientMessenger::Instance().initialized_; } static std::string CmAddr() { return ClientMessenger::Instance().cm_addr_; } + + static sicl::rpc::SiRPC *SwapRpcClient(sicl::rpc::SiRPC *replacement) { + auto &messenger = ClientMessenger::Instance(); + auto *original = messenger.rpc_client_; + messenger.rpc_client_ = replacement; + return original; + } + + static error_code_t CallSyncLookup(uint16_t shard_id) { + auto ctx = std::make_shared(); + sicl::rpc::RpcContext *ctx_p = nullptr; + sicl::rpc::RpcContext::newInstance(ctx_p); + auto rpc_ctx = std::shared_ptr(ctx_p); + ctx->set_rpc_ctx(rpc_ctx); + rpc_ctx->set_timeout(sicl::transport::TimerTick::TIMER_1S); + + KVLookupRequestPB req; + req.set_shard_id(shard_id); + req.set_key("mock-key"); + auto resp = std::make_shared(); + return ClientMessenger::Instance().call_sync( + shard_id, static_cast(simm::ds::KVServerRpcType::RPC_CLIENT_KV_LOOKUP), req, resp, ctx); + } + + static error_code_t BuildConnectionWait(const std::string &addr) { + return ClientMessenger::Instance().build_connection(addr, ClientMessenger::BuildConnWaitMode::kWaitForInflight); + } + + static error_code_t BuildConnectionNoWait(const std::string &addr) { + return ClientMessenger::Instance().build_connection(addr, ClientMessenger::BuildConnWaitMode::kNoWait); + } }; class ClientMessengerUnitTest : public ::testing::Test { protected: - void SetUp() override { ClientMessengerTestPeer::ResetState(); } + void SetUp() override { + ClientMessengerTestPeer::ResetState(); + original_rpc_client_ = ClientMessengerTestPeer::SwapRpcClient(nullptr); + } + + void TearDown() override { + if (injected_rpc_client_ != nullptr) { + auto *to_delete = ClientMessengerTestPeer::SwapRpcClient(original_rpc_client_); + if (to_delete != nullptr && to_delete != original_rpc_client_) { + delete to_delete; + } + injected_rpc_client_ = nullptr; + } else { + ClientMessengerTestPeer::SwapRpcClient(original_rpc_client_); + } + ClientMessengerTestPeer::ResetState(); + } + + FakeSiRPC *InstallFakeRpcClient() { + auto *fake = new FakeSiRPC(); + ClientMessengerTestPeer::SwapRpcClient(fake); + injected_rpc_client_ = fake; + return fake; + } - void TearDown() override { ClientMessengerTestPeer::ResetState(); } + private: + sicl::rpc::SiRPC *original_rpc_client_{nullptr}; + sicl::rpc::SiRPC *injected_rpc_client_{nullptr}; }; TEST_F(ClientMessengerUnitTest, PruneStaleConnectionsRemovesDeadDataservers) { @@ -262,6 +419,111 @@ TEST_F(ClientMessengerUnitTest, ConnectionAccessorsStayConsistentDuringConcurren EXPECT_TRUE(final_conn == conn_a || final_conn == conn_b); } +TEST_F(ClientMessengerUnitTest, SyncRetryFlagControlsRetryCountForTimeoutErrors) { + auto *fake_rpc = InstallFakeRpcClient(); + ClientMessengerTestPeer::InstallDsContext("10.0.0.1:1001", true, 1, {0}, std::make_shared("ready")); + + const auto old_enable_retry = FLAGS_clnt_syncreq_enable_retry; + const auto old_retry_count = FLAGS_clnt_syncreq_retry_count; + + fake_rpc->SetSyncSendHandler([](const std::shared_ptr &ctx) { + ctx->SetError(sicl::transport::SICL_ERR_TIMEOUT, "timeout"); + }); + + FLAGS_clnt_syncreq_enable_retry = false; + FLAGS_clnt_syncreq_retry_count = 2; + EXPECT_EQ(ClientMessengerTestPeer::CallSyncLookup(0), ClntErr::ClntSendRPCFailed); + EXPECT_EQ(fake_rpc->SyncSendCalls(), 1); + + FLAGS_clnt_syncreq_enable_retry = true; + FLAGS_clnt_syncreq_retry_count = 2; + EXPECT_EQ(ClientMessengerTestPeer::CallSyncLookup(0), ClntErr::ClntSendRPCFailed); + EXPECT_EQ(fake_rpc->SyncSendCalls(), 4); + + FLAGS_clnt_syncreq_enable_retry = old_enable_retry; + FLAGS_clnt_syncreq_retry_count = old_retry_count; +} + +TEST_F(ClientMessengerUnitTest, ConcurrentForegroundBuildConnectionFailsFastFollowers) { + auto *fake_rpc = InstallFakeRpcClient(); + fake_rpc->SetConnectDelay(std::chrono::milliseconds(200)); + fake_rpc->SetConnectHandler([]() { return std::make_shared("connected"); }); + ClientMessengerTestPeer::InstallDsContext("10.0.0.8:1008", false, 0, {0}, nullptr); + + std::atomic ready{0}; + std::atomic go{false}; + std::atomic ok_count{0}; + std::atomic fail_count{0}; + std::vector threads; + for (int i = 0; i < 8; ++i) { + threads.emplace_back([&]() { + ready.fetch_add(1); + while (!go.load()) { + std::this_thread::yield(); + } + auto ret = ClientMessengerTestPeer::BuildConnectionNoWait("10.0.0.8:1008"); + if (ret == CommonErr::OK) { + ok_count.fetch_add(1); + } else { + EXPECT_EQ(ret, ClntErr::BuildConnectionFailed); + fail_count.fetch_add(1); + } + }); + } + + while (ready.load() != 8) { + std::this_thread::yield(); + } + go.store(true); + + for (auto &t : threads) { + t.join(); + } + + EXPECT_EQ(fake_rpc->ConnectCalls(), 1); + EXPECT_GE(ok_count.load(), 1); + EXPECT_GE(fail_count.load(), 1); + EXPECT_TRUE(ClientMessengerTestPeer::IsDsActive("10.0.0.8:1008")); +} + +TEST_F(ClientMessengerUnitTest, ConcurrentBackgroundBuildConnectionWaitsAndSharesSingleConnect) { + auto *fake_rpc = InstallFakeRpcClient(); + fake_rpc->SetConnectDelay(std::chrono::milliseconds(100)); + fake_rpc->SetConnectHandler([]() { return std::make_shared("connected"); }); + ClientMessengerTestPeer::InstallDsContext("10.0.0.9:1009", false, 0, {0}, nullptr); + + std::atomic ready{0}; + std::atomic go{false}; + std::atomic ok_count{0}; + std::vector threads; + for (int i = 0; i < 6; ++i) { + threads.emplace_back([&]() { + ready.fetch_add(1); + while (!go.load()) { + std::this_thread::yield(); + } + auto ret = ClientMessengerTestPeer::BuildConnectionWait("10.0.0.9:1009"); + EXPECT_EQ(ret, CommonErr::OK); + if (ret == CommonErr::OK) { + ok_count.fetch_add(1); + } + }); + } + + while (ready.load() != 6) { + std::this_thread::yield(); + } + go.store(true); + + for (auto &t : threads) { + t.join(); + } + + EXPECT_EQ(fake_rpc->ConnectCalls(), 1); + EXPECT_EQ(ok_count.load(), 6); + EXPECT_TRUE(ClientMessengerTestPeer::IsDsActive("10.0.0.9:1009")); +} + TEST_F(ClientMessengerUnitTest, ReInitRefreshesRoutingBuildsConnectionsAndPrunesStaleServers) { auto healthy_conn = std::make_shared("healthy"); auto inactive_conn = std::make_shared("inactive");