Skip to content

Commit 30ed0df

Browse files
committed
Strengthen merge key DoS timeout tests
1 parent a838219 commit 30ed0df

1 file changed

Lines changed: 24 additions & 17 deletions

File tree

tests/legacy_tests/test_constructor.py

Lines changed: 24 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -296,37 +296,45 @@ def test_subclass_blacklist_types(data_filename, verbose=False):
296296

297297
test_subclass_blacklist_types.unittest = ['.subclass_blacklist']
298298

299+
def run_with_timeout(payload, seconds, message):
300+
if not hasattr(signal, 'SIGALRM'):
301+
yaml.safe_load(payload)
302+
return
303+
304+
def timeout_handler(signum, frame):
305+
raise AssertionError(message)
306+
307+
previous_handler = signal.signal(signal.SIGALRM, timeout_handler)
308+
try:
309+
signal.alarm(seconds)
310+
yaml.safe_load(payload)
311+
finally:
312+
signal.alarm(0)
313+
signal.signal(signal.SIGALRM, previous_handler)
314+
299315
def test_merge_key_dos_prevention(verbose=False):
300316
# Warm up to trigger lazy loading / imports
301317
yaml.safe_load("L0: &L0 { x: 1 }")
302-
318+
303319
lines = ["L0: &L0 { x: 1 }"]
304-
for level in range(1, 21):
320+
for level in range(1, 27):
305321
lines.append(f"L{level}: &L{level} {{ <<: [*L{level-1}, *L{level-1}] }}")
306322
payload = "\n".join(lines)
307-
308-
import time
309-
start = time.time()
310-
yaml.safe_load(payload)
311-
elapsed = time.time() - start
312-
assert elapsed < 0.3, f"Merge key DoS took too long: {elapsed:.4f}s"
323+
324+
run_with_timeout(payload, 1, "Merge key DoS took too long")
313325

314326
test_merge_key_dos_prevention.unittest = True
315327

316328
def test_multi_merge_key_dos_prevention(verbose=False):
317329
# Warm up to trigger lazy loading / imports
318330
yaml.safe_load("L0: &L0 { x: 1 }")
319-
331+
320332
lines = ["L0: &L0 { x: 1 }"]
321-
for level in range(1, 19):
333+
for level in range(1, 27):
322334
lines.append(f"L{level}: &L{level}\n <<: *L{level-1}\n <<: *L{level-1}")
323335
payload = "\n".join(lines)
324-
325-
import time
326-
start = time.time()
327-
yaml.safe_load(payload)
328-
elapsed = time.time() - start
329-
assert elapsed < 0.3, f"Multi-merge key DoS took too long: {elapsed:.4f}s"
336+
337+
run_with_timeout(payload, 1, "Multi-merge key DoS took too long")
330338

331339
test_multi_merge_key_dos_prevention.unittest = True
332340

@@ -335,4 +343,3 @@ def test_multi_merge_key_dos_prevention(verbose=False):
335343
sys.modules['test_constructor'] = sys.modules['__main__']
336344
import test_appliance
337345
test_appliance.run(globals())
338-

0 commit comments

Comments
 (0)