import asyncio from datetime import UTC, datetime from pathlib import Path import pytest from haiku.rag.config import ( CircuitBreakerConfig, FSSourceConfig, HTTPSourceConfig, S3SourceConfig, WebDAVSourceConfig, ) from haiku.rag.ingester.pollers.circuit_breaker import CircuitBreaker from haiku.rag.ingester.pollers.manager import PollerManager from haiku.rag.ingester.pollers.periodic import PeriodicPoller from haiku.rag.ingester.queue.models import JobOp, JobStatus from haiku.rag.ingester.sources.base import ( FetchResult, SourceEvent, SourceEventKind, ) class _StubSource: """Test double that yields a scripted sequence of events on each discover() call. `fetch` and `supports` aren't exercised by pollers.""" def __init__(self, source_id: str, sweeps: list[list[SourceEvent]]): self.source_id = source_id self._sweeps = list(sweeps) self.discover_calls = 0 self.fail_with: Exception | None = None def supports(self, uri: str) -> bool: # pragma: no cover - unused here return True async def head(self, uri: str) -> str | None: # pragma: no cover return None async def fetch(self, uri: str) -> FetchResult: # pragma: no cover raise NotImplementedError async def discover(self, since=None, *, known_uris=None): self.discover_calls += 1 if self.fail_with is not None: raise self.fail_with events = self._sweeps.pop(0) if self._sweeps else [] for event in events: yield event def _event( uri: str, kind=SourceEventKind.UPSERT, revision: str | None = "v1", source_id: str = "src", ): return SourceEvent( source_id=source_id, uri=uri, kind=kind, revision=None if kind is SourceEventKind.DELETE else revision, discovered_at=datetime.now(UTC), ) @pytest.fixture def fs_config(tmp_path): return FSSourceConfig( type="fs", id="src", root=tmp_path, delete_orphans=True, poll_interval_s=0.05, ) def _periodic(source, config, jobs, sync, **kwargs): return PeriodicPoller( source=source, config=config, job_repo=jobs, sync_repo=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 --- @pytest.mark.asyncio async def test_upsert_event_enqueues_job_and_touches_sync_state(fs_config, jobs, sync): source = _StubSource("src", [[_event("file:///a.md", revision="r1")]]) poller = _periodic(source, fs_config, jobs, sync) ok = await poller._sweep_once() assert ok is True queued = await jobs.list_jobs(source_id="src") assert len(queued) == 1 assert queued[0].op is JobOp.UPSERT assert queued[0].revision == "r1" # Pollers DO NOT write revision to sync_state — the worker does that # after a successful ingest. The URI shows up in list_known_uris from # the moment the poller emits the event, but the revision_snapshot # stays empty until ingestion completes. assert await sync.get_revision_snapshot("src") == {} assert await sync.list_known_uris("src") == {"file:///a.md"} @pytest.mark.asyncio async def test_unchanged_event_touches_sync_state_no_job(fs_config, jobs, sync): source = _StubSource( "src", [[_event("file:///a.md", kind=SourceEventKind.UNCHANGED, revision="r1")]] ) poller = _periodic(source, fs_config, jobs, sync) await poller._sweep_once() assert await jobs.list_jobs(source_id="src") == [] assert await sync.get_revision_snapshot("src") == {"file:///a.md": "r1"} @pytest.mark.asyncio async def test_delete_event_enqueues_delete_job(fs_config, jobs, sync): source = _StubSource( "src", [[_event("file:///gone.md", kind=SourceEventKind.DELETE)]] ) poller = _periodic(source, fs_config, jobs, sync) await poller._sweep_once() queued = await jobs.list_jobs(source_id="src") assert len(queued) == 1 assert queued[0].op is JobOp.DELETE @pytest.mark.asyncio async def test_delete_event_skipped_when_delete_orphans_false(fs_config, jobs, sync): fs_config = fs_config.model_copy(update={"delete_orphans": False}) source = _StubSource( "src", [[_event("file:///gone.md", kind=SourceEventKind.DELETE)]] ) poller = _periodic(source, fs_config, jobs, sync) await poller._sweep_once() assert await jobs.list_jobs(source_id="src") == [] @pytest.mark.asyncio async def test_repeated_sweep_skipped_when_queue_has_pending(fs_config, jobs, sync): """Backpressure: once a job is queued/claimed, the next sweep skips discover() entirely instead of churning the listing operation.""" event = _event("file:///a.md", revision="r1") source = _StubSource("src", [[event], [event]]) poller = _periodic(source, fs_config, jobs, sync) assert await poller._sweep_once() is True assert source.discover_calls == 1 # Second sweep: queue still has the live job → skip without calling discover. assert await poller._sweep_once() is False assert source.discover_calls == 1 @pytest.mark.asyncio async def test_skipped_sweep_records_pending_work_reason(fs_config, jobs, sync): """last_skip_reason surfaces 'pending_work' while the queue is saturated and clears once the next sweep actually polls.""" event = _event("file:///a.md", revision="r1") source = _StubSource("src", [[event], [event], []]) poller = _periodic(source, fs_config, jobs, sync) await poller._sweep_once() # first sweep enqueues, succeeds assert poller.last_skip_reason is None await poller._sweep_once() # backpressure skips assert poller.last_skip_reason == "pending_work" # Drain the queue, sweep again, reason clears. claimed = await jobs.claim_next("worker") assert claimed is not None await jobs.mark_succeeded(claimed.id, "worker") await poller._sweep_once() assert poller.last_skip_reason is None @pytest.mark.asyncio async def test_dead_job_does_not_clear_sync_state_revision(fs_config, jobs, sync): """When a job dies, the URI's previously-ingested revision must remain in sync_state so subsequent sweeps still see the URI as known.""" await sync.upsert("src", "file:///a.md", revision="r1", content_hash="h1") changed = _event("file:///a.md", revision="r2") source = _StubSource("src", [[changed], [changed]]) poller = _periodic(source, fs_config, jobs, sync) await poller._sweep_once() queued = await jobs.list_jobs() assert len(queued) == 1 assert queued[0].revision == "r2" claimed = await jobs.claim_next("w") assert claimed is not None await jobs.mark_dead(claimed.id, "transient blew up", "w") row = await sync.get_row("src", "file:///a.md") assert row is not None assert row.revision == "r1" await poller._sweep_once() assert await sync.get_revision_snapshot("src") == {"file:///a.md": "r1"} @pytest.mark.asyncio async def test_sweep_resumes_after_queue_drains(fs_config, jobs, sync): """Once the queue clears (success, dead, or cancel), sweeps resume.""" event = _event("file:///a.md", revision="r1") source = _StubSource("src", [[event], []]) poller = _periodic(source, fs_config, jobs, sync) await poller._sweep_once() claimed = await jobs.claim_next("worker") assert claimed is not None await jobs.mark_succeeded(claimed.id, "worker") assert await poller._sweep_once() is True assert source.discover_calls == 2 @pytest.mark.asyncio async def test_circuit_breaker_pauses_sweeps_after_failures(fs_config, jobs, sync): class _Clock: now = 0.0 def __call__(self): return self.now clock = _Clock() breaker = CircuitBreaker( CircuitBreakerConfig(failure_threshold=2, cooldown_s=30.0), now_fn=clock, ) source = _StubSource("src", []) source.fail_with = RuntimeError("upstream down") poller = _periodic(source, fs_config, jobs, sync, breaker=breaker) # Two failures open the breaker. assert await poller._sweep_once() is False assert await poller._sweep_once() is False assert breaker.is_open is True # Third call should be skipped — discover() is not invoked. before = source.discover_calls assert await poller._sweep_once() is False assert source.discover_calls == before @pytest.mark.asyncio async def test_sweep_records_last_polled_at_on_success(fs_config, jobs, sync): source = _StubSource("src", [[]]) poller = _periodic(source, fs_config, jobs, sync) assert poller.last_polled_at is None await poller._sweep_once() assert poller.last_polled_at is not None assert poller.last_polled_at.tzinfo is not None @pytest.mark.asyncio async def test_per_source_retry_policy_overrides_default(jobs, sync, tmp_path): from haiku.rag.config import RetryPolicyConfig cfg = FSSourceConfig( type="fs", id="src", root=tmp_path, retry=RetryPolicyConfig(max_attempts=9), ) source = _StubSource("src", [[_event("file:///a.md")]]) poller = _periodic(source, cfg, jobs, sync, default_max_attempts=3) await poller._sweep_once() queued = await jobs.list_jobs(source_id="src") assert queued[0].max_attempts == 9 # --- dry-run collection --- @pytest.mark.asyncio async def test_dry_run_reports_changes_without_mutating_queue_or_sync( fs_config, jobs, sync ): source = _StubSource( "src", [ [ _event("file:///a.md", revision="r1"), _event( "file:///b.md", kind=SourceEventKind.UNCHANGED, revision="r2", ), _event("file:///gone.md", kind=SourceEventKind.DELETE), ] ], ) poller = _periodic(source, fs_config, jobs, sync) ok, summary, changes = await poller._dry_run_once() assert ok is True assert summary.upsert_count == 1 assert summary.delete_count == 1 assert summary.unchanged_count == 1 assert summary.ignored_delete_count == 0 assert [(c.op, c.uri, c.revision) for c in changes] == [ (JobOp.UPSERT, "file:///a.md", "r1"), (JobOp.DELETE, "file:///gone.md", None), ] assert await jobs.list_jobs(source_id="src") == [] assert await sync.list_known_uris("src") == set() @pytest.mark.asyncio async def test_dry_run_counts_ignored_deletes_when_orphan_delete_disabled( fs_config, jobs, sync ): cfg = fs_config.model_copy(update={"delete_orphans": False}) source = _StubSource( "src", [[_event("file:///gone.md", kind=SourceEventKind.DELETE)]] ) poller = _periodic(source, cfg, jobs, sync) ok, summary, changes = await poller._dry_run_once() assert ok is True assert summary.delete_count == 0 assert summary.ignored_delete_count == 1 assert changes == [] assert await jobs.list_jobs(source_id="src") == [] @pytest.mark.asyncio async def test_dry_run_skips_when_queue_has_pending_work(fs_config, jobs, sync): await jobs.enqueue("src", "file:///already.md", op=JobOp.UPSERT) source = _StubSource("src", [[_event("file:///a.md")]]) poller = _periodic(source, fs_config, jobs, sync) ok, summary, changes = await poller._dry_run_once() assert ok is False assert summary.source_id == "src" assert changes == [] assert source.discover_calls == 0 assert poller.last_skip_reason == "pending_work" # --- PollerManager lifecycle --- @pytest.mark.asyncio async def test_manager_builds_pollers_per_source(tmp_path, jobs, sync): """When SourceConfig.id is set, the poller's source uses it verbatim; when omitted, the adapter auto-derives one from its target.""" configs = [ FSSourceConfig(type="fs", root=tmp_path), S3SourceConfig(type="s3", uri="s3://bucket/"), HTTPSourceConfig(type="http", id="urls", urls=[]), WebDAVSourceConfig( type="webdav", id="nc", base_url="https://nc.example.com/dav/" ), ] manager = PollerManager( configs=configs, job_repo=jobs, sync_repo=sync, ) built = manager.pollers assert len(built) == 4 assert {p.source_id for p in built} == { f"fs:{tmp_path.resolve()}", "s3:bucket/", "urls", "nc", } @pytest.mark.asyncio async def test_manager_sources_available_at_construction(tmp_path, jobs, sync): """PollerManager builds Sources eagerly so callers (WorkerPool) can receive them by plain construction order.""" from haiku.rag.config import SourceConfig from haiku.rag.ingester.sources.http import HTTPSource configs: list[SourceConfig] = [ FSSourceConfig(type="fs", id="docs", root=tmp_path), HTTPSourceConfig( type="http", id="urls", urls=[], headers={"Authorization": "Bearer abc"}, ), ] manager = PollerManager(configs=configs, job_repo=jobs, sync_repo=sync) sources = manager.sources assert len(sources) == 2 assert {s.source_id for s in sources} == {"docs", "urls"} http = next(s for s in sources if s.source_id == "urls") assert isinstance(http, HTTPSource) assert http.headers == {"Authorization": "Bearer abc"} @pytest.mark.asyncio async def test_manager_double_start_raises(tmp_path, jobs, sync): cfg = FSSourceConfig( type="fs", id="local", root=tmp_path, poll_interval_s=60.0, ) manager = PollerManager(configs=[cfg], job_repo=jobs, sync_repo=sync) await manager.start() try: with pytest.raises(RuntimeError, match="already started"): await manager.start() finally: await manager.stop() @pytest.mark.asyncio async def test_manager_restart_resumes_polling(tmp_path, jobs, sync): """stop() then start() produces a working poller that stays alive after its initial sweep — the second cycle's stop event is fresh, not the set state left over from the previous stop().""" (tmp_path / "a.md").write_text("hello") cfg = FSSourceConfig( type="fs", id="local", root=tmp_path, poll_interval_s=60.0, ) manager = PollerManager( configs=[cfg], job_repo=jobs, sync_repo=sync, supported_extensions=[".md"] ) await manager.start() await asyncio.sleep(0.1) await manager.stop() await manager.start() try: # The poller task must STAY alive after its initial sweep so the # watchfiles + periodic-sweep loops keep running. live_pollers # drops to 0 immediately if run() exited because _stop was set. await asyncio.sleep(0.1) assert manager.live_pollers == 1 finally: await manager.stop() # --- FSPoller._handle_watch_change --- def _fs_poller(tmp_path, jobs, sync): """Construct an FSPoller for unit-testing the watch-change handler. Doesn't start the watch loop — tests call `_handle_watch_change` directly.""" from haiku.rag.ingester.pollers.fs import FSPoller from haiku.rag.ingester.sources.fs import FSSource cfg = FSSourceConfig( type="fs", id="local", root=tmp_path, delete_orphans=True, poll_interval_s=60.0, ) source = FSSource(root=tmp_path, supported_extensions=[".md"], source_id="local") return FSPoller( source=source, config=cfg, job_repo=jobs, sync_repo=sync, ) @pytest.mark.asyncio async def test_watch_deleted_enqueues_when_file_truly_gone(tmp_path, jobs, sync): from watchfiles import Change poller = _fs_poller(tmp_path, jobs, sync) missing = tmp_path / "gone.md" # file doesn't exist; deleted event should enqueue DELETE await poller._handle_watch_change(Change.deleted, missing) queued = await jobs.list_jobs(source_id="local") assert len(queued) == 1 assert queued[0].op is JobOp.DELETE @pytest.mark.asyncio async def test_watch_deleted_skipped_when_file_already_back(tmp_path, jobs, sync): """`git checkout` and atomic-rename saves fire (deleted, added) in quick succession; by the time the deleted event reaches us the file is back. Skipping the DELETE keeps the follow-up Change.added's UPSERT from being blocked by the live-row index.""" from watchfiles import Change poller = _fs_poller(tmp_path, jobs, sync) present = tmp_path / "restored.md" present.write_text("restored") await poller._handle_watch_change(Change.deleted, present) assert await jobs.list_jobs(source_id="local") == [] @pytest.mark.asyncio async def test_watch_deleted_then_added_enqueues_upsert(tmp_path, jobs, sync): """End-to-end of the git-checkout scenario: (deleted, added) for an existing file leaves a single UPSERT job, not a DELETE.""" from watchfiles import Change poller = _fs_poller(tmp_path, jobs, sync) path = tmp_path / "doc.md" path.write_text("contents") await poller._handle_watch_change(Change.deleted, path) await poller._handle_watch_change(Change.added, path) queued = await jobs.list_jobs(source_id="local") assert len(queued) == 1 assert queued[0].op is JobOp.UPSERT @pytest.mark.asyncio async def test_watch_added_file_deleted_before_stat_does_not_crash( tmp_path, jobs, sync ): """If a file is deleted between the watchfiles event and the stat() call, the handler should return silently instead of raising FileNotFoundError and killing the watch loop.""" from watchfiles import Change poller = _fs_poller(tmp_path, jobs, sync) missing = tmp_path / "vanished.md" # File doesn't exist — simulate Change.added arriving after deletion. await poller._handle_watch_change(Change.added, missing) # No job should be enqueued, and no exception should have propagated. assert await jobs.list_jobs(source_id="local") == [] # --- FSPoller end-to-end smoke --- @pytest.mark.asyncio async def test_fs_poller_enqueues_initial_files(tmp_path, jobs, sync): (tmp_path / "a.md").write_text("hello") (tmp_path / "b.md").write_text("world") cfg = FSSourceConfig(type="fs", id="local", root=tmp_path, poll_interval_s=60.0) manager = PollerManager( configs=[cfg], job_repo=jobs, sync_repo=sync, supported_extensions=[".md"] ) await manager.start() try: # Wait for the initial sweep to land jobs. for _ in range(40): queued = await jobs.list_jobs(source_id="local") if len(queued) == 2: break await asyncio.sleep(0.05) finally: await manager.stop() queued = await jobs.list_jobs(source_id="local") assert {Path(j.uri).name for j in queued} == {"a.md", "b.md"} assert all(j.status is JobStatus.QUEUED for j in queued)