From fd4db45ae10e619a4d3be0e7a89fb47c39a5ee3e Mon Sep 17 00:00:00 2001 From: xavierk Date: Mon, 28 Sep 2026 15:02:58 +0530 Subject: [PATCH] fix: bound pending publication admission (#101) --- src/fenris/collector.py | 186 ++++++++++++++++++++--------- tests/test_issue_101.py | 256 ++++++++++++++++++++++++++++++++++++++++ tests/test_issue_97.py | 1 + 3 files changed, 384 insertions(+), 59 deletions(-) create mode 100644 tests/test_issue_101.py diff --git a/src/fenris/collector.py b/src/fenris/collector.py index 460b5d5..99ec28e 100644 --- a/src/fenris/collector.py +++ b/src/fenris/collector.py @@ -32,6 +32,26 @@ class InvariantViolationError(Exception): pass +class _PendingDerivationError(Exception): + """A valid staged observation failed while deriving dependent evidence.""" + + def __init__(self, cause: Exception): + super().__init__(str(cause)) + self.cause = cause + + +class _PendingRecoveryFailure(Exception): + """A publication attempt failed after earlier pending rows committed.""" + + def __init__(self, error: Exception, published_count: int): + derivation_error = isinstance(error, _PendingDerivationError) + cause = error.cause if derivation_error else error + super().__init__(str(cause)) + self.cause = cause + self.published_count = published_count + self.derivation_error = derivation_error + + PENDING_PUBLICATION_LIMIT = 6720 @@ -300,47 +320,50 @@ def _publish_observation( current_id = seg_info["sample_id"] prev = find_previous_sample(conn, seg_info.get("segment_id"), current_id) - if prev is not None: - current = { - "id": current_id, - "ts": sample["ts"], - "bytes_written": sample["bytes_written"], - "bytes_read": sample["bytes_read"], - "power_on_hours": sample["power_on_hours"], - "temperature_c": sample["temperature_c"], - "data_units_written": sample["data_units_written"], - "data_units_read": sample["data_units_read"], - "local_tz": tz_name, - } - derive_hours_from_interval(conn, prev, current) - from .local_day import record_local_activity_interval - record_local_activity_interval( - conn, - prev, - current, - start_sample_id=prev["id"], - end_sample_id=current_id, - segment_id=seg_info.get("segment_id"), - ) - else: - previous_any_segment = find_previous_sample(conn, None, current_id) - if previous_any_segment is not None: - from .local_day import mark_local_activity_gap - mark_local_activity_gap( + try: + if prev is not None: + current = { + "id": current_id, + "ts": sample["ts"], + "bytes_written": sample["bytes_written"], + "bytes_read": sample["bytes_read"], + "power_on_hours": sample["power_on_hours"], + "temperature_c": sample["temperature_c"], + "data_units_written": sample["data_units_written"], + "data_units_read": sample["data_units_read"], + "local_tz": tz_name, + } + derive_hours_from_interval(conn, prev, current) + from .local_day import record_local_activity_interval + record_local_activity_interval( conn, - sample["ts"], - tz_name, - previous_any_segment, + prev, + current, + start_sample_id=prev["id"], + end_sample_id=current_id, + segment_id=seg_info.get("segment_id"), ) + else: + previous_any_segment = find_previous_sample(conn, None, current_id) + if previous_any_segment is not None: + from .local_day import mark_local_activity_gap + mark_local_activity_gap( + conn, + sample["ts"], + tz_name, + previous_any_segment, + ) - from .day_aggregate import derive_all_days, persist_day_aggregate - for aggregate in derive_all_days(conn): - persist_day_aggregate(conn, aggregate) + from .day_aggregate import derive_all_days, persist_day_aggregate + for aggregate in derive_all_days(conn): + persist_day_aggregate(conn, aggregate) - from .local_day import derive_local_day_summary, persist_local_day - local_summary = derive_local_day_summary(conn, tz_name, observed_at) - if local_summary is not None: - persist_local_day(conn, local_summary) + from .local_day import derive_local_day_summary, persist_local_day + local_summary = derive_local_day_summary(conn, tz_name, observed_at) + if local_summary is not None: + persist_local_day(conn, local_summary) + except Exception as exc: # noqa: BLE001 - preserve valid staged evidence for retry. + raise _PendingDerivationError(exc) from exc def _pending_count(conn: sqlite3.Connection) -> int: @@ -381,12 +404,15 @@ def _recover_pending(conn: sqlite3.Connection) -> int: pending_id, payload = row observation = json.loads(payload) - _publish_observation( - conn, - observation["sample"], - observation["identity"], - observation["tz_name"], - ) + try: + _publish_observation( + conn, + observation["sample"], + observation["identity"], + observation["tz_name"], + ) + except Exception as exc: + raise _PendingRecoveryFailure(exc, recovered) from exc conn.execute("DELETE FROM pending_publications WHERE id = ?", (pending_id,)) conn.commit() recovered += 1 @@ -395,26 +421,61 @@ def _recover_pending(conn: sqlite3.Connection) -> int: raise +def _is_store_or_invariant_failure(exc: Exception) -> bool: + return ( + isinstance(exc, sqlite3.Error) + and not isinstance(exc, sqlite3.IntegrityError) + ) or isinstance(exc, InvariantViolationError) + + +def _pending_capacity_error(reason: str) -> RuntimeError: + return RuntimeError( + "pending publication capacity full " + f"({PENDING_PUBLICATION_LIMIT} observations); no new observation acquired; " + f"{reason}" + ) + + def _recover_pending_for_collection(conn: sqlite3.Connection) -> int: """Retry queued work and report exhausted capacity without masking store faults.""" try: return _recover_pending(conn) - except Exception as exc: - if isinstance(exc, sqlite3.Error) and not isinstance(exc, sqlite3.IntegrityError): - raise - if isinstance(exc, InvariantViolationError): - raise + except Exception as exc: # noqa: BLE001 - preserve pending evidence on recovery errors. + cause = exc.cause if isinstance(exc, _PendingRecoveryFailure) else exc + if _is_store_or_invariant_failure(cause): + raise cause try: capacity_full = _pending_count(conn) >= PENDING_PUBLICATION_LIMIT except sqlite3.Error: - raise exc + raise cause if capacity_full: - raise RuntimeError( - "pending publication capacity full " - f"({PENDING_PUBLICATION_LIMIT} observations); no new observation acquired; " - f"recovery failed: {exc}" - ) from exc - raise + raise _pending_capacity_error(f"recovery failed: {cause}") from cause + raise cause + + +def _recover_pending_for_admission(conn: sqlite3.Connection) -> int: + """Recover in order, then admit one observation if bounded space remains.""" + try: + recovered = _recover_pending(conn) + except _PendingRecoveryFailure as failure: + if ( + not failure.derivation_error + or _is_store_or_invariant_failure(failure.cause) + ): + raise failure.cause + try: + pending_count = _pending_count(conn) + except sqlite3.Error: + raise failure.cause + if pending_count >= PENDING_PUBLICATION_LIMIT: + raise _pending_capacity_error( + f"recovery failed: {failure.cause}" + ) from failure.cause + return failure.published_count + + if _pending_count(conn) >= PENDING_PUBLICATION_LIMIT: + raise _pending_capacity_error("recovery left the queue full") + return recovered def run_collection( @@ -444,16 +505,23 @@ def run_collection( if history_path.exists(): import_legacy_history(conn, history_path, clock=clock) - published_count = _recover_pending_for_collection(conn) + published_count = _recover_pending_for_admission(conn) + checked_pending_count = _pending_count(conn) # Hold the writer reservation across the capacity check and acquisition. - # A concurrent collector will recheck pending work before it acquires. + # Recheck recovery if another collector appended work during preflight. while True: conn.execute("BEGIN IMMEDIATE") waiting = _pending_count(conn) - if waiting: + if waiting >= PENDING_PUBLICATION_LIMIT: conn.rollback() - published_count += _recover_pending_for_collection(conn) + published_count += _recover_pending_for_admission(conn) + checked_pending_count = _pending_count(conn) + continue + if waiting > checked_pending_count: + conn.rollback() + published_count += _recover_pending_for_admission(conn) + checked_pending_count = _pending_count(conn) continue from .tz_util import detect_system_tz diff --git a/tests/test_issue_101.py b/tests/test_issue_101.py new file mode 100644 index 0000000..9078792 --- /dev/null +++ b/tests/test_issue_101.py @@ -0,0 +1,256 @@ +"""Bound pending-publication admission before device acquisition.""" +import json +import sqlite3 +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).parent.parent / "src")) + +from fenris import collector + + +class FakeClock: + def __init__(self, now): + self.now = now + + def utcnow(self): + return self.now + + +@pytest.fixture +def sysfs_controller(tmp_path): + controller_path = tmp_path / "sys" / "class" / "nvme" / "nvme0" + controller_path.mkdir(parents=True) + (controller_path / "subsysnqn").write_text("nqn.test:drive\n") + (controller_path / "model").write_text("Test NVMe\n") + (controller_path / "serial").write_text("test-serial\n") + (controller_path / "firmware_rev").write_text("1.0\n") + transport = controller_path / "transport" + transport.mkdir() + (transport / "trstring").write_text("pcie\n") + return controller_path + + +def _smartctl(units_written): + return { + "nvme_smart_health_information_log": { + "critical_warning": 0, + "temperature": 35, + "available_spare": 100, + "percentage_used": 5, + "data_units_written": units_written, + "data_units_read": 500, + "power_on_hours": 100, + }, + "user_capacity": {"bytes": 1_024_000_000_000}, + "model_name": "Test NVMe", + "serial_number": "test-serial", + "firmware_version": "1.0", + } + + +def _collect(config, controller, now, units_written): + return collector.run_collection( + config=config, + clock=FakeClock(now), + acquire=lambda: (_smartctl(units_written), controller), + ) + + +def test_partial_recovery_frees_admission_and_keeps_pending_order( + tmp_path, sysfs_controller, monkeypatch, +): + """A freed slot admits one observation behind older pending evidence.""" + assert collector.PENDING_PUBLICATION_LIMIT == 6_720 + monkeypatch.setenv("TZ", "UTC") + monkeypatch.setattr(collector, "PENDING_PUBLICATION_LIMIT", 2) + + store_path = tmp_path / "observations.db" + config = {"device": "/dev/nvme0", "store_path": str(store_path)} + old_time = datetime(2026, 8, 1, tzinfo=timezone.utc) + controller = sysfs_controller + + assert _collect(config, controller, old_time, 1000)["ok"] is True + + original_derive = collector.derive_hours_from_interval + + def unavailable_derivation(*_args): + raise RuntimeError("derivation unavailable") + + monkeypatch.setattr(collector, "derive_hours_from_interval", unavailable_derivation) + first_pending = _collect( + config, controller, old_time + timedelta(minutes=5), 1010, + ) + assert first_pending["ok"] is False + + second_pending = _collect( + config, controller, old_time + timedelta(minutes=10), 1020, + ) + assert second_pending["ok"] is False + with sqlite3.connect(store_path) as conn: + assert conn.execute( + "SELECT COUNT(*) FROM pending_publications" + ).fetchone()[0] == 2 + + called = False + + def must_not_acquire(): + nonlocal called + called = True + return _smartctl(1030), controller + + full = collector.run_collection( + config=config, + clock=FakeClock(old_time + timedelta(days=40)), + acquire=must_not_acquire, + ) + assert full["ok"] is False + assert "capacity full (2 observations)" in full["error"] + assert called is False + with sqlite3.connect(store_path) as conn: + assert conn.execute( + "SELECT COUNT(*), MIN(sample_ts) FROM pending_publications" + ).fetchone() == ( + 2, + (old_time + timedelta(minutes=5)).isoformat(), + ) + assert conn.execute("SELECT COUNT(*) FROM samples").fetchone()[0] == 1 + + def fail_second_pending(conn, previous, current): + if current["data_units_written"] == 1020: + raise RuntimeError("second observation still blocked") + return original_derive(conn, previous, current) + + monkeypatch.setattr(collector, "derive_hours_from_interval", fail_second_pending) + acquired = False + + def acquire_after_partial_recovery(): + nonlocal acquired + acquired = True + return _smartctl(1030), controller + + partial = collector.run_collection( + config=config, + clock=FakeClock(old_time + timedelta(days=40, minutes=5)), + acquire=acquire_after_partial_recovery, + ) + assert partial["ok"] is False + assert acquired is True, partial + + with sqlite3.connect(store_path) as conn: + assert conn.execute( + "SELECT data_units_written FROM samples ORDER BY id" + ).fetchall() == [(1000,), (1010,)] + queued_units = [ + json.loads(row[0])["sample"]["data_units_written"] + for row in conn.execute( + "SELECT payload FROM pending_publications ORDER BY id" + ) + ] + assert queued_units == [1020, 1030] + assert conn.execute("SELECT COUNT(*) FROM monitoring_periods").fetchone()[0] == 1 + + recovered_in_order = [] + + def record_recovery_order(conn, previous, current): + recovered_in_order.append(current["data_units_written"]) + return original_derive(conn, previous, current) + + monkeypatch.setattr(collector, "derive_hours_from_interval", record_recovery_order) + recovered = collector.run_collection( + config=config, + clock=FakeClock(old_time + timedelta(days=40, minutes=10)), + acquire=lambda: (_smartctl(1040), controller), + ) + assert recovered["ok"] is True, recovered + assert recovered_in_order == [1020, 1030, 1040] + + with sqlite3.connect(store_path) as conn: + assert conn.execute( + "SELECT data_units_written FROM samples ORDER BY id DESC LIMIT 1" + ).fetchone() == (1040,) + assert conn.execute("SELECT COUNT(*) FROM pending_publications").fetchone()[0] == 0 + + +def test_actual_pending_limit_fails_before_acquisition_and_keeps_old_evidence( + tmp_path, sysfs_controller, monkeypatch, +): + """The production limit retains every old row and refuses device acquisition.""" + monkeypatch.setenv("TZ", "UTC") + store_path = tmp_path / "full-observations.db" + config = {"device": "/dev/nvme0", "store_path": str(store_path)} + controller = sysfs_controller + old_time = datetime(2026, 7, 1, tzinfo=timezone.utc) + assert _collect(config, controller, old_time, 1000)["ok"] is True + + def unavailable_derivation(*_args): + raise RuntimeError("derivation unavailable") + + monkeypatch.setattr(collector, "derive_hours_from_interval", unavailable_derivation) + failed = _collect( + config, controller, old_time + timedelta(minutes=5), 1010, + ) + assert failed["ok"] is False + + with sqlite3.connect(store_path) as conn: + payload_row = conn.execute( + "SELECT payload FROM pending_publications ORDER BY id LIMIT 1" + ).fetchone() + assert payload_row is not None + base = json.loads(payload_row[0]) + base_observed_at = datetime.fromisoformat(base["sample"]["ts"]) + base_units = base["sample"]["data_units_written"] + rows = [] + for offset in range(1, collector.PENDING_PUBLICATION_LIMIT): + observation = json.loads(payload_row[0]) + units_written = base_units + offset + observed_at = base_observed_at + timedelta(minutes=5 * offset) + observation["sample"]["ts"] = observed_at.isoformat() + observation["sample"]["data_units_written"] = units_written + observation["sample"]["bytes_written"] = units_written * 512_000 + rows.append((observed_at.isoformat(), json.dumps(observation))) + conn.executemany( + "INSERT INTO pending_publications (sample_ts, payload) VALUES (?, ?)", + rows, + ) + conn.commit() + + called = False + + def must_not_acquire(): + nonlocal called + called = True + return _smartctl(9000), controller + + result = collector.run_collection( + config=config, + clock=FakeClock(old_time + timedelta(days=40)), + acquire=must_not_acquire, + ) + assert result["ok"] is False + assert f"capacity full ({collector.PENDING_PUBLICATION_LIMIT} observations)" in result[ + "error" + ] + assert called is False + + with sqlite3.connect(store_path) as conn: + pending_rows = conn.execute( + "SELECT sample_ts, payload FROM pending_publications ORDER BY id" + ).fetchall() + assert len(pending_rows) == collector.PENDING_PUBLICATION_LIMIT + assert (pending_rows[0][0], pending_rows[-1][0]) == ( + (old_time + timedelta(minutes=5)).isoformat(), + ( + old_time + + timedelta(minutes=5 * collector.PENDING_PUBLICATION_LIMIT) + ).isoformat(), + ) + assert [ + json.loads(payload)["sample"]["data_units_written"] + for _, payload in pending_rows + ] == list(range(1010, 1010 + collector.PENDING_PUBLICATION_LIMIT)) + assert conn.execute("SELECT COUNT(*) FROM samples").fetchone()[0] == 1 + assert conn.execute("SELECT COUNT(*) FROM monitoring_periods").fetchone()[0] == 1 diff --git a/tests/test_issue_97.py b/tests/test_issue_97.py index f3b6643..785a11e 100644 --- a/tests/test_issue_97.py +++ b/tests/test_issue_97.py @@ -168,6 +168,7 @@ def test_failed_publication_survives_restart_and_readers_keep_last_consistent_vi called = True return _smartctl(1030), sysfs_path + monkeypatch.setattr(collector, "PENDING_PUBLICATION_LIMIT", 1) blocked = collector.run_collection( config=config, clock=FakeClock(now + timedelta(minutes=5)),