已合并
test(distributed): add test for EtcdStore APIs for v2.12.0 #35276
zf_zhang创建于 5月11日
test(distributed): add test for EtcdStore APIs for v2.12.0 #35276
已合并
共 1 个文件变更+442-0
| @@ -0,0 +1,442 @@ | |||
| 1 | +""" | ||
| 2 | +PyTorch community lacks some torch.distributed.elastic.rendezvous.etcd_store APIs validation cases, so this file is added. | ||
| 3 | + | ||
| 4 | +This file validate following APIs: | ||
| 5 | +torch.distributed.elastic.rendezvous.etcd_store.EtcdStore | ||
| 6 | +torch.distributed.elastic.rendezvous.etcd_store.EtcdStore.set | ||
| 7 | +torch.distributed.elastic.rendezvous.etcd_store.EtcdStore.get | ||
| 8 | +torch.distributed.elastic.rendezvous.etcd_store.EtcdStore.add | ||
| 9 | +torch.distributed.elastic.rendezvous.etcd_store.EtcdStore.check | ||
| 10 | +torch.distributed.elastic.rendezvous.etcd_store.EtcdStore.wait | ||
| 11 | +torch.distributed.elastic.rendezvous.etcd_store.EtcdStore.set_timeout | ||
| 12 | +(Extendable) | ||
| 13 | +""" | ||
| 14 | + | ||
| 15 | +import base64 | ||
| 16 | +import concurrent.futures | ||
| 17 | +import datetime | ||
| 18 | +import subprocess | ||
| 19 | +import threading | ||
| 20 | +import time | ||
| 21 | + | ||
| 22 | +from torch.distributed.elastic.rendezvous.etcd_server import EtcdServer | ||
| 23 | +from torch.distributed.elastic.rendezvous.etcd_store import EtcdStore | ||
| 24 | +from torch.testing._internal.common_utils import run_tests, TestCase | ||
| 25 | + | ||
| 26 | + | ||
| 27 | +class EtcdTestBase(TestCase): | ||
| 28 | + """Base class with common setup - fails fast on any error""" | ||
| 29 | + | ||
| 30 | + | ||
| 31 | + def setUpClass(cls): | ||
| 32 | + """Class-level setup: Create and start a single EtcdServer for all tests""" | ||
| 33 | + cls.server = EtcdServer() | ||
| 34 | + cls.server.start(stderr=subprocess.DEVNULL) | ||
| 35 | + cls.client = cls.server.get_client() | ||
| 36 | + | ||
| 37 | + | ||
| 38 | + def tearDownClass(cls): | ||
| 39 | + """Class-level teardown: Stop the global EtcdServer""" | ||
| 40 | + cls.server.stop() | ||
| 41 | + | ||
| 42 | + def setUp(self): | ||
| 43 | + """Test-level setup: Create a clean EtcdStore with isolated prefix""" | ||
| 44 | + self.store = EtcdStore(self.client, f"/test_prefix_{self._testMethodName}/") | ||
| 45 | + | ||
| 46 | + def tearDown(self): | ||
| 47 | + pass | ||
| 48 | + | ||
| 49 | + | ||
| 50 | +class TestEtcdStoreInit(EtcdTestBase): | ||
| 51 | + """Test EtcdStore constructor and basic configuration""" | ||
| 52 | + | ||
| 53 | + def test_init_with_timeout(self): | ||
| 54 | + """Test passing a custom timeout to the constructor""" | ||
| 55 | + timeout = datetime.timedelta(seconds=30) | ||
| 56 | + store = EtcdStore(self.client, "/timeout_test/", timeout=timeout) | ||
| 57 | + | ||
| 58 | + start = time.time() | ||
| 59 | + with self.assertRaises(LookupError): | ||
| 60 | + store.get("nonexistent") | ||
| 61 | + elapsed = time.time() - start | ||
| 62 | + | ||
| 63 | + self.assertLess(elapsed, 35) | ||
| 64 | + | ||
| 65 | + def test_prefix_auto_append_slash(self): | ||
| 66 | + """Test that prefix automatically appends '/' if missing""" | ||
| 67 | + store = EtcdStore(self.client, "/no_slash") | ||
| 68 | + store.set("k", "v") | ||
| 69 | + | ||
| 70 | + encoded_key = base64.b64encode(b"k").decode() | ||
| 71 | + raw_value = self.client.get(f"/no_slash/{encoded_key}") | ||
| 72 | + | ||
| 73 | + self.assertEqual(raw_value.value, base64.b64encode(b"v").decode()) | ||
| 74 | + | ||
| 75 | + def test_prefix_with_existing_slash(self): | ||
| 76 | + """Test that prefix does not duplicate '/' when already present""" | ||
| 77 | + store = EtcdStore(self.client, "/with_slash/") | ||
| 78 | + store.set("k", "v") | ||
| 79 | + | ||
| 80 | + encoded_key = base64.b64encode(b"k").decode() | ||
| 81 | + raw_value = self.client.get(f"/with_slash/{encoded_key}") | ||
| 82 | + | ||
| 83 | + self.assertEqual(raw_value.value, base64.b64encode(b"v").decode()) | ||
| 84 | + | ||
| 85 | + | ||
| 86 | +class TestEtcdStoreSetGet(EtcdTestBase): | ||
| 87 | + """Test set and get API functionality""" | ||
| 88 | + | ||
| 89 | + def test_set_and_get_string(self): | ||
| 90 | + """Test storing and retrieving string values""" | ||
| 91 | + self.store.set("str_key", "hello") | ||
| 92 | + self.assertEqual(self.store.get("str_key"), b"hello") | ||
| 93 | + | ||
| 94 | + def test_set_and_get_bytes(self): | ||
| 95 | + """Test storing and retrieving raw bytes values""" | ||
| 96 | + binary_data = b"\x00\x01\x02\xff" | ||
| 97 | + self.store.set("bytes_key", binary_data) | ||
| 98 | + self.assertEqual(self.store.get("bytes_key"), binary_data) | ||
| 99 | + | ||
| 100 | + def test_set_overwrite(self): | ||
| 101 | + """Test overwriting an existing key""" | ||
| 102 | + self.store.set("overwrite_key", "v1") | ||
| 103 | + self.store.set("overwrite_key", "v2") | ||
| 104 | + self.assertEqual(self.store.get("overwrite_key"), b"v2") | ||
| 105 | + | ||
| 106 | + def test_set_invalid_type_raises(self): | ||
| 107 | + """Test that set raises ValueError for invalid value types""" | ||
| 108 | + invalid_values = [123, 45.6, ["list"], {"dict": "value"}, None] | ||
| 109 | + | ||
| 110 | + for val in invalid_values: | ||
| 111 | + with self.subTest(value=val): | ||
| 112 | + with self.assertRaises(ValueError): | ||
| 113 | + self.store.set("k", val) | ||
| 114 | + | ||
| 115 | + def test_set_special_characters(self): | ||
| 116 | + """Test keys with special characters (slashes, spaces, unicode, etc.)""" | ||
| 117 | + special_cases = { | ||
| 118 | + "key/with/nested/path": "slash_value", | ||
| 119 | + "key#hash": "hash_value", | ||
| 120 | + "key with space": "space_value", | ||
| 121 | + "key中文测试": "chinese_value", | ||
| 122 | + "key:colon": "colon_value", | ||
| 123 | + "key.dots": "dots_value", | ||
| 124 | + } | ||
| 125 | + | ||
| 126 | + for key, val in special_cases.items(): | ||
| 127 | + with self.subTest(key=key): | ||
| 128 | + self.store.set(key, val) | ||
| 129 | + self.assertEqual(self.store.get(key), val.encode()) | ||
| 130 | + | ||
| 131 | + def test_get_nonexistent_timeout(self): | ||
| 132 | + """Test that get raises LookupError on timeout for non-existent key""" | ||
| 133 | + self.store.set_timeout(datetime.timedelta(seconds=1)) | ||
| 134 | + | ||
| 135 | + start = time.time() | ||
| 136 | + with self.assertRaises(LookupError): | ||
| 137 | + self.store.get("never_exists_key") | ||
| 138 | + | ||
| 139 | + elapsed = time.time() - start | ||
| 140 | + self.assertGreaterEqual(elapsed, 0.9) | ||
| 141 | + self.assertLess(elapsed, 2.0) | ||
| 142 | + | ||
| 143 | + def test_get_blocks_until_available(self): | ||
| 144 | + """Test get blocks until key is set (using thread event, not sleep)""" | ||
| 145 | + key = "blocking_key" | ||
| 146 | + ready_event = threading.Event() | ||
| 147 | + | ||
| 148 | + def delayed_set(): | ||
| 149 | + ready_event.wait() | ||
| 150 | + time.sleep(0.2) | ||
| 151 | + self.store.set(key, "delayed_value") | ||
| 152 | + | ||
| 153 | + writer = threading.Thread(target=delayed_set) | ||
| 154 | + writer.start() | ||
| 155 | + ready_event.set() | ||
| 156 | + | ||
| 157 | + result = self.store.get(key) | ||
| 158 | + self.assertEqual(result, b"delayed_value") | ||
| 159 | + | ||
| 160 | + writer.join() | ||
| 161 | + | ||
| 162 | + def test_get_timeout_none_permanent_block(self): | ||
| 163 | + """Test get blocks permanently when timeout=None (released by thread)""" | ||
| 164 | + key = "permanent_block_key" | ||
| 165 | + | ||
| 166 | + def release_after(): | ||
| 167 | + time.sleep(0.5) | ||
| 168 | + self.store.set(key, "released") | ||
| 169 | + | ||
| 170 | + releaser = threading.Thread(target=release_after) | ||
| 171 | + releaser.start() | ||
| 172 | + | ||
| 173 | + result = self.store.get(key) | ||
| 174 | + self.assertEqual(result, b"released") | ||
| 175 | + | ||
| 176 | + releaser.join() | ||
| 177 | + | ||
| 178 | + def test_multiple_keys_isolation(self): | ||
| 179 | + """Test data isolation between multiple distinct keys""" | ||
| 180 | + keys = [f"isolated_key_{i}" for i in range(10)] | ||
| 181 | + | ||
| 182 | + for i, k in enumerate(keys): | ||
| 183 | + self.store.set(k, f"value_{i}") | ||
| 184 | + | ||
| 185 | + for i, k in enumerate(keys): | ||
| 186 | + self.assertEqual(self.store.get(k), f"value_{i}".encode()) | ||
| 187 | + | ||
| 188 | + | ||
| 189 | +class TestEtcdStoreAdd(EtcdTestBase): | ||
| 190 | + """Test add API (atomic increment)""" | ||
| 191 | + | ||
| 192 | + def test_add_initializes_when_key_absent(self): | ||
| 193 | + """Test add initializes key to 0 and increments when key does not exist""" | ||
| 194 | + result = self.store.add("init_counter", 5) | ||
| 195 | + self.assertEqual(result, 5) | ||
| 196 | + self.assertEqual(self.store.get("init_counter"), b"5") | ||
| 197 | + | ||
| 198 | + def test_add_increments_existing(self): | ||
| 199 | + """Test add increments an existing numeric value""" | ||
| 200 | + self.store.add("inc_counter", 10) | ||
| 201 | + result = self.store.add("inc_counter", 7) | ||
| 202 | + | ||
| 203 | + self.assertEqual(result, 17) | ||
| 204 | + self.assertEqual(self.store.get("inc_counter"), b"17") | ||
| 205 | + | ||
| 206 | + def test_add_zero(self): | ||
| 207 | + """Test add with 0 (idempotent operation)""" | ||
| 208 | + self.store.add("zero_test", 100) | ||
| 209 | + result = self.store.add("zero_test", 0) | ||
| 210 | + | ||
| 211 | + self.assertEqual(result, 100) | ||
| 212 | + self.assertEqual(self.store.get("zero_test"), b"100") | ||
| 213 | + | ||
| 214 | + def test_add_negative_decrement(self): | ||
| 215 | + """Test add with negative values (decrement behavior)""" | ||
| 216 | + self.store.add("dec_counter", 20) | ||
| 217 | + result = self.store.add("dec_counter", -5) | ||
| 218 | + | ||
| 219 | + self.assertEqual(result, 15) | ||
| 220 | + self.assertEqual(self.store.get("dec_counter"), b"15") | ||
| 221 | + | ||
| 222 | + def test_add_concurrent_20_threads(self): | ||
| 223 | + """Test concurrent add from 20 threads (exceptions propagate to main thread)""" | ||
| 224 | + key = "concurrent_20" | ||
| 225 | + | ||
| 226 | + def worker(): | ||
| 227 | + return self.store.add(key, 1) | ||
| 228 | + | ||
| 229 | + with concurrent.futures.ThreadPoolExecutor(max_workers=20) as executor: | ||
| 230 | + futures = [executor.submit(worker) for _ in range(20)] | ||
| 231 | + results = [f.result() for f in futures] | ||
| 232 | + | ||
| 233 | + self.assertEqual(len(results), 20) | ||
| 234 | + self.assertEqual(int(self.store.get(key)), 20) | ||
| 235 | + | ||
| 236 | + def test_add_high_contention_100_threads(self): | ||
| 237 | + """Test high-contention add with 100 threads (atomicity validation)""" | ||
| 238 | + key = "high_contention_100" | ||
| 239 | + | ||
| 240 | + def worker(): | ||
| 241 | + return self.store.add(key, 1) | ||
| 242 | + | ||
| 243 | + start_time = time.time() | ||
| 244 | + with concurrent.futures.ThreadPoolExecutor(max_workers=100) as executor: | ||
| 245 | + futures = [executor.submit(worker) for _ in range(100)] | ||
| 246 | + for f in futures: | ||
| 247 | + f.result() | ||
| 248 | + elapsed = time.time() - start_time | ||
| 249 | + | ||
| 250 | + final_value = int(self.store.get(key)) | ||
| 251 | + | ||
| 252 | + self.assertEqual( | ||
| 253 | + final_value, 100, f"Atomicity violation! Expected 100, got {final_value}" | ||
| 254 | + ) | ||
| 255 | + self.assertLess(elapsed, 60) | ||
| 256 | + | ||
| 257 | + def test_add_large_number(self): | ||
| 258 | + """Test add with large numeric values""" | ||
| 259 | + self.store.add("big_num", 10**9) | ||
| 260 | + result = self.store.add("big_num", 10**9) | ||
| 261 | + | ||
| 262 | + self.assertEqual(result, 2 * 10**9) | ||
| 263 | + | ||
| 264 | + | ||
| 265 | +class TestEtcdStoreWait(EtcdTestBase): | ||
| 266 | + """Test wait API (block until all keys exist)""" | ||
| 267 | + | ||
| 268 | + def test_wait_success_multiple_keys(self): | ||
| 269 | + """Test wait succeeds when all required keys are set""" | ||
| 270 | + keys = ["wait_k1", "wait_k2", "wait_k3"] | ||
| 271 | + ready = threading.Event() | ||
| 272 | + | ||
| 273 | + def delayed_write(): | ||
| 274 | + ready.wait() | ||
| 275 | + time.sleep(0.2) | ||
| 276 | + for k in keys: | ||
| 277 | + self.store.set(k, f"{k}_value") | ||
| 278 | + | ||
| 279 | + writer = threading.Thread(target=delayed_write) | ||
| 280 | + writer.start() | ||
| 281 | + ready.set() | ||
| 282 | + | ||
| 283 | + self.store.wait(keys) | ||
| 284 | + | ||
| 285 | + for k in keys: | ||
| 286 | + self.assertEqual(self.store.get(k), f"{k}_value".encode()) | ||
| 287 | + | ||
| 288 | + writer.join() | ||
| 289 | + | ||
| 290 | + def test_wait_timeout_global(self): | ||
| 291 | + """Test wait times out using the global store timeout""" | ||
| 292 | + self.store.set_timeout(datetime.timedelta(seconds=1)) | ||
| 293 | + | ||
| 294 | + start = time.time() | ||
| 295 | + with self.assertRaises(LookupError): | ||
| 296 | + self.store.wait(["never_key_1", "never_key_2"]) | ||
| 297 | + | ||
| 298 | + elapsed = time.time() - start | ||
| 299 | + self.assertGreaterEqual(elapsed, 0.9) | ||
| 300 | + self.assertLess(elapsed, 2.0) | ||
| 301 | + | ||
| 302 | + def test_wait_override_timeout(self): | ||
| 303 | + """Test wait with a custom override_timeout parameter""" | ||
| 304 | + start = time.time() | ||
| 305 | + | ||
| 306 | + with self.assertRaises(LookupError): | ||
| 307 | + self.store.wait( | ||
| 308 | + ["override_key"], override_timeout=datetime.timedelta(seconds=1) | ||
| 309 | + ) | ||
| 310 | + | ||
| 311 | + elapsed = time.time() - start | ||
| 312 | + self.assertGreaterEqual(elapsed, 0.9) | ||
| 313 | + self.assertLess(elapsed, 2.0) | ||
| 314 | + | ||
| 315 | + def test_wait_partial_keys_timeout(self): | ||
| 316 | + """Test wait times out if only some keys exist""" | ||
| 317 | + self.store.set("partial_exist", "yes") | ||
| 318 | + self.store.set_timeout(datetime.timedelta(seconds=1)) | ||
| 319 | + | ||
| 320 | + start = time.time() | ||
| 321 | + with self.assertRaises(LookupError): | ||
| 322 | + self.store.wait(["partial_exist", "missing_key"]) | ||
| 323 | + | ||
| 324 | + elapsed = time.time() - start | ||
| 325 | + self.assertGreaterEqual(elapsed, 0.9) | ||
| 326 | + | ||
| 327 | + | ||
| 328 | +class TestEtcdStoreCheck(EtcdTestBase): | ||
| 329 | + """Test check API (non-blocking existence check)""" | ||
| 330 | + | ||
| 331 | + def test_check_all_keys_exist(self): | ||
| 332 | + """Test check returns True when all keys exist""" | ||
| 333 | + self.store.set("check_a", "1") | ||
| 334 | + self.store.set("check_b", "2") | ||
| 335 | + | ||
| 336 | + result = self.store.check(["check_a", "check_b"]) | ||
| 337 | + self.assertTrue(result) | ||
| 338 | + | ||
| 339 | + def test_check_partial_exist(self): | ||
| 340 | + """Test check returns False when only some keys exist""" | ||
| 341 | + self.store.set("partial_a", "1") | ||
| 342 | + | ||
| 343 | + result = self.store.check(["partial_a", "partial_b"]) | ||
| 344 | + self.assertFalse(result) | ||
| 345 | + | ||
| 346 | + def test_check_none_exist(self): | ||
| 347 | + """Test check returns False when no keys exist""" | ||
| 348 | + result = self.store.check(["none_1", "none_2"]) | ||
| 349 | + self.assertFalse(result) | ||
| 350 | + | ||
| 351 | + def test_check_single_key(self): | ||
| 352 | + """Test check with a single key""" | ||
| 353 | + self.store.set("single", "1") | ||
| 354 | + | ||
| 355 | + self.assertTrue(self.store.check(["single"])) | ||
| 356 | + self.assertFalse(self.store.check(["not_exist"])) | ||
| 357 | + | ||
| 358 | + | ||
| 359 | +class TestEtcdStoreSetTimeout(EtcdTestBase): | ||
| 360 | + """Test set_timeout API for default timeout configuration""" | ||
| 361 | + | ||
| 362 | + def test_set_timeout_changes_default(self): | ||
| 363 | + """Test set_timeout modifies the default blocking timeout for get""" | ||
| 364 | + self.store.set_timeout(datetime.timedelta(seconds=1)) | ||
| 365 | + | ||
| 366 | + start = time.time() | ||
| 367 | + with self.assertRaises(LookupError): | ||
| 368 | + self.store.get("timeout_test_key") | ||
| 369 | + elapsed = time.time() - start | ||
| 370 | + | ||
| 371 | + self.assertLess(elapsed, 2.0) | ||
| 372 | + | ||
| 373 | + def test_set_timeout_zero_immediate(self): | ||
| 374 | + """Test timeout=0 causes immediate timeout (no blocking)""" | ||
| 375 | + self.store.set_timeout(datetime.timedelta(seconds=0)) | ||
| 376 | + | ||
| 377 | + start = time.time() | ||
| 378 | + with self.assertRaises(LookupError): | ||
| 379 | + self.store.get("immediate_timeout") | ||
| 380 | + elapsed = time.time() - start | ||
| 381 | + | ||
| 382 | + self.assertLess(elapsed, 0.5) | ||
| 383 | + | ||
| 384 | + def test_set_timeout_affects_wait(self): | ||
| 385 | + """Test set_timeout modifies the default timeout for wait""" | ||
| 386 | + self.store.set_timeout(datetime.timedelta(seconds=1)) | ||
| 387 | + | ||
| 388 | + start = time.time() | ||
| 389 | + with self.assertRaises(LookupError): | ||
| 390 | + self.store.wait(["wait_timeout_test"]) | ||
| 391 | + elapsed = time.time() - start | ||
| 392 | + | ||
| 393 | + self.assertGreaterEqual(elapsed, 0.9) | ||
| 394 | + self.assertLess(elapsed, 2.0) | ||
| 395 | + | ||
| 396 | + | ||
| 397 | +class TestEtcdStoreIntegration(EtcdTestBase): | ||
| 398 | + """Integration and end-to-end scenario tests""" | ||
| 399 | + | ||
| 400 | + def test_full_workflow_rendezvous_simulation(self): | ||
| 401 | + """Simulate distributed rendezvous: multiple nodes coordinate via EtcdStore""" | ||
| 402 | + num_nodes = 5 | ||
| 403 | + node_ids = [] | ||
| 404 | + lock = threading.Lock() | ||
| 405 | + | ||
| 406 | + def node_worker(node_id): | ||
| 407 | + self.store.add("node_count", 1) | ||
| 408 | + self.store.set(f"node_{node_id}_ready", "yes") | ||
| 409 | + self.store.wait([f"node_{i}_ready" for i in range(num_nodes)]) | ||
| 410 | + | ||
| 411 | + with lock: | ||
| 412 | + node_ids.append(node_id) | ||
| 413 | + | ||
| 414 | + with concurrent.futures.ThreadPoolExecutor(max_workers=num_nodes) as executor: | ||
| 415 | + futures = [executor.submit(node_worker, i) for i in range(num_nodes)] | ||
| 416 | + [f.result() for f in futures] | ||
| 417 | + | ||
| 418 | + self.assertEqual(len(node_ids), num_nodes) | ||
| 419 | + self.assertEqual(int(self.store.get("node_count")), num_nodes) | ||
| 420 | + | ||
| 421 | + def test_stress_multiple_operations(self): | ||
| 422 | + """Stress test: mixed operations across 50 concurrent workers""" | ||
| 423 | + | ||
| 424 | + def mixed_worker(worker_id): | ||
| 425 | + key = f"worker_{worker_id}" | ||
| 426 | + self.store.set(key, "initial") | ||
| 427 | + self.store.add(f"counter_{worker_id}", 1) | ||
| 428 | + self.store.check([key]) | ||
| 429 | + self.store.get(key) | ||
| 430 | + self.store.add(f"counter_{worker_id}", 1) | ||
| 431 | + return int(self.store.get(f"counter_{worker_id}")) | ||
| 432 | + | ||
| 433 | + with concurrent.futures.ThreadPoolExecutor(max_workers=50) as executor: | ||
| 434 | + futures = [executor.submit(mixed_worker, i) for i in range(50)] | ||
| 435 | + results = [f.result() for f in futures] | ||
| 436 | + | ||
| 437 | + for val in results: | ||
| 438 | + self.assertEqual(val, 2) | ||
| 439 | + | ||
| 440 | + | ||
| 441 | +if __name__ == "__main__": | ||
| 442 | + run_tests() | ||