diff --git a/rf4_research/community_cli.py b/rf4_research/community_cli.py index ac99d50..e8e7b4e 100644 --- a/rf4_research/community_cli.py +++ b/rf4_research/community_cli.py @@ -87,19 +87,61 @@ def fetch_site_key(url: str) -> str: def enforce_fetch_interval( source: str, *, state_file: Path, now: float | None = None, ) -> None: - now = time.time() if now is None else now - state = _read_state(state_file) - last_fetch = state.get(source) - if isinstance(last_fetch, (int, float)) and now - last_fetch < MIN_FETCH_INTERVAL_SECONDS: - wait = int(MIN_FETCH_INTERVAL_SECONDS - (now - last_fetch)) - raise RuntimeError(f"source cooldown is active; retry in {wait} seconds") + """Legacy: use check_and_reserve instead for atomic check-and-reserve.""" + check_and_reserve(source, state_file=state_file, now=now) def mark_fetch(source: str, *, state_file: Path, now: float | None = None) -> None: - now = time.time() if now is None else now - state = _read_state(state_file) + """Legacy: use check_and_reserve instead for atomic check-and-reserve.""" + check_and_reserve(source, state_file=state_file, now=now) + + +def check_and_reserve( + source: str, *, state_file: Path, now: float | None = None, +) -> None: + """A02: Atomic check-and-reserve under a single exclusive lock. + + Opens state file with exclusive lock, reads state, checks cooldown, + reserves if allowed — all in one critical section. Uses lockfile + pattern for cross-process coordination. + """ + if now is None: + now = time.time() + state_file.parent.mkdir(parents=True, exist_ok=True) + lock_file = state_file.with_suffix(".lock") + # Create lock file if not exists + lock_file.touch(exist_ok=True) + with open(lock_file, "w") as lf: + fcntl.flock(lf, fcntl.LOCK_EX) + try: + # Read state under lock + try: + with open(state_file, "r") as sf: + state = json.loads(sf.read()) or {} + except (FileNotFoundError, json.JSONDecodeError, ValueError, OSError): + state = {} + # Check cooldown + last_fetch = state.get(source) + if isinstance(last_fetch, (int, float)) and now - last_fetch < MIN_FETCH_INTERVAL_SECONDS: + wait = int(MIN_FETCH_INTERVAL_SECONDS - (now - last_fetch)) + raise RuntimeError(f"source cooldown is active; retry in {wait} seconds") + # Reserve + state[source] = now + # Write atomically via temp file + temp_file = state_file.with_suffix(".tmp") + with open(temp_file, "w") as sf: + sf.write(json.dumps(state, sort_keys=True)) + sf.flush() + os.fsync(sf.fileno()) + temp_file.replace(state_file) + finally: + fcntl.flock(lf, fcntl.LOCK_UN) + + +def _mark_fetch_only(source: str, state: dict, now: float) -> dict: + """Internal: update state without cooldown check (for internal use).""" state[source] = now - _write_state(state_file, state) + return state MAX_REDIRECT_HOPS = 5 diff --git a/tests/test_community_cli.py b/tests/test_community_cli.py index 400e95b..65ba805 100644 --- a/tests/test_community_cli.py +++ b/tests/test_community_cli.py @@ -8,13 +8,19 @@ from rf4_research.community_cli import enforce_fetch_interval, fetch_site_key, m def test_fetch_cooldown_is_persistent_per_source(tmp_path: Path) -> None: - state_file = tmp_path / "fetch-state.json" - mark_fetch("rf4map-point", state_file=state_file, now=1_000) + """A02: check_and_reserve is atomic — one reservation per interval per source.""" + from rf4_research.community_cli import check_and_reserve + state_file = tmp_path / "fetch-state.json" + check_and_reserve("rf4map-point", state_file=state_file, now=1_000) + + # Second call for same source within interval should fail with pytest.raises(RuntimeError, match="retry in 1800 seconds"): - enforce_fetch_interval("rf4map-point", state_file=state_file, now=1_000) - enforce_fetch_interval("rf4posts-spot", state_file=state_file, now=1_000) - enforce_fetch_interval("rf4map-point", state_file=state_file, now=2_800) + check_and_reserve("rf4map-point", state_file=state_file, now=1_000) + # Different source should succeed + check_and_reserve("rf4posts-spot", state_file=state_file, now=1_000) + # After interval expires, should succeed again + check_and_reserve("rf4map-point", state_file=state_file, now=2_800) def test_fetch_site_key_groups_endpoints_and_normalizes_www() -> None: @@ -77,3 +83,61 @@ def test_write_state_is_atomic_with_flush(tmp_path: Path) -> None: state = _read_state(state_file) assert state == {"key1": "value2", "key2": "value3"} assert not (state_file.with_suffix(".tmp")).exists() + + +def test_check_and_reserve_allows_only_one_per_interval(tmp_path: Path) -> None: + """A02: Atomic check-and-reserve — only one success per interval.""" + from rf4_research.community_cli import check_and_reserve + + state_file = tmp_path / "state.json" + base_time = 1000.0 + + # First call should succeed + check_and_reserve("test-source", state_file=state_file, now=base_time) + state = json.loads(state_file.read_text(encoding="utf-8")) + assert state["test-source"] == base_time + + # Second call within interval should fail + with pytest.raises(RuntimeError, match="retry in 1799"): + check_and_reserve("test-source", state_file=state_file, now=base_time + 1) + + # After interval expires, should succeed again + check_and_reserve("test-source", state_file=state_file, now=base_time + 1800) + + +def _try_reserve_for_test(args: tuple) -> tuple: + """Helper for multiprocessing — must be at module level.""" + pid, state_file_str, results_list = args + from rf4_research.community_cli import check_and_reserve + from pathlib import Path + try: + check_and_reserve("shared-source", state_file=Path(state_file_str), now=1000.0) + results_list.append((pid, "ok")) + except RuntimeError as e: + results_list.append((pid, str(e))) + return (pid, "ok") + + +def test_check_and_reserve_atomic_under_concurrent_access(tmp_path: Path) -> None: + """A02: Real multi-process test — concurrent processes get at most one reservation.""" + import multiprocessing + + state_file = tmp_path / "concurrent.json" + results = multiprocessing.Manager().list() + + # Launch 3 processes simultaneously + processes = [] + for i in range(3): + args = (i, str(state_file), results) + p = multiprocessing.Process(target=_try_reserve_for_test, args=(args,)) + processes.append(p) + for p in processes: + p.start() + for p in processes: + p.join(timeout=10) + + # At most one should succeed + ok_count = sum(1 for _, r in results if r == "ok") + assert ok_count == 1, f"Expected exactly 1 ok, got {ok_count}: {results}" + denied_count = sum(1 for _, r in results if "cooldown" in r) + assert denied_count == 2, f"Expected 2 denied, got {denied_count}: {results}"