已合并
test(distributed): add test for EtcdStore APIs for v2.12.0 #35276
test(distributed): add test for EtcdStore APIs for v2.12.0 #35276
已合并
zf_zhang创建于 5月11日
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+ @classmethod
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+ @classmethod
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()