From 9b50b354b39403659ed796ef1bf26ac5640c5804 Mon Sep 17 00:00:00 2001 From: CAICAII <3360776475@qq.com> Date: Tue, 2 Jun 2026 10:33:08 +0800 Subject: [PATCH] [Store] Robustify ConfigDict size parsing (#2206) --- mooncake-store/include/utils.h | 29 +++++-- mooncake-store/src/real_client.cpp | 67 ++++++++------- mooncake-store/tests/pybind_client_test.cpp | 82 ++++++++++++++++--- mooncake-store/tests/utils_test.cpp | 14 ++++ .../tests/test_distributed_object_store.py | 40 +++++++++ 5 files changed, 179 insertions(+), 53 deletions(-) diff --git a/mooncake-store/include/utils.h b/mooncake-store/include/utils.h index 8e4d0994..d89c952f 100644 --- a/mooncake-store/include/utils.h +++ b/mooncake-store/include/utils.h @@ -171,14 +171,15 @@ std::string expected_to_str(const tl::expected& expected) { } /** - * @brief Convert a string representation of size to bytes + * @brief Parse a string representation of size to bytes * @param str String representation of size (e.g., "1.5 GB", "1024 MB", * "1048576") - * @return uint64_t Number of bytes, or 0 if parsing fails + * @return Parsed byte size, or std::nullopt if parsing fails */ -[[nodiscard]] inline uint64_t string_to_byte_size(const std::string& str) { +[[nodiscard]] inline std::optional try_string_to_byte_size( + const std::string& str) { if (str.empty()) { - return 0; + return std::nullopt; } // Create a copy for manipulation @@ -189,7 +190,7 @@ std::string expected_to_str(const tl::expected& expected) { s.erase(s.find_last_not_of(" \t\r\n") + 1); if (s.empty()) { - return 0; + return std::nullopt; } // Handle special case for "infinite" @@ -204,7 +205,10 @@ std::string expected_to_str(const tl::expected& expected) { try { value = std::stod(s, &pos); } catch (const std::exception&) { - return 0; // Failed to parse number + return std::nullopt; // Failed to parse number + } + if (value < 0) { + return std::nullopt; } if (pos >= s.length()) { @@ -238,10 +242,21 @@ std::string expected_to_str(const tl::expected& expected) { return static_cast(value); } else { // Unknown unit - return 0; + return std::nullopt; } } +/** + * @brief Convert a string representation of size to bytes + * @param str String representation of size (e.g., "1.5 GB", "1024 MB", + * "1048576") + * @return uint64_t Number of bytes, or 0 if parsing fails + */ +[[nodiscard]] inline uint64_t string_to_byte_size(const std::string& str) { + auto parsed = try_string_to_byte_size(str); + return parsed.value_or(0); +} + /** * @brief Convert a boolean-like string to a bool * @param str String representation ("1"/"true"/"yes"/"on" or diff --git a/mooncake-store/src/real_client.cpp b/mooncake-store/src/real_client.cpp index 9f017345..dbc126da 100644 --- a/mooncake-store/src/real_client.cpp +++ b/mooncake-store/src/real_client.cpp @@ -954,30 +954,21 @@ inline std::string get_config(const ConfigDict &config, const std::string &key, return (it != config.end()) ? it->second : default_value; } -inline size_t get_config_size(const ConfigDict &config, const std::string &key, - size_t default_value) { +inline std::optional get_config_size(const ConfigDict &config, + const std::string &key, + size_t default_value) { auto it = config.find(key); if (it == config.end()) { return default_value; } - const std::string &value = it->second; - // Check for negative numbers (stoull incorrectly parses "-1" as large val) - if (!value.empty() && value[0] == '-') { - LOG(WARNING) << "Invalid negative value for config key '" << key - << "': " << value << ", using default: " << default_value; - return default_value; - } - try { - return std::stoull(value); - } catch (const std::invalid_argument &e) { - LOG(WARNING) << "Invalid non-numeric value for config key '" << key - << "': " << value << ", using default: " << default_value; - return default_value; - } catch (const std::out_of_range &e) { - LOG(WARNING) << "Value out of range for config key '" << key - << "': " << value << ", using default: " << default_value; - return default_value; + + auto parsed_size_opt = try_string_to_byte_size(it->second); + if (!parsed_size_opt.has_value()) { + LOG(ERROR) << "Invalid size value for config key '" << key + << "': " << it->second; + return std::nullopt; } + return static_cast(parsed_size_opt.value()); } } // namespace @@ -999,10 +990,18 @@ tl::expected RealClient::setup_internal( } // Extract optional parameters with defaults - size_t global_segment_size = get_config_size( + auto global_segment_size_opt = get_config_size( config, CONFIG_KEY_GLOBAL_SEGMENT_SIZE, DEFAULT_GLOBAL_SEGMENT_SIZE); - size_t local_buffer_size = get_config_size( + if (!global_segment_size_opt.has_value()) { + return tl::unexpected(ErrorCode::INVALID_PARAMS); + } + auto local_buffer_size_opt = get_config_size( config, CONFIG_KEY_LOCAL_BUFFER_SIZE, DEFAULT_LOCAL_BUFFER_SIZE); + if (!local_buffer_size_opt.has_value()) { + return tl::unexpected(ErrorCode::INVALID_PARAMS); + } + size_t global_segment_size = global_segment_size_opt.value(); + size_t local_buffer_size = local_buffer_size_opt.value(); std::string protocol = get_config(config, CONFIG_KEY_PROTOCOL, DEFAULT_PROTOCOL); std::string rdma_devices = get_config(config, CONFIG_KEY_RDMA_DEVICES); @@ -1011,19 +1010,19 @@ tl::expected RealClient::setup_internal( std::string ipc_socket_path = get_config(config, CONFIG_KEY_IPC_SOCKET_PATH); - // Validate size parameters are within acceptable ranges - if (global_segment_size < MIN_SEGMENT_SIZE || - global_segment_size > MAX_SEGMENT_SIZE) { - LOG(ERROR) << "Invalid " << CONFIG_KEY_GLOBAL_SEGMENT_SIZE << ": " - << global_segment_size << ", must be between " - << MIN_SEGMENT_SIZE << " and " << MAX_SEGMENT_SIZE; - return tl::unexpected(ErrorCode::INVALID_PARAMS); - } - if (local_buffer_size < MIN_SEGMENT_SIZE || - local_buffer_size > MAX_SEGMENT_SIZE) { - LOG(ERROR) << "Invalid " << CONFIG_KEY_LOCAL_BUFFER_SIZE << ": " - << local_buffer_size << ", must be between " - << MIN_SEGMENT_SIZE << " and " << MAX_SEGMENT_SIZE; + // A size of 0 keeps the pure client/server setup semantics. + auto validate_size = [](const char *key, size_t value) { + if ((value != 0 && value < MIN_SEGMENT_SIZE) || + value > MAX_SEGMENT_SIZE) { + LOG(ERROR) << "Invalid " << key << ": " << value + << ", must be 0 or between " << MIN_SEGMENT_SIZE + << " and " << MAX_SEGMENT_SIZE; + return false; + } + return true; + }; + if (!validate_size(CONFIG_KEY_GLOBAL_SEGMENT_SIZE, global_segment_size) || + !validate_size(CONFIG_KEY_LOCAL_BUFFER_SIZE, local_buffer_size)) { return tl::unexpected(ErrorCode::INVALID_PARAMS); } diff --git a/mooncake-store/tests/pybind_client_test.cpp b/mooncake-store/tests/pybind_client_test.cpp index dec33176..28cf616b 100644 --- a/mooncake-store/tests/pybind_client_test.cpp +++ b/mooncake-store/tests/pybind_client_test.cpp @@ -85,6 +85,23 @@ class RealClientTest : public ::testing::Test { 0); } + ConfigDict MakeConfigDict(const std::string& local_hostname, + const std::string& global_segment_size, + const std::string& local_buffer_size) const { + const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) + ? FLAGS_device_name + : std::string(""); + ConfigDict config; + config[CONFIG_KEY_LOCAL_HOSTNAME] = local_hostname; + config[CONFIG_KEY_METADATA_SERVER] = "P2PHANDSHAKE"; + config[CONFIG_KEY_GLOBAL_SEGMENT_SIZE] = global_segment_size; + config[CONFIG_KEY_LOCAL_BUFFER_SIZE] = local_buffer_size; + config[CONFIG_KEY_PROTOCOL] = FLAGS_protocol; + config[CONFIG_KEY_RDMA_DEVICES] = rdma_devices; + config[CONFIG_KEY_MASTER_SERVER_ADDR] = master_address_; + return config; + } + std::string CreateTempSegmentFile(size_t size) { std::string path = "/tmp/mooncake_real_client_segment_XXXXXX"; int fd = mkstemp(path.data()); @@ -889,22 +906,12 @@ TEST_F(RealClientTest, SetupWithConfigDict) { master_address_ = master_.master_address(); LOG(INFO) << "Started in-proc master at " << master_address_; - // Setup the client using ConfigDict - const std::string rdma_devices = (FLAGS_protocol == std::string("rdma")) - ? FLAGS_device_name - : std::string(""); - ConfigDict config; auto result = py_client_->setup_internal(config); ASSERT_FALSE(result.has_value()) << "Setup with empty config should fail"; - config[CONFIG_KEY_LOCAL_HOSTNAME] = "localhost:17813"; - config[CONFIG_KEY_METADATA_SERVER] = "P2PHANDSHAKE"; - config[CONFIG_KEY_GLOBAL_SEGMENT_SIZE] = std::to_string(16 * 1024 * 1024); - config[CONFIG_KEY_LOCAL_BUFFER_SIZE] = std::to_string(16 * 1024 * 1024); - config[CONFIG_KEY_PROTOCOL] = FLAGS_protocol; - config[CONFIG_KEY_RDMA_DEVICES] = rdma_devices; - config[CONFIG_KEY_MASTER_SERVER_ADDR] = master_address_; + config = MakeConfigDict("localhost:17813", std::to_string(16 * 1024 * 1024), + std::to_string(16 * 1024 * 1024)); result = py_client_->setup_internal(config); ASSERT_TRUE(result.has_value()) << "Setup with ConfigDict should succeed"; @@ -929,6 +936,57 @@ TEST_F(RealClientTest, SetupWithConfigDict) { EXPECT_EQ(retrieved_data, test_data) << "Retrieved data should match"; } +TEST_F(RealClientTest, SetupWithConfigDictHumanReadableSizes) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + ConfigDict config = MakeConfigDict("localhost:17814", "16MB", "16 MB"); + auto result = py_client_->setup_internal(config); + ASSERT_TRUE(result.has_value()) + << "Setup should accept human-readable size strings"; +} + +TEST_F(RealClientTest, SetupWithConfigDictAllowsZeroSizes) { + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + ConfigDict config = MakeConfigDict("localhost:17815", "0", "0"); + auto result = py_client_->setup_internal(config); + ASSERT_TRUE(result.has_value()) + << "Setup should preserve zero-size pure client/server semantics"; +} + +TEST_F(RealClientTest, ErrSetupWithInvalidConfigDictSize) { + GLogMuter muter; + ASSERT_TRUE(master_.Start(InProcMasterConfigBuilder().build())) + << "Failed to start in-proc master"; + master_address_ = master_.master_address(); + + struct InvalidSizeCase { + const char* local_hostname; + const char* global_segment_size; + const char* local_buffer_size; + }; + + const InvalidSizeCase invalid_size_cases[] = { + {"localhost:17816", "50%", "16MB"}, + {"localhost:17817", "16MB", "16XB"}, + {"localhost:17818", "-5", "16MB"}, + }; + + for (const auto& test_case : invalid_size_cases) { + ConfigDict config = MakeConfigDict(test_case.local_hostname, + test_case.global_segment_size, + test_case.local_buffer_size); + auto result = py_client_->setup_internal(config); + EXPECT_FALSE(result.has_value()) + << "Invalid explicit size values should fail instead of being " + "partially parsed or silently defaulted"; + } +} + TEST_F(RealClientTest, ErrSetupWithInvalidArgument) { GLogMuter muter; // Case 1: Setup with unreachable master address diff --git a/mooncake-store/tests/utils_test.cpp b/mooncake-store/tests/utils_test.cpp index 32f703f4..b3102fd9 100644 --- a/mooncake-store/tests/utils_test.cpp +++ b/mooncake-store/tests/utils_test.cpp @@ -23,6 +23,20 @@ TEST(UtilsTest, ByteSizeToString) { EXPECT_EQ(byte_size_to_string(15 * 1024 * 1024 + 44048), "15.04 MB"); } +TEST(UtilsTest, StringToByteSize) { + auto parsed = try_string_to_byte_size("16 MB"); + ASSERT_TRUE(parsed.has_value()); + EXPECT_EQ(parsed.value(), 16ULL * 1024 * 1024); + + parsed = try_string_to_byte_size("0"); + ASSERT_TRUE(parsed.has_value()); + EXPECT_EQ(parsed.value(), 0); + + EXPECT_FALSE(try_string_to_byte_size("-5").has_value()); + EXPECT_FALSE(try_string_to_byte_size("16XB").has_value()); + EXPECT_EQ(string_to_byte_size("-5"), 0); +} + TEST(UtilsTest, StringToBool) { EXPECT_EQ(string_to_bool("1"), true); EXPECT_EQ(string_to_bool("true"), true); diff --git a/mooncake-wheel/tests/test_distributed_object_store.py b/mooncake-wheel/tests/test_distributed_object_store.py index 52530d3b..9e5d9e96 100644 --- a/mooncake-wheel/tests/test_distributed_object_store.py +++ b/mooncake-wheel/tests/test_distributed_object_store.py @@ -38,6 +38,46 @@ def get_client(store, local_buffer_size_param=None): if retcode: raise RuntimeError(f"Failed to setup store client. Return code: {retcode}") + +def get_config_dict(global_segment_size, local_buffer_size): + """Build a config dictionary for the MooncakeDistributedStore setup wrapper.""" + return { + "local_hostname": os.getenv("LOCAL_HOSTNAME", "localhost"), + "metadata_server": os.getenv( + "MC_METADATA_SERVER", "http://127.0.0.1:8080/metadata" + ), + "global_segment_size": global_segment_size, + "local_buffer_size": local_buffer_size, + "protocol": os.getenv("PROTOCOL", "tcp"), + "rdma_devices": os.getenv("DEVICE_NAME", "ibp6s0"), + "master_server_addr": os.getenv("MASTER_SERVER", "127.0.0.1:50051"), + } + + +class TestConfigDictSetup(unittest.TestCase): + """Test configuration-dictionary setup through the Python store wrapper.""" + + def test_human_readable_sizes(self): + store = MooncakeDistributedStore() + self.addCleanup(store.close) + + retcode = store.setup(get_config_dict("16MB", "16 MB")) + self.assertEqual(retcode, 0) + + test_data = b"test_config_dict_human_readable_value" + key = f"test_config_dict_human_readable_key_{os.getpid()}" + + self.assertEqual(store.put(key, test_data), 0) + self.assertEqual(store.get(key), test_data) + + def test_unsupported_percentage_size(self): + store = MooncakeDistributedStore() + self.addCleanup(store.close) + + retcode = store.setup(get_config_dict("50%", "16MB")) + self.assertNotEqual(retcode, 0) + + class TestZeroLocalBufferSize(unittest.TestCase): """Test class for zero local buffer size scenarios."""