diff --git a/changelog.d/549.changed.md b/changelog.d/549.changed.md new file mode 100644 index 00000000..37ca9e04 --- /dev/null +++ b/changelog.d/549.changed.md @@ -0,0 +1 @@ +Raise the first simulation cache memory warning from 8 GiB to 16 GiB while retaining the 32 GiB warning. diff --git a/src/policyengine/core/cache.py b/src/policyengine/core/cache.py index 410301e4..4761ff2f 100644 --- a/src/policyengine/core/cache.py +++ b/src/policyengine/core/cache.py @@ -6,7 +6,7 @@ logger = logging.getLogger(__name__) -_MEMORY_THRESHOLDS_GB = [8, 16, 32] +_MEMORY_THRESHOLDS_GIB = [16, 32] _warned_thresholds: set[int] = set() T = TypeVar("T") @@ -50,10 +50,11 @@ def _check_memory_usage(self) -> None: process = psutil.Process() memory_gb = process.memory_info().rss / (1024**3) - for threshold in _MEMORY_THRESHOLDS_GB: + for threshold in _MEMORY_THRESHOLDS_GIB: if memory_gb >= threshold and threshold not in _warned_thresholds: logger.warning( - f"Memory usage has reached {memory_gb:.2f}GB (threshold: {threshold}GB). " + f"Memory usage has reached {memory_gb:.2f} GiB " + f"(threshold: {threshold} GiB). " f"Cache contains {len(self._cache)} items." ) _warned_thresholds.add(threshold) diff --git a/tests/test_cache.py b/tests/test_cache.py index cb16ef65..4839551e 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -1,10 +1,13 @@ import os import tempfile +from types import SimpleNamespace import pandas as pd +import pytest from microdf import MicroDataFrame from policyengine.core import Simulation +from policyengine.core import cache as cache_module from policyengine.core.cache import LRUCache from policyengine.tax_benefit_models.uk import ( PolicyEngineUKDataset, @@ -143,3 +146,74 @@ def test_lru_cache_clear(): assert cache.get("a") is None assert cache.get("b") is None assert cache.get("c") is None + + +@pytest.mark.parametrize("memory_gib", [0, 8, 15.99]) +def test_lru_cache_does_not_warn_below_16_gib( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, + memory_gib: float, +) -> None: + cache_module._warned_thresholds.clear() + monkeypatch.setattr( + cache_module.psutil, + "Process", + lambda: SimpleNamespace( + memory_info=lambda: SimpleNamespace(rss=memory_gib * 1024**3) + ), + ) + + with caplog.at_level("WARNING", logger=cache_module.__name__): + LRUCache[str]().add("key", "value") + + assert caplog.records == [] + + +def test_lru_cache_warns_once_at_16_gib_and_again_at_32_gib( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + cache_module._warned_thresholds.clear() + memory = SimpleNamespace(gib=16.0) + monkeypatch.setattr( + cache_module.psutil, + "Process", + lambda: SimpleNamespace( + memory_info=lambda: SimpleNamespace(rss=memory.gib * 1024**3) + ), + ) + cache = LRUCache[str]() + + with caplog.at_level("WARNING", logger=cache_module.__name__): + cache.add("first", "value") + cache.add("second", "value") + memory.gib = 32.0 + cache.add("third", "value") + + messages = [record.getMessage() for record in caplog.records] + assert len(messages) == 2 + assert "threshold: 16 GiB" in messages[0] + assert "threshold: 32 GiB" in messages[1] + + +def test_lru_cache_clear_resets_memory_warning_state( + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + cache_module._warned_thresholds.clear() + monkeypatch.setattr( + cache_module.psutil, + "Process", + lambda: SimpleNamespace(memory_info=lambda: SimpleNamespace(rss=16 * 1024**3)), + ) + cache = LRUCache[str]() + + with caplog.at_level("WARNING", logger=cache_module.__name__): + cache.add("first", "value") + cache.clear() + cache.add("second", "value") + + assert ( + sum("threshold: 16 GiB" in record.getMessage() for record in caplog.records) + == 2 + )