From 07abf329da416e60fb7b6c7467cd0b5a1b79b880 Mon Sep 17 00:00:00 2001 From: Wang Xiaofeng Date: Tue, 6 Oct 2026 17:33:50 +0800 Subject: [PATCH] Reject truncated consistent hashing replica keys Check snprintf results before hashing replica keys in the default and Ketama policies. Reject formatting errors and truncated keys without partially updating the hash ring. Keep existing mappings for keys that fit the buffer. Cover key-length boundaries, batch additions, and disabled tag hashing for Murmur3, MD5, and Ketama. --- .../consistent_hashing_load_balancer.cpp | 8 +++ test/brpc_load_balancer_unittest.cpp | 56 +++++++++++++++++++ 2 files changed, 64 insertions(+) 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",