From f89cc998eb88709fd4103a10e2a3a25c57b23c1c Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Tue, 26 May 2026 10:26:24 +0300 Subject: [PATCH] add sources_provider to WorkerPool. workers now resolve extra info through these sources --- haiku_rag_slim/haiku/rag/client/__init__.py | 3 + haiku_rag_slim/haiku/rag/client/documents.py | 8 +- haiku_rag_slim/haiku/rag/ingester/app.py | 20 ++-- .../haiku/rag/ingester/pollers/base.py | 26 ++--- .../haiku/rag/ingester/pollers/fs.py | 4 +- .../haiku/rag/ingester/pollers/manager.py | 76 ++++++++------- .../haiku/rag/ingester/workers/pipeline.py | 23 +++-- .../haiku/rag/ingester/workers/pool.py | 5 +- tests/ingester/test_api.py | 2 - tests/ingester/test_pipeline.py | 19 ++-- tests/ingester/test_pollers.py | 97 +++++++++++-------- tests/ingester/test_serve_integration.py | 51 +++++++++- tests/ingester/test_workers.py | 21 ++++ 13 files changed, 230 insertions(+), 125 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/client/__init__.py b/haiku_rag_slim/haiku/rag/client/__init__.py index cba36edc..4cadb1ee 100644 --- a/haiku_rag_slim/haiku/rag/client/__init__.py +++ b/haiku_rag_slim/haiku/rag/client/__init__.py @@ -30,6 +30,7 @@ if TYPE_CHECKING: from docling_core.types.doc.document import DoclingDocument from PIL import Image as PILImage + from haiku.rag.ingester.sources.base import Source from haiku.rag.sandbox import AnalysisResult from haiku.rag.store.models.citation import Citation @@ -214,6 +215,7 @@ class HaikuRAG: metadata: dict | None = None, uri: str | None = None, storage_options: dict[str, str] | None = None, + sources: "list[Source] | None" = None, ) -> Document | list[Document]: from haiku.rag.client.documents import create_document_from_source @@ -224,6 +226,7 @@ class HaikuRAG: metadata, uri=uri, storage_options=storage_options, + sources=sources, ) async def update_document( diff --git a/haiku_rag_slim/haiku/rag/client/documents.py b/haiku_rag_slim/haiku/rag/client/documents.py index efc48165..c51face9 100644 --- a/haiku_rag_slim/haiku/rag/client/documents.py +++ b/haiku_rag_slim/haiku/rag/client/documents.py @@ -20,6 +20,7 @@ if TYPE_CHECKING: from docling_core.types.doc.document import DoclingDocument from haiku.rag.client import HaikuRAG + from haiku.rag.ingester.sources.base import Source async def _store_document_with_chunks( @@ -306,6 +307,7 @@ async def create_document_from_source( metadata: dict | None = None, uri: str | None = None, storage_options: dict[str, str] | None = None, + sources: "list[Source] | None" = None, ) -> Document | list[Document]: """Create or update document(s) from a file path, directory, or URL. @@ -370,7 +372,11 @@ async def create_document_from_source( ) # Single resource — resolve the right Source adapter for this URI. - fetcher = resolve_fetcher(source_str, storage_options=storage_options) + # `sources` (configured, in-order) wins over scheme-based adhoc adapters + # so worker fetches reuse the authenticated source the poller used. + fetcher = resolve_fetcher( + source_str, sources=sources, storage_options=storage_options + ) # The stored URI is what we look up + persist by. For an explicit uri # override, use it as-is. For a file:// input the source string is diff --git a/haiku_rag_slim/haiku/rag/ingester/app.py b/haiku_rag_slim/haiku/rag/ingester/app.py index 62e31e91..be62684b 100644 --- a/haiku_rag_slim/haiku/rag/ingester/app.py +++ b/haiku_rag_slim/haiku/rag/ingester/app.py @@ -61,6 +61,13 @@ class IngesterApp: self._db_path, config=self._config, create=True ) as client: self._client = client + self._pollers = PollerManager( + configs=ingester_cfg.sources, + job_repo=self._jobs, + sync_repo=self._sync, + supported_extensions=supported_extensions, + default_max_attempts=ingester_cfg.workers.retry.max_attempts, + ) self._pool = WorkerPool( client=client, job_repo=self._jobs, @@ -71,13 +78,10 @@ class IngesterApp: poll_idle_interval_s=ingester_cfg.workers.poll_idle_interval_s, claim_timeout_s=ingester_cfg.workers.claim_timeout_s, reaper_interval_s=ingester_cfg.workers.reaper_interval_s, - ) - self._pollers = PollerManager( - configs=ingester_cfg.sources, - job_repo=self._jobs, - sync_repo=self._sync, - supported_extensions=supported_extensions, - default_max_attempts=ingester_cfg.workers.retry.max_attempts, + # Same Source instances the pollers discover with — + # workers resolve URIs through them so authenticated + # HTTP / WebDAV / S3 fetches reuse credentials. + sources=self._pollers.sources, ) stop_event = asyncio.Event() @@ -89,8 +93,8 @@ class IngesterApp: # Windows; signal handlers unavailable in asyncio. pass - await self._pool.start() await self._pollers.start() + await self._pool.start() # Log the docling-serve fleet size when relevant so the # operator can eyeball the worker/instance ratio. The convert # phase is usually the throughput ceiling. diff --git a/haiku_rag_slim/haiku/rag/ingester/pollers/base.py b/haiku_rag_slim/haiku/rag/ingester/pollers/base.py index f9348d4d..1ee107e7 100644 --- a/haiku_rag_slim/haiku/rag/ingester/pollers/base.py +++ b/haiku_rag_slim/haiku/rag/ingester/pollers/base.py @@ -16,22 +16,14 @@ from haiku.rag.telemetry import get_context, logfire logger = logging.getLogger(__name__) -def _enqueue_extra(cfg: SourceConfig) -> dict | None: - """Per-source state worth carrying into the job (so the worker can rebuild - the same fetch context when it processes), plus the current logfire trace - context so the worker's `ingester.job` span nests under the - `ingester.poller.sweep` that enqueued it.""" - extra: dict = {} - storage_options = getattr(cfg, "storage_options", None) - if storage_options: - extra["storage_options"] = dict(storage_options) - headers = getattr(cfg, "headers", None) - if headers: - extra["headers"] = dict(headers) +def _enqueue_extra() -> dict | None: + """Per-job context the worker can't reconstruct from config alone. + Currently only the active logfire trace carrier so `ingester.job` + nests under the sweep/watch span that enqueued it. Connection details + (headers, auth, storage_options) come from the configured Source + instance the worker resolves at run time.""" carrier = get_context() - if carrier: - extra["_otel"] = dict(carrier) - return extra or None + return {"_otel": dict(carrier)} if carrier else None def _max_attempts(cfg: SourceConfig, default: int) -> int: @@ -162,7 +154,7 @@ class BasePoller: op=JobOp.UPSERT, revision=event.revision, max_attempts=_max_attempts(self.config, self._default_max_attempts), - extra=_enqueue_extra(self.config), + extra=_enqueue_extra(), ) # Don't write revision to sync_state here — the worker writes it # after a successful ingestion. last_seen_at gets bumped to keep @@ -189,5 +181,5 @@ class BasePoller: event.uri, op=JobOp.DELETE, max_attempts=_max_attempts(self.config, self._default_max_attempts), - extra=_enqueue_extra(self.config), + extra=_enqueue_extra(), ) diff --git a/haiku_rag_slim/haiku/rag/ingester/pollers/fs.py b/haiku_rag_slim/haiku/rag/ingester/pollers/fs.py index b01a50df..09a6ec72 100644 --- a/haiku_rag_slim/haiku/rag/ingester/pollers/fs.py +++ b/haiku_rag_slim/haiku/rag/ingester/pollers/fs.py @@ -118,7 +118,7 @@ class FSPoller(BasePoller): uri, op=JobOp.DELETE, max_attempts=self._max_attempts(), - extra=_enqueue_extra(self._fs_config), + extra=_enqueue_extra(), ) return @@ -130,7 +130,7 @@ class FSPoller(BasePoller): op=JobOp.UPSERT, revision=revision, max_attempts=self._max_attempts(), - extra=_enqueue_extra(self._fs_config), + extra=_enqueue_extra(), ) await self._sync.upsert( self.source_id, uri, revision=None, content_hash=None diff --git a/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py b/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py index d058eb71..c44793a4 100644 --- a/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py +++ b/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py @@ -11,6 +11,7 @@ from haiku.rag.ingester.pollers.periodic import PeriodicPoller if TYPE_CHECKING: from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo + from haiku.rag.ingester.sources.base import Source logger = logging.getLogger(__name__) @@ -29,51 +30,51 @@ class PollerManager: supported_extensions: list[str] | None = None, default_max_attempts: int = 5, ): - self._configs = configs self._jobs = job_repo self._sync = sync_repo self._supported_extensions = supported_extensions self._default_max_attempts = default_max_attempts - self._pollers: list[BasePoller] = [] + # Build eagerly so `sources` is available before `start()` — any + # downstream component that holds the configured Source list (e.g. + # WorkerPool) can do so via plain construction order. + self._pollers: list[BasePoller] = [self._build_poller(cfg) for cfg in configs] self._tasks: list[asyncio.Task] = [] + self._started = False - def build_pollers(self) -> list[BasePoller]: - pollers: list[BasePoller] = [] - for cfg in self._configs: - source = build_source(cfg, supported_extensions=self._supported_extensions) - breaker = CircuitBreaker(cfg.circuit_breaker) - if isinstance(cfg, FSSourceConfig): - from haiku.rag.ingester.sources.fs import FSSource + def _build_poller(self, cfg: SourceConfig) -> BasePoller: + source = build_source(cfg, supported_extensions=self._supported_extensions) + breaker = CircuitBreaker(cfg.circuit_breaker) + if isinstance(cfg, FSSourceConfig): + from haiku.rag.ingester.sources.fs import FSSource - assert isinstance(source, FSSource) - pollers.append( - FSPoller( - source=source, - config=cfg, - job_repo=self._jobs, - sync_repo=self._sync, - breaker=breaker, - default_max_attempts=self._default_max_attempts, - ) - ) - else: - pollers.append( - PeriodicPoller( - source=source, - config=cfg, - job_repo=self._jobs, - sync_repo=self._sync, - breaker=breaker, - default_max_attempts=self._default_max_attempts, - ) - ) - return pollers + assert isinstance(source, FSSource) + return FSPoller( + source=source, + config=cfg, + job_repo=self._jobs, + sync_repo=self._sync, + breaker=breaker, + default_max_attempts=self._default_max_attempts, + ) + return PeriodicPoller( + source=source, + config=cfg, + job_repo=self._jobs, + sync_repo=self._sync, + breaker=breaker, + default_max_attempts=self._default_max_attempts, + ) async def start(self) -> None: - if self._pollers: + if self._started: raise RuntimeError("PollerManager already started") - self._pollers = self.build_pollers() + self._started = True for poller in self._pollers: + # Reset the stop signal synchronously *before* scheduling the + # task. If a poller is being restarted (stop() set the event + # on the previous cycle) clearing inside run() would race with + # any concurrent stop() and could deadlock. + poller._stop.clear() self._tasks.append(asyncio.create_task(poller.run())) async def stop(self) -> None: @@ -82,12 +83,17 @@ class PollerManager: if self._tasks: await asyncio.gather(*self._tasks, return_exceptions=True) self._tasks.clear() - self._pollers.clear() + self._started = False @property def pollers(self) -> list[BasePoller]: return list(self._pollers) + @property + def sources(self) -> list["Source"]: + """Configured Source adapters, one per poller, in config order.""" + return [p.source for p in self._pollers] + @property def live_pollers(self) -> int: """Poller tasks that are still running. Equal to len(pollers) under diff --git a/haiku_rag_slim/haiku/rag/ingester/workers/pipeline.py b/haiku_rag_slim/haiku/rag/ingester/workers/pipeline.py index fade3726..c4b566d2 100644 --- a/haiku_rag_slim/haiku/rag/ingester/workers/pipeline.py +++ b/haiku_rag_slim/haiku/rag/ingester/workers/pipeline.py @@ -12,6 +12,7 @@ from haiku.rag.telemetry import attach_context, logfire if TYPE_CHECKING: from haiku.rag.client import HaikuRAG + from haiku.rag.ingester.sources.base import Source class JobResult(BaseModel): @@ -66,14 +67,19 @@ def _classify(exc: BaseException) -> Exception: return TransientError(f"unexpected: {exc!r}") -async def run_job(client: "HaikuRAG", job: Job) -> JobResult: - """Execute the work described by `job`. Raises PermanentError or - TransientError; the worker uses that to decide dead vs retry.""" +async def run_job( + client: "HaikuRAG", + job: Job, + *, + sources: list["Source"] | None = None, +) -> JobResult: + """Execute the work described by `job`. `sources` is the list of + configured Source adapters — `resolve_fetcher` prefers them over + URI-scheme adhoc adapters so workers reuse the same authenticated / + pre-configured fetch context the pollers used at discovery. Raises + PermanentError or TransientError; the worker uses that to decide + dead vs retry.""" extra = job.extra or {} - storage_options = extra.get("storage_options") - user_metadata = extra.get("metadata", {}) - # Restore the poller's trace context (if any) so the job span nests - # under the `ingester.poller.sweep` that enqueued it. parent_ctx = extra.get("_otel") attach = attach_context(parent_ctx) if parent_ctx else nullcontext() @@ -97,8 +103,7 @@ async def run_job(client: "HaikuRAG", job: Job) -> JobResult: result = await client.create_document_from_source( job.uri, - metadata=user_metadata, - storage_options=storage_options, + sources=sources, ) # Directory ingestion returns list[Document] — workers ingest single # resources, so a list here is a programming error in the caller. diff --git a/haiku_rag_slim/haiku/rag/ingester/workers/pool.py b/haiku_rag_slim/haiku/rag/ingester/workers/pool.py index a8dbf475..709c7898 100644 --- a/haiku_rag_slim/haiku/rag/ingester/workers/pool.py +++ b/haiku_rag_slim/haiku/rag/ingester/workers/pool.py @@ -11,6 +11,7 @@ from haiku.rag.ingester.workers.retry import RetryPolicy, compute_backoff if TYPE_CHECKING: from haiku.rag.client import HaikuRAG + from haiku.rag.ingester.sources.base import Source logger = logging.getLogger(__name__) @@ -36,6 +37,7 @@ class WorkerPool: poll_idle_interval_s: float = 1.0, claim_timeout_s: int = 1800, reaper_interval_s: int = 60, + sources: "list[Source] | None" = None, ): self._client = client self._jobs = job_repo @@ -46,6 +48,7 @@ class WorkerPool: self._poll_idle_s = poll_idle_interval_s self._claim_timeout_s = claim_timeout_s self._reaper_interval_s = reaper_interval_s + self._sources: list[Source] = list(sources) if sources else [] self._stop = asyncio.Event() self._workers: list[asyncio.Task] = [] self._reaper: asyncio.Task | None = None @@ -123,7 +126,7 @@ class WorkerPool: started = time.monotonic() logger.info("Processing %s %s (job %s)", job.op.value, job.uri, job.id) try: - result = await run_job(self._client, job) + result = await run_job(self._client, job, sources=self._sources) except asyncio.CancelledError: # Graceful shutdown cancelled us mid-flight. Release the claim so # the next process can pick the job up immediately instead of diff --git a/tests/ingester/test_api.py b/tests/ingester/test_api.py index 974fa88f..2a39f531 100644 --- a/tests/ingester/test_api.py +++ b/tests/ingester/test_api.py @@ -319,8 +319,6 @@ def _build_pollers_state(tmp_path, jobs, sync, source_id: str = "local"): cfg = FSSourceConfig(type="fs", id=source_id, root=tmp_path) manager = PollerManager(configs=[cfg], job_repo=jobs, sync_repo=sync) - # Build pollers without starting tasks — we want to inspect/refresh directly. - manager._pollers = manager.build_pollers() state = APIState( config=AppConfig(), job_repo=jobs, diff --git a/tests/ingester/test_pipeline.py b/tests/ingester/test_pipeline.py index 17180384..be16c61a 100644 --- a/tests/ingester/test_pipeline.py +++ b/tests/ingester/test_pipeline.py @@ -51,31 +51,32 @@ async def test_upsert_calls_create_document_from_source_and_returns_metadata(): }, ) - result = await run_job( - client, _job(extra={"metadata": {"k": "v"}, "storage_options": {"o": "1"}}) - ) + result = await run_job(client, _job()) assert result.document_id == "doc-42" assert result.revision == "xyz" assert result.content_hash == "abcd" assert result.deleted is False client.create_document_from_source.assert_awaited_once_with( - "https://example.com/a.pdf", - metadata={"k": "v"}, - storage_options={"o": "1"}, + "https://example.com/a.pdf", sources=None ) @pytest.mark.asyncio -async def test_upsert_without_extra_passes_empty_metadata(): +async def test_upsert_threads_configured_sources_to_client(): + """The list of configured Source adapters reaches the client so + resolve_fetcher can pick the authenticated one over an adhoc adapter.""" + from haiku.rag.ingester.sources.http import HTTPSource + client = _mock_client() client.create_document_from_source.return_value = Document( id="d", content="x", uri="u", metadata={} ) + configured = HTTPSource(source_id="urls", headers={"Authorization": "Bearer abc"}) - await run_job(client, _job()) + await run_job(client, _job(), sources=[configured]) client.create_document_from_source.assert_awaited_once_with( - "https://example.com/a.pdf", metadata={}, storage_options=None + "https://example.com/a.pdf", sources=[configured] ) diff --git a/tests/ingester/test_pollers.py b/tests/ingester/test_pollers.py index e4744813..9c4534c7 100644 --- a/tests/ingester/test_pollers.py +++ b/tests/ingester/test_pollers.py @@ -300,45 +300,6 @@ async def test_per_source_retry_policy_overrides_default(jobs, sync, tmp_path): assert queued[0].max_attempts == 9 -@pytest.mark.asyncio -async def test_storage_options_thread_through_to_job_extra(jobs, sync): - cfg = S3SourceConfig( - type="s3", - id="bucket", - uri="s3://bucket/", - storage_options={"endpoint": "http://seaweed:8333"}, - ) - source = _StubSource( - "bucket", [[_event("s3://bucket/file.md", source_id="bucket")]] - ) - poller = _periodic(source, cfg, jobs, sync) - await poller._sweep_once() - queued = await jobs.list_jobs(source_id="bucket") - # _otel is also threaded into extra so the worker's `ingester.job` span - # can nest under the sweep that enqueued it; assert the source-specific - # keys we care about and ignore the trace context payload. - assert queued[0].extra is not None - assert queued[0].extra["storage_options"] == {"endpoint": "http://seaweed:8333"} - - -@pytest.mark.asyncio -async def test_http_headers_thread_through_to_job_extra(jobs, sync): - cfg = HTTPSourceConfig( - type="http", - id="auth", - urls=["https://example.com/a.md"], - headers={"Authorization": "Bearer abc"}, - ) - source = _StubSource( - "auth", [[_event("https://example.com/a.md", source_id="auth")]] - ) - poller = _periodic(source, cfg, jobs, sync) - await poller._sweep_once() - queued = await jobs.list_jobs(source_id="auth") - assert queued[0].extra is not None - assert queued[0].extra["headers"] == {"Authorization": "Bearer abc"} - - # --- PollerManager lifecycle --- @@ -359,7 +320,7 @@ async def test_manager_builds_pollers_per_source(tmp_path, jobs, sync): job_repo=jobs, sync_repo=sync, ) - built = manager.build_pollers() + built = manager.pollers assert len(built) == 4 assert {p.source_id for p in built} == { f"fs:{tmp_path.resolve()}", @@ -369,6 +330,31 @@ async def test_manager_builds_pollers_per_source(tmp_path, jobs, sync): } +@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( @@ -386,6 +372,37 @@ async def test_manager_double_start_raises(tmp_path, jobs, sync): 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 end-to-end smoke --- diff --git a/tests/ingester/test_serve_integration.py b/tests/ingester/test_serve_integration.py index c1716501..e4a10ba7 100644 --- a/tests/ingester/test_serve_integration.py +++ b/tests/ingester/test_serve_integration.py @@ -7,11 +7,12 @@ import aiosqlite import pytest from haiku.rag.client import HaikuRAG -from haiku.rag.config import FSSourceConfig +from haiku.rag.config import FSSourceConfig, HTTPSourceConfig from haiku.rag.ingester.pollers.manager import PollerManager from haiku.rag.ingester.queue.migrations import apply_migrations from haiku.rag.ingester.queue.models import JobOp from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo +from haiku.rag.ingester.sources.http import HTTPSource from haiku.rag.ingester.workers.pool import WorkerPool from haiku.rag.store.models.document import Document @@ -237,3 +238,51 @@ async def test_e2e_watchfiles_push_event_lands_as_job(tmp_path, jobs, sync): assert len(queued) == 1 assert queued[0].op is JobOp.UPSERT assert queued[0].uri == (tmp_path / "new.md").as_uri() + + +@pytest.mark.asyncio +async def test_pre_existing_job_resolves_through_configured_source( + tmp_path, jobs, sync +): + """A job already in the queue at startup is processed through the + configured Source adapter (with its headers / auth), not an adhoc + HTTPSource. The Source list is built at PollerManager construction so + the worker holds it before any start() call — no ordering required.""" + await jobs.enqueue("auth", "https://example.com/a.md", JobOp.UPSERT) + + client = _mock_client(tmp_path) + cfg = HTTPSourceConfig( + type="http", + id="auth", + urls=["https://example.com/a.md"], + headers={"Authorization": "Bearer secret"}, + ) + manager = PollerManager(configs=[cfg], job_repo=jobs, sync_repo=sync) + pool = WorkerPool( + client=client, + job_repo=jobs, + sync_repo=sync, + worker_count=1, + max_concurrent=1, + poll_idle_interval_s=0.05, + sources=manager.sources, + ) + + await manager.start() + await pool.start() + try: + + async def _one_succeeded() -> bool: + counts = await jobs.counts_by_status() + return counts.get("succeeded", 0) == 1 + + await _wait_for(_one_succeeded, timeout=5.0) + finally: + await pool.stop() + await manager.stop() + + kwargs = client.create_document_from_source.await_args.kwargs + sources = kwargs.get("sources") + assert sources is not None and len(sources) == 1 + assert isinstance(sources[0], HTTPSource) + assert sources[0].headers == {"Authorization": "Bearer secret"} diff --git a/tests/ingester/test_workers.py b/tests/ingester/test_workers.py index 37cf8045..4f229af1 100644 --- a/tests/ingester/test_workers.py +++ b/tests/ingester/test_workers.py @@ -50,6 +50,7 @@ def _pool(client, jobs, sync, **kwargs) -> WorkerPool: reaper_interval_s=kwargs.pop("reaper_interval_s", 60), claim_timeout_s=kwargs.pop("claim_timeout_s", 60), retry_policy=kwargs.pop("retry_policy", RetryPolicy()), + sources=kwargs.pop("sources", None), ) @@ -195,6 +196,26 @@ async def test_keyboard_interrupt_propagates_not_classified(client, jobs, sync): assert refreshed.status is JobStatus.CLAIMED +@pytest.mark.asyncio +async def test_drain_passes_configured_sources_to_client(client, jobs, sync): + """The pool's `sources` list flows through run_job to + client.create_document_from_source so resolve_fetcher can pick the + configured authenticated source over an adhoc adapter.""" + from haiku.rag.ingester.sources.http import HTTPSource + + client.create_document_from_source.return_value = Document( + id="d", content="x", uri="u", metadata={"md5": "m", "source_revision": "r"} + ) + configured = HTTPSource(source_id="urls", headers={"Authorization": "Bearer abc"}) + await jobs.enqueue("src", "https://example.com/x", JobOp.UPSERT) + + pool = _pool(client, jobs, sync, sources=[configured]) + await pool.drain_once() + + kwargs = client.create_document_from_source.await_args.kwargs + assert kwargs["sources"] == [configured] + + # --- start / stop lifecycle ---