diff --git a/tests/unit/_utils/test_system.py b/tests/unit/_utils/test_system.py index 26ec8a7204..0c2451fde6 100644 --- a/tests/unit/_utils/test_system.py +++ b/tests/unit/_utils/test_system.py @@ -2,10 +2,10 @@ import logging import sys -from multiprocessing import get_context, synchronize +from multiprocessing import get_context, resource_tracker, synchronize from multiprocessing.shared_memory import SharedMemory from types import SimpleNamespace -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, cast from unittest.mock import Mock import proclimits @@ -19,6 +19,7 @@ if TYPE_CHECKING: from collections.abc import Callable + from multiprocessing.context import ForkContext, ForkServerContext, SpawnContext HOST_TOTAL_BYTES = 8 * 1024**3 HOST_AVAILABLE_BYTES = 3 * 1024**3 @@ -424,11 +425,51 @@ def test_log_resource_limits_lets_a_failing_sensor_surface(monkeypatch: pytest.M snapshot.assert_called_once() +EXTRA_MEMORY_SIZE = 1024 * 1024 * 100 # 100 MB + + +# The children of the estimation test below live at module level, so that every start method can pickle them. +def no_extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None: + ready.wait() + measured.wait() + + +def extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None: + memory = SharedMemory(size=EXTRA_MEMORY_SIZE, create=True) + assert memory.buf is not None + fill_buffer(memory.buf, EXTRA_MEMORY_SIZE) + print(f'Using the memory... {memory.buf[-1]}') + ready.wait() + measured.wait() + memory.close() + memory.unlink() + + +def shared_extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier, memory: SharedMemory) -> None: + assert memory.buf is not None + # Fault every page in: untouched pages never enter the RSS (hiding the overcount this test guards + # against) and are reclaimed first under memory pressure, which drops them from every mapper's PSS. + page_sum = sum(memory.buf[::4096]) + print(f'Using the memory... {page_sum}') + ready.wait() + measured.wait() + + # The estimation is asserted on absolute memory readings, which hold only as long as nothing else on the machine makes # the kernel reclaim the pages allocated below. Running alongside the other test workers is enough to break that. @pytest.mark.run_alone @pytest.mark.skipif(sys.platform != 'linux', reason='Improved estimation available only on Linux') -def test_memory_estimation_does_not_overestimate_due_to_shared_memory() -> None: +# The start methods differ in how much memory the children share with the rest of the process tree, which is what the +# estimation has to account for. The default is `fork` up to Python 3.13 and `forkserver` from 3.14 on. +@pytest.mark.parametrize( + 'start_method', + [ + pytest.param('fork', id='fork'), + pytest.param('forkserver', id='forkserver'), + pytest.param('spawn', id='spawn'), + ], +) +def test_memory_estimation_does_not_overestimate_due_to_shared_memory(start_method: str) -> None: """Test that memory usage estimation is not overestimating memory usage by counting shared memory multiple times. In this test, the parent process is started and its memory usage is measured in situations where it is running @@ -440,41 +481,17 @@ def test_memory_estimation_does_not_overestimate_due_to_shared_memory() -> None: the same as the unshared memory. """ - ctx = get_context('fork') - estimated_memory_expectation = ctx.Value('b', False) # noqa: FBT003 # Common usage pattern for multiprocessing.Value + # The measuring process is a closure, which only `fork` can start. The children are started by `start_method`. + fork_ctx = get_context('fork') + estimated_memory_expectation = fork_ctx.Value('b', False) # noqa: FBT003 # Common usage pattern for multiprocessing.Value def parent_process() -> None: - extra_memory_size = 1024 * 1024 * 100 # 100 MB + ctx = cast('ForkContext | ForkServerContext | SpawnContext', get_context(start_method)) children_count = 4 # Memory calculation is not exact, so allow for some tolerance. test_tolerance = 0.3 measurement_rounds = 3 - def no_extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None: - ready.wait() - measured.wait() - - def extra_memory_child(ready: synchronize.Barrier, measured: synchronize.Barrier) -> None: - memory = SharedMemory(size=extra_memory_size, create=True) - assert memory.buf is not None - fill_buffer(memory.buf, extra_memory_size) - print(f'Using the memory... {memory.buf[-1]}') - ready.wait() - measured.wait() - memory.close() - memory.unlink() - - def shared_extra_memory_child( - ready: synchronize.Barrier, measured: synchronize.Barrier, memory: SharedMemory - ) -> None: - assert memory.buf is not None - # Fault every page in: untouched pages never enter the RSS (hiding the overcount this test guards - # against) and are reclaimed first under memory pressure, which drops them from every mapper's PSS. - page_sum = sum(memory.buf[::4096]) - print(f'Using the memory... {page_sum}') - ready.wait() - measured.wait() - def get_additional_memory_estimation_while_running_processes( *, target: Callable, count: int = 1, use_shared_memory: bool = False ) -> float: @@ -485,9 +502,9 @@ def get_additional_memory_estimation_while_running_processes( memory_before = get_memory_info().current_size if use_shared_memory: - shared_memory = SharedMemory(size=extra_memory_size, create=True) + shared_memory = SharedMemory(size=EXTRA_MEMORY_SIZE, create=True) assert shared_memory.buf is not None - fill_buffer(shared_memory.buf, extra_memory_size) + fill_buffer(shared_memory.buf, EXTRA_MEMORY_SIZE) extra_args = [shared_memory] else: extra_args = [] @@ -510,6 +527,18 @@ def get_additional_memory_estimation_while_running_processes( return (memory_during - memory_before).to_mb() / count + # Some start methods launch long-lived helper processes (the fork server, the resource tracker) on first use. + # They belong to the process tree and so to the estimate, but started inside a round they would inflate its + # baseline. Start them ahead of the measurements, keeping the barriers referenced until the child is done. The + # resource tracker is started explicitly: under `fork` no barrier registers with it, and every child creating + # shared memory would start one of its own. + resource_tracker.ensure_running() + ready, measured = ctx.Barrier(parties=1), ctx.Barrier(parties=1) + warm_up = ctx.Process(target=no_extra_memory_child, args=[ready, measured]) + warm_up.start() + warm_up.join() + assert warm_up.exitcode == 0 + # Under memory pressure the kernel reclaims cold pages, which silently leave the PSS readings and skew a # round's deltas, so a distorted round is re-measured. A genuine overcount of shared memory misses the # expectation several times over in every round, so the retries cannot mask it. @@ -549,10 +578,10 @@ def get_additional_memory_estimation_while_running_processes( f'{memory_estimation_difference_ratio=}' ) - process = ctx.Process(target=parent_process) + process = fork_ctx.Process(target=parent_process) process.start() process.join() assert estimated_memory_expectation.value, ( - 'Estimated memory usage for process with shared memory does not meet the expectation.' + f'Estimated memory usage for process with shared memory does not meet the expectation under {start_method}.' )