diff --git a/src/brpc/policy/consistent_hashing_load_balancer.cpp b/src/brpc/policy/consistent_hashing_load_balancer.cpp index 31d1ad744a..4b2f980bfc 100644 --- a/src/brpc/policy/consistent_hashing_load_balancer.cpp +++ b/src/brpc/policy/consistent_hashing_load_balancer.cpp @@ -94,6 +94,10 @@ bool DefaultReplicaPolicy::Build(ServerId server, len = snprintf(host, sizeof(host), "%s-%lu-%s", endpoint2str(ptr->remote_side()).c_str(), i, server.tag.c_str()); } + if (len < 0 || static_cast(len) >= sizeof(host)) { + LOG(ERROR) << "Invalid consistent hashing replica key length=" << len; + return false; + } ConsistentHashingLoadBalancer::Node node; node.hash = _hash_func(host, len); node.server_sock = server; @@ -133,6 +137,10 @@ bool KetamaReplicaPolicy::Build(ServerId server, len = snprintf(host, sizeof(host), "%s-%lu-%s", endpoint2str(ptr->remote_side()).c_str(), i, server.tag.c_str()); } + if (len < 0 || static_cast(len) >= sizeof(host)) { + LOG(ERROR) << "Invalid consistent hashing replica key length=" << len; + return false; + } unsigned char digest[MD5_DIGEST_LENGTH]; MD5HashSignature(host, len, digest); for (size_t j = 0; j < points_per_hash; ++j) { diff --git a/test/brpc_load_balancer_unittest.cpp b/test/brpc_load_balancer_unittest.cpp index 8748625c73..b0fd11c3f7 100644 --- a/test/brpc_load_balancer_unittest.cpp +++ b/test/brpc_load_balancer_unittest.cpp @@ -21,6 +21,7 @@ #include #include +#include #include #include "bthread/bthread.h" #include "gperftools_helper.h" @@ -47,6 +48,8 @@ namespace brpc { DECLARE_int32(health_check_interval); DECLARE_int64(detect_available_server_interval_ms); namespace policy { +DECLARE_int32(chash_num_replicas); +DECLARE_bool(consistent_hashing_enable_server_tag); extern uint32_t CRCHash32(const char *key, size_t len); extern const char* GetHashName(uint32_t (*hasher)(const void* key, size_t len)); }} @@ -825,6 +828,59 @@ TEST_F(LoadBalancerTest, consistent_hashing) { } } +TEST_F(LoadBalancerTest, consistent_hashing_replica_key_length) { + GFLAGS_NAMESPACE::FlagSaver flags_saver; + brpc::policy::FLAGS_chash_num_replicas = 100; + brpc::SocketOptions options; + ASSERT_EQ(0, butil::str2endpoint("127.0.0.1:8000", &options.remote_side)); + brpc::ServerId id; + ASSERT_EQ(0, brpc::Socket::Create(options, &id.id)); + const size_t prefix_size = + std::string(butil::endpoint2str(options.remote_side).c_str()).size() + 3; + const size_t tag_sizes[] = { + 0, 8, 254 - prefix_size, 255 - prefix_size, 256 - prefix_size, 4096 + }; + const brpc::policy::ConsistentHashingLoadBalancerType types[] = { + brpc::policy::CONS_HASH_LB_MURMUR3, + brpc::policy::CONS_HASH_LB_MD5, + brpc::policy::CONS_HASH_LB_KETAMA + }; + for (auto type : types) { + SCOPED_TRACE(type); + brpc::policy::FLAGS_consistent_hashing_enable_server_tag = true; + brpc::policy::ConsistentHashingLoadBalancer lb(type); + for (size_t tag_size : tag_sizes) { + SCOPED_TRACE(tag_size); + id.tag.assign(tag_size, 'x'); + // The two-digit replica indices use one more byte than "-0-". + const bool fits = prefix_size + tag_size + 1 < 256; + EXPECT_EQ(fits, lb.AddServer(id)); + if (fits) { + EXPECT_TRUE(lb.RemoveServer(id)); + } + brpc::SocketUniquePtr selected; + brpc::LoadBalancer::SelectIn in = { 0, false, true, 0, nullptr }; + brpc::LoadBalancer::SelectOut out(&selected); + // Even late truncation must not leave partial replicas on the ring. + EXPECT_EQ(ENODATA, lb.SelectServer(in, &out)); + } + + brpc::ServerId valid = id; + valid.tag = "valid"; + id.tag.assign(4096, 'x'); + std::vector servers = {id, valid}; + EXPECT_EQ(1u, lb.AddServersInBatch(servers)); + EXPECT_TRUE(lb.RemoveServer(valid)); + EXPECT_FALSE(lb.RemoveServer(id)); + + // A long tag must not affect the key when tags are disabled. + brpc::policy::FLAGS_consistent_hashing_enable_server_tag = false; + EXPECT_TRUE(lb.AddServer(id)); + EXPECT_TRUE(lb.RemoveServer(id)); + } + EXPECT_EQ(0, brpc::Socket::SetFailed(id.id)); +} + TEST_F(LoadBalancerTest, weighted_round_robin) { const char* servers[] = { "10.92.115.19:8831",