Extract _stagger_start helper into BasePoller, add tests
The jitter-before-first-sleep block was duplicated verbatim in PeriodicPoller.run() and FSPoller._sweep_loop(). Move it to BasePoller._stagger_start() with a named _STAGGER_FRACTION constant. This also gives a testable seam outside the pragma-no-cover event-loop glue methods.
This commit is contained in:
parent
a6b3e7f1f6
commit
004d59563c
4 changed files with 44 additions and 23 deletions
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import logging
|
||||
import random
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from haiku.rag.config import SourceConfig
|
||||
|
|
@ -15,6 +16,8 @@ from haiku.rag.telemetry import get_context, logfire
|
|||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_STAGGER_FRACTION = 0.25
|
||||
|
||||
|
||||
def _enqueue_extra() -> dict | None:
|
||||
"""Per-job context the worker can't reconstruct from config alone.
|
||||
|
|
@ -84,6 +87,17 @@ class BasePoller:
|
|||
await asyncio.gather(self._task, return_exceptions=True)
|
||||
self._task = None
|
||||
|
||||
async def _stagger_start(self) -> bool:
|
||||
"""Sleep a random fraction of the interval so pollers sharing an
|
||||
interval don't sweep in lockstep. Returns True if stop was
|
||||
signalled during the wait."""
|
||||
jitter = random.uniform(0, self.config.poll_interval_s * _STAGGER_FRACTION)
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), timeout=jitter)
|
||||
return True
|
||||
except TimeoutError:
|
||||
return False
|
||||
|
||||
async def _sweep_once(self) -> bool:
|
||||
"""One discover() sweep. Returns True on success, False if the
|
||||
breaker is open, the source has pending work already queued, or the
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import asyncio
|
||||
import logging
|
||||
import random
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
|
|
@ -65,19 +64,12 @@ class FSPoller(BasePoller):
|
|||
"""Periodic full sweep. Catches files modified while the watcher
|
||||
wasn't running (gaps between starts, races, FS events the OS dropped).
|
||||
Sweep behaviour is unit-tested via `_sweep_once()` directly."""
|
||||
# Stagger the first sleep so multiple FS pollers don't all sweep
|
||||
# at exactly the same moment after startup.
|
||||
interval = self.config.poll_interval_s
|
||||
jitter = random.uniform(0, interval * 0.25)
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), timeout=jitter)
|
||||
if await self._stagger_start():
|
||||
return
|
||||
except TimeoutError:
|
||||
pass
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._stop.wait(), timeout=interval
|
||||
self._stop.wait(), timeout=self.config.poll_interval_s
|
||||
)
|
||||
return
|
||||
except TimeoutError:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,4 @@
|
|||
import asyncio
|
||||
import logging
|
||||
import random
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from haiku.rag.ingester.pollers.base import BasePoller
|
||||
|
|
@ -8,8 +6,6 @@ from haiku.rag.ingester.pollers.base import BasePoller
|
|||
if TYPE_CHECKING:
|
||||
from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PeriodicPoller(BasePoller):
|
||||
"""Runs `source.discover()` on a fixed interval. Used for HTTP, S3, WebDAV
|
||||
|
|
@ -39,19 +35,12 @@ class PeriodicPoller(BasePoller):
|
|||
# immediately instead of waiting one full interval. The sweep
|
||||
# behaviour itself is exercised via `_sweep_once()` unit tests.
|
||||
await self._sweep_once()
|
||||
# Stagger the first sleep so pollers that share the same interval
|
||||
# don't all wake up and sweep at exactly the same moment.
|
||||
interval = self.config.poll_interval_s
|
||||
jitter = random.uniform(0, interval * 0.25)
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), timeout=jitter)
|
||||
if await self._stagger_start():
|
||||
return
|
||||
except TimeoutError:
|
||||
pass
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._stop.wait(), timeout=interval
|
||||
self._stop.wait(), timeout=self.config.poll_interval_s
|
||||
)
|
||||
return
|
||||
except TimeoutError:
|
||||
|
|
|
|||
|
|
@ -114,6 +114,32 @@ def _periodic(source, config, jobs, sync, **kwargs):
|
|||
)
|
||||
|
||||
|
||||
# --- _stagger_start ---
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stagger_start_sleeps_fraction_of_interval(jobs, sync, fs_config, monkeypatch):
|
||||
"""_stagger_start should sleep for a random fraction of poll_interval_s
|
||||
and return False (not stopped)."""
|
||||
monkeypatch.setattr("random.uniform", lambda a, b: b) # max jitter
|
||||
source = _StubSource("src", [])
|
||||
poller = _periodic(source, fs_config, jobs, sync)
|
||||
# poll_interval_s=0.05, so max jitter = 0.05 * 0.25 = 0.0125s
|
||||
stopped = await poller._stagger_start()
|
||||
assert stopped is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stagger_start_returns_true_when_stopped(jobs, sync, fs_config, monkeypatch):
|
||||
"""If _stop is set before the jitter elapses, _stagger_start returns True."""
|
||||
monkeypatch.setattr("random.uniform", lambda a, b: 10.0) # long jitter
|
||||
source = _StubSource("src", [])
|
||||
poller = _periodic(source, fs_config, jobs, sync)
|
||||
poller._stop.set()
|
||||
stopped = await poller._stagger_start()
|
||||
assert stopped is True
|
||||
|
||||
|
||||
# --- _sweep_once / event handling on the base class via PeriodicPoller ---
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue