diff --git a/test/web_interface/test_cache.py b/test/web_interface/test_cache.py index 46f6b0ee..90d260bb 100644 --- a/test/web_interface/test_cache.py +++ b/test/web_interface/test_cache.py @@ -1,9 +1,14 @@ """Tests for the web interface's in-memory cache helpers.""" +import sys +import threading from typing import Iterator import pytest -from web_interface.cache import delete_cached, get_cached, invalidate_cache, set_cached +from web_interface import cache as cache_module +from web_interface.cache import ( + TTLCache, delete_cached, get_cached, invalidate_cache, set_cached, +) @pytest.fixture(autouse=True) @@ -44,3 +49,142 @@ def test_invalidate_cache_pattern() -> None: invalidate_cache('fonts') assert get_cached('fonts_catalog') is None assert get_cached('plugins_list') == 2 + + +# --------------------------------------------------------------------------- +# Expiry. set_cached used to accept ttl_seconds and ignore it; only the TTL a +# reader passed to get_cached counted, and get_cached defaulted to 60s. +# --------------------------------------------------------------------------- + +class _Clock: + def __init__(self) -> None: + self.now = 1000.0 + + def __call__(self) -> float: + return self.now + + +@pytest.fixture +def clock() -> _Clock: + return _Clock() + + +def test_entry_expires_after_its_ttl(clock: _Clock) -> None: + c = TTLCache(clock=clock) + c.set('k', 'v', ttl=10) + clock.now += 9.9 + assert c.get('k') == 'v' + clock.now += 0.1 + assert c.get('k') is None + + +def test_default_ttl_applies_when_none_given(clock: _Clock) -> None: + c = TTLCache(default_ttl=5, clock=clock) + c.set('k', 'v') + clock.now += 4.9 + assert c.get('k') == 'v' + clock.now += 0.1 + assert c.get('k') is None + + +def test_reader_max_age_can_only_shorten(clock: _Clock) -> None: + c = TTLCache(clock=clock) + c.set('k', 'v', ttl=10) + clock.now += 5 + assert c.get('k', max_age=6) == 'v' + assert c.get('k', max_age=5) is None + clock.now += 5 + assert c.get('k', max_age=60) is None, "a reader extended a 10s entry" + + +def test_set_cached_ttl_is_honoured(monkeypatch: pytest.MonkeyPatch, clock: _Clock) -> None: + monkeypatch.setattr(cache_module, '_default_cache', TTLCache(clock=clock)) + set_cached('short', 1, ttl_seconds=2) + set_cached('long', 2, ttl_seconds=300) + clock.now += 2 + assert get_cached('short') is None, "set_cached ignored its ttl_seconds" + clock.now += 100 # past the old implicit 60s read default + assert get_cached('long') == 2 + + +def test_get_cached_ttl_still_bounds_the_read(monkeypatch: pytest.MonkeyPatch, clock: _Clock) -> None: + """The existing callers pass the TTL on both sides; that keeps working.""" + monkeypatch.setattr(cache_module, '_default_cache', TTLCache(clock=clock)) + set_cached('system_status', {'cpu': 1}, ttl_seconds=10) + clock.now += 9 + assert get_cached('system_status', ttl_seconds=10) == {'cpu': 1} + clock.now += 1 + assert get_cached('system_status', ttl_seconds=10) is None + + +def test_peek_returns_the_last_value_after_expiry(clock: _Clock) -> None: + c = TTLCache(clock=clock) + assert c.peek('k', 'fallback') == 'fallback' + c.set('k', True, ttl=1) + clock.now += 5 + assert c.get('k') is None + assert c.peek('k', False) is True + + +def test_falsy_values_are_cached(clock: _Clock) -> None: + c = TTLCache(clock=clock) + c.set('k', False, ttl=10) + assert c.get('k', default='miss') is False + + +def test_clear_pattern_on_instance() -> None: + c = TTLCache() + c.set('fonts_catalog', 1) + c.set('system_status', 2) + c.clear('fonts') + assert c.peek('fonts_catalog') is None + assert c.get('system_status') == 2 + c.clear() + assert c.peek('system_status') is None + + +def test_concurrent_expiry_reads_and_writes_do_not_raise() -> None: + """The old dicts deleted expired keys inside get; two threads reading the + same expired key (or one reading while another invalidated) could raise + KeyError, which the endpoints turned into a 500.""" + c = TTLCache() + keys = [f'k{n}' for n in range(8)] + errors = [] + stop = threading.Event() + + def reader() -> None: + try: + while not stop.is_set(): + for key in keys: + c.get(key, max_age=0) # always expired for this reader + c.get(key) + c.peek(key) + c.clear('k1') + except Exception as exc: # pragma: no cover - the failure being tested + errors.append(exc) + + def writer() -> None: + try: + for i in range(20000): + key = keys[i % len(keys)] + c.set(key, i, ttl=0 if i % 2 else 60) + if i % 7 == 0: + c.delete(key) + except Exception as exc: # pragma: no cover + errors.append(exc) + + # Switch threads as often as possible so an unlocked check-then-act + # actually gets interleaved within the test's run time. + old_interval = sys.getswitchinterval() + sys.setswitchinterval(1e-6) + try: + readers = [threading.Thread(target=reader) for _ in range(4)] + for t in readers: + t.start() + writer() + stop.set() + for t in readers: + t.join() + finally: + sys.setswitchinterval(old_interval) + assert errors == [] diff --git a/web_interface/app.py b/web_interface/app.py index f1230ca5..886e5220 100644 --- a/web_interface/app.py +++ b/web_interface/app.py @@ -295,34 +295,41 @@ try: except ImportError: pass -# Cached AP mode check — avoids creating a WiFiManager per request -_ap_mode_cache = {'value': False, 'timestamp': 0} +# systemctl answers, memoised so they are not a subprocess fork per request +# (AP mode) or per SSE tick (display service). A failed check keeps the last +# known answer for the same TTL rather than retrying on every request. +from web_interface.cache import TTLCache +_service_status_cache = TTLCache() _AP_MODE_CACHE_TTL = 30 # seconds — AP mode is user-initiated; 30s is fine - -# Cached ledmatrix service status for SSE stats stream -_ledmatrix_service_cache = {'active': False, 'timestamp': 0} _LEDMATRIX_SERVICE_CACHE_TTL = 15 # seconds +def _unit_is_active(unit, ttl): + """`systemctl is-active `, cached for ``ttl`` seconds. + + False where there is no systemctl (a dev machine); on a failed check, the + last known answer. + """ + active = _service_status_cache.get(unit) + if active is not None: + return active + active = _service_status_cache.peek(unit, False) + if _SYSTEMCTL: + try: + result = subprocess.run([_SYSTEMCTL, 'is-active', unit], + capture_output=True, text=True, timeout=2) + active = result.stdout.strip() == 'active' + except (subprocess.SubprocessError, OSError) as e: + logging.getLogger('web_interface').warning( + "systemctl is-active %s failed: %s", unit, e) + _service_status_cache.set(unit, active, ttl=ttl) + return active + def is_ap_mode_active(): """ Check if access point mode is currently active (cached, 30s TTL). Uses a direct systemctl check instead of instantiating WiFiManager. """ - now = time.time() - if (now - _ap_mode_cache['timestamp']) < _AP_MODE_CACHE_TTL: - return _ap_mode_cache['value'] - try: - result = subprocess.run( - ['systemctl', 'is-active', 'hostapd'], - capture_output=True, text=True, timeout=2 - ) - active = result.stdout.strip() == 'active' - _ap_mode_cache['value'] = active - _ap_mode_cache['timestamp'] = now - return active - except (subprocess.SubprocessError, OSError) as e: - logging.getLogger('web_interface').error(f"AP mode check failed: {e}") - return _ap_mode_cache['value'] + return _unit_is_active('hostapd', _AP_MODE_CACHE_TTL) # Captive portal detection endpoints # When AP mode is active, return responses that TRIGGER the captive portal popup. @@ -672,17 +679,7 @@ def system_status_generator(): cpu_temp = metrics['cpu_temp'] # Check if display service is running (cached to avoid per-client subprocess forks) - now = time.time() - if (now - _ledmatrix_service_cache['timestamp']) >= _LEDMATRIX_SERVICE_CACHE_TTL: - if _SYSTEMCTL: - try: - result = subprocess.run([_SYSTEMCTL, 'is-active', 'ledmatrix'], - capture_output=True, text=True, timeout=2) - _ledmatrix_service_cache['active'] = result.stdout.strip() == 'active' - except (subprocess.SubprocessError, OSError) as e: - app.logger.warning("systemctl status check failed: %s", e) - _ledmatrix_service_cache['timestamp'] = now - service_active = _ledmatrix_service_cache['active'] + service_active = _unit_is_active('ledmatrix', _LEDMATRIX_SERVICE_CACHE_TTL) status = { 'timestamp': time.time(), diff --git a/web_interface/cache.py b/web_interface/cache.py index f7aad3ea..977600a2 100644 --- a/web_interface/cache.py +++ b/web_interface/cache.py @@ -1,48 +1,107 @@ """ -Simple in-memory cache for expensive operations. -Separated from app.py to avoid circular import issues. +In-process TTL cache for the web interface. + +The one place the web process memoises cheap-to-recompute values for a few +seconds or minutes (the font catalog, the system-status snapshot, systemctl +checks). It is per-process and in-memory only; data shared with the display +service goes through ``src.cache_manager.CacheManager`` instead. + +Separated from app.py to avoid circular imports: blueprints import the +module-level helpers below lazily, inside their request handlers. """ +import threading import time -from typing import Any, Optional +from typing import Any, Callable, Dict, Optional, Tuple -# Simple in-memory cache for expensive operations -_cache = {} -_cache_timestamps = {} +class TTLCache: + """A small thread-safe key/value store whose entries expire. + + Each entry keeps the TTL it was stored with. A reader may additionally + pass ``max_age`` to ask for something fresher than that; an entry is only + returned while it is younger than both. + + Expired entries are not dropped on read: :meth:`peek` still returns them, + which is what a "keep the last known answer if the refresh fails" caller + needs. They are replaced by the next :meth:`set` of the same key, so this + is meant for a small, fixed set of keys, not an unbounded key space. + + Ages are measured with ``time.monotonic`` so a wall-clock jump (NTP sync + on a Pi that booted without an RTC) neither expires nor immortalises + everything at once. + """ + + def __init__(self, default_ttl: float = 60, + clock: Callable[[], float] = time.monotonic): + self._default_ttl = default_ttl + self._clock = clock + self._lock = threading.Lock() + # key -> (value, stored_at, ttl) + self._entries: Dict[str, Tuple[Any, float, float]] = {} + + def get(self, key: str, default: Any = None, + max_age: Optional[float] = None) -> Any: + """The value for ``key`` if it is still fresh, else ``default``.""" + with self._lock: + entry = self._entries.get(key) + if entry is None: + return default + value, stored_at, ttl = entry + age = self._clock() - stored_at + if age >= ttl or (max_age is not None and age >= max_age): + return default + return value + + def peek(self, key: str, default: Any = None) -> Any: + """The last value stored for ``key``, fresh or not.""" + with self._lock: + entry = self._entries.get(key) + return default if entry is None else entry[0] + + def set(self, key: str, value: Any, ttl: Optional[float] = None) -> None: + """Store ``value`` for ``ttl`` seconds (the cache default if None).""" + ttl = self._default_ttl if ttl is None else ttl + with self._lock: + self._entries[key] = (value, self._clock(), ttl) + + def delete(self, key: str) -> None: + """Remove ``key`` if present.""" + with self._lock: + self._entries.pop(key, None) + + def clear(self, pattern: Optional[str] = None) -> None: + """Remove every entry, or only those whose key contains ``pattern``.""" + with self._lock: + if pattern is None: + self._entries.clear() + else: + for key in [k for k in self._entries if pattern in k]: + del self._entries[key] -def get_cached(key: str, ttl_seconds: int = 60) -> Optional[Any]: - """Get value from cache if not expired.""" - if key in _cache: - if time.time() - _cache_timestamps[key] < ttl_seconds: - return _cache[key] - else: - # Expired, remove - del _cache[key] - del _cache_timestamps[key] - return None +# The shared cache behind the functional helpers the blueprints use. +_default_cache = TTLCache(default_ttl=60) -def set_cached(key: str, value: Any, ttl_seconds: int = 60) -> None: - """Set value in cache with TTL.""" - _cache[key] = value - _cache_timestamps[key] = time.time() +def get_cached(key: str, ttl_seconds: Optional[float] = None) -> Optional[Any]: + """Get a value from the cache if it has not expired. + + The entry expires after the TTL it was stored with; ``ttl_seconds``, when + given, is an extra upper bound on its age for this read. + """ + return _default_cache.get(key, max_age=ttl_seconds) + + +def set_cached(key: str, value: Any, ttl_seconds: float = 60) -> None: + """Store a value in the cache for ``ttl_seconds``.""" + _default_cache.set(key, value, ttl=ttl_seconds) def delete_cached(key: str) -> None: """Remove a single key from the cache if present.""" - _cache.pop(key, None) - _cache_timestamps.pop(key, None) + _default_cache.delete(key) def invalidate_cache(pattern: Optional[str] = None) -> None: """Invalidate cache entries matching pattern, or all if pattern is None.""" - if pattern is None: - _cache.clear() - _cache_timestamps.clear() - else: - keys_to_remove = [k for k in _cache.keys() if pattern in k] - for key in keys_to_remove: - del _cache[key] - del _cache_timestamps[key] - + _default_cache.clear(pattern)