Wire metadata providers into ingester ingestion

This commit is contained in:
Yiorgis Gozadinos 2026-06-15 08:21:46 +03:00
parent 5722260857
commit cc73a8629a
No known key found for this signature in database
11 changed files with 312 additions and 17 deletions

View file

@ -1,6 +1,10 @@
# Changelog
## [Unreleased]
### Added
- Ingester metadata providers: a source config's `metadata_provider` names a callable registered under the `haiku.rag.metadata_providers` entry-point group; the ingester calls it per document with `(source_id, uri)` and attaches the returned dict as document metadata. System-derived keys (`md5`, `source_revision`, `content_type`) take precedence on collision.
### Changed
- Mutable document attributes (`uri`, `title`, `metadata`, `created_at`, `updated_at`) moved from the `documents` table into a new `document_meta` table (1:1 on `document_id`); metadata/title/`source_revision` updates no longer rewrite the docling blobs. Migration `v0_58_0` relocates existing data and runs a one-time `vacuum` to reclaim prior bloat.

View file

@ -174,6 +174,51 @@ relies on a `Content-Length` response header; a server that omits it (for
example a chunked response) is fetched in full and the limit does not
apply.
### Metadata providers
A source can attach custom metadata to every document it ingests by
naming a `metadata_provider`. The provider is a callable that an external
package registers under the `haiku.rag.metadata_providers` entry-point
group; the ingester calls it per document with `(source_id, uri)` and
merges the returned dict into the document's metadata.
```yaml
- type: webdav
id: handbook
base_url: https://dav.example.com/remote.php/dav/files/svc
metadata_provider: example-provider
```
The provider is a zero-argument callable returning the provider instance,
so a class is its own factory:
```python
# example_pkg/__init__.py
from urllib.parse import urlparse
class Provider:
async def __call__(self, source_id: str, uri: str) -> dict:
path = urlparse(uri).path
return {
"collection": source_id,
"folder": path.rsplit("/", 1)[0] or "/",
}
```
```toml
# in the provider package's pyproject.toml
[project.entry-points."haiku.rag.metadata_providers"]
example-provider = "example_pkg:Provider"
```
The provider is built once at startup, so it can hold a client or cache
across calls. The source-derived keys (`md5`, `source_revision`,
`content_type`) are stripped from provider output, so a provider cannot
override them. A `metadata_provider` name with no installed entry point
fails at startup. A provider exception is classified like any other
ingestion error (network and timeout errors retry; others go to the DLQ).
## Workers and retry
```yaml

View file

@ -398,6 +398,13 @@ class _SourceBase(BaseModel):
description="Maximum file size in bytes to fetch. Files larger than "
"this are rejected with a PermanentError. None = no limit.",
)
metadata_provider: str | None = Field(
default=None,
description="Name of a metadata provider registered under the "
"'haiku.rag.metadata_providers' entry-point group. When set, the "
"provider is called per document with (source_id, uri) and its result "
"is attached as document metadata. None = no provider.",
)
class FSSourceConfig(_SourceBase):

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING
from pydantic import BaseModel
from haiku.rag.config import AppConfig
from haiku.rag.ingester.metadata import build_providers, load_metadata_providers
from haiku.rag.ingester.pollers.manager import PollerManager
from haiku.rag.ingester.queue.migrations import open_queue
from haiku.rag.ingester.queue.repository import JobRepo, SyncStateRepo
@ -89,6 +90,15 @@ class IngesterApp:
supported_extensions=supported_extensions,
default_max_attempts=ingester_cfg.workers.retry.max_attempts,
)
metadata_providers = build_providers(
[
(source.source_id, cfg.metadata_provider)
for cfg, source in zip(
ingester_cfg.sources, self._pollers.sources
)
],
load_metadata_providers(),
)
self._pool = WorkerPool(
client=client,
job_repo=self._jobs,
@ -107,6 +117,7 @@ class IngesterApp:
# workers resolve URIs through them so authenticated
# HTTP / WebDAV / S3 fetches reuse credentials.
sources=self._pollers.sources,
metadata_providers=metadata_providers,
)
yield
finally:

View file

@ -1,4 +1,4 @@
from collections.abc import Callable
from collections.abc import Callable, Iterable, Mapping
from importlib.metadata import entry_points
from typing import Protocol, runtime_checkable
@ -18,6 +18,42 @@ class MetadataProvider(Protocol):
MetadataProviderFactory = Callable[[], MetadataProvider]
def load_metadata_providers() -> dict[str, MetadataProviderFactory]:
"""Discover registered metadata-provider factories, keyed by entry-point name."""
return {ep.name: ep.load() for ep in entry_points(group=ENTRY_POINT_GROUP)}
@runtime_checkable
class LoadableEntryPoint(Protocol):
"""The slice of ``importlib.metadata.EntryPoint`` ``build_providers`` needs:
a deferred ``load()`` returning the provider factory."""
def load(self) -> MetadataProviderFactory: ...
def load_metadata_providers() -> dict[str, LoadableEntryPoint]:
"""Discover registered metadata-provider entry points, keyed by name. The
entry points are not imported here; ``build_providers`` loads only the ones
a source references, so an unused provider with a broken import does not
fail the ingester at startup."""
return {ep.name: ep for ep in entry_points(group=ENTRY_POINT_GROUP)}
def build_providers(
sources: Iterable[tuple[str, str | None]],
discovered: Mapping[str, LoadableEntryPoint],
) -> dict[str, MetadataProvider]:
"""Load and instantiate the provider named by each ``(source_id, name)``
pair, keyed by source id. Pairs with no name are skipped, and only
referenced entry points are loaded. Raises ValueError if a name has no
registered entry point so a misconfigured source fails at startup rather
than silently dropping metadata."""
providers: dict[str, MetadataProvider] = {}
for source_id, name in sources:
if name is None:
continue
try:
entry_point = discovered[name]
except KeyError:
raise ValueError(
f"Source {source_id!r} references unknown metadata provider "
f"{name!r}; no entry point registered under {ENTRY_POINT_GROUP!r}."
) from None
factory: MetadataProviderFactory = entry_point.load()
providers[source_id] = factory()
return providers

View file

@ -13,10 +13,20 @@ from haiku.rag.ingester.sources.registry import resolve_configured_source
from haiku.rag.telemetry import attach_context, logfire
if TYPE_CHECKING:
from collections.abc import Mapping
from haiku.rag.client import HaikuRAG
from haiku.rag.ingester.metadata import MetadataProvider
from haiku.rag.ingester.sources.base import Source
# Keys the source pipeline owns (content_type/md5/source_revision and the
# source_revision/md5 that drive sync_state). A provider must not set them, or
# the metadata-only refresh path would let provider values overwrite the real
# source-derived ones. Stripped before provider metadata reaches the client.
_RESERVED_METADATA_KEYS = frozenset({"content_type", "md5", "source_revision"})
class JobResult(BaseModel):
"""What the worker needs after a successful job: enough metadata to
update sync_state. document_id is None for DELETE ops."""
@ -82,13 +92,15 @@ async def run_job(
job: Job,
*,
sources: list["Source"] | None = None,
metadata_providers: "Mapping[str, MetadataProvider] | None" = None,
) -> JobResult:
"""Execute the work described by `job`. `sources` is the list of
configured Source adapters; the client looks up `job.source_id`
against it via `resolve_configured_source` so workers reuse the
authenticated/pre-configured fetch context the pollers used at
discovery. Raises PermanentError or TransientError; the worker
uses that to decide dead vs retry."""
discovery. `metadata_providers` maps source_id to a provider whose
output is attached as document metadata on UPSERT. Raises PermanentError
or TransientError; the worker uses that to decide dead vs retry."""
extra = job.extra or {}
parent_ctx = extra.get("_otel")
attach = attach_context(parent_ctx) if parent_ctx else nullcontext()
@ -122,10 +134,19 @@ async def run_job(
await client.delete_document(doc.id)
return JobResult(deleted=True)
provider = (metadata_providers or {}).get(job.source_id)
extra_metadata: dict | None = None
if provider is not None:
extra_metadata = {
k: v
for k, v in (await provider(job.source_id, job.uri)).items()
if k not in _RESERVED_METADATA_KEYS
}
result = await client.create_document_from_source(
job.uri,
sources=sources,
source_id=job.source_id,
metadata=extra_metadata,
)
# Directory ingestion returns list[Document] — workers ingest single
# resources, so a list here is a programming error in the caller.

View file

@ -12,7 +12,10 @@ from haiku.rag.ingester.workers.pipeline import run_job
from haiku.rag.ingester.workers.retry import RetryPolicy, compute_backoff
if TYPE_CHECKING:
from collections.abc import Mapping
from haiku.rag.client import HaikuRAG
from haiku.rag.ingester.metadata import MetadataProvider
from haiku.rag.ingester.sources.base import Source
logger = logging.getLogger(__name__)
@ -41,6 +44,7 @@ class WorkerPool:
reaper_interval_s: int = 60,
retention_s: int | None = None,
sources: "list[Source] | None" = None,
metadata_providers: "Mapping[str, MetadataProvider] | None" = None,
):
self._client = client
self._jobs = job_repo
@ -52,6 +56,9 @@ class WorkerPool:
self._reaper_interval_s = reaper_interval_s
self._retention_s = retention_s
self._sources: list[Source] = list(sources) if sources else []
self._metadata_providers: dict[str, MetadataProvider] = (
dict(metadata_providers) if metadata_providers else {}
)
self._stop = asyncio.Event()
self._workers: list[asyncio.Task] = []
self._reaper: asyncio.Task | None = None
@ -181,7 +188,12 @@ 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, sources=self._sources)
result = await run_job(
self._client,
job,
sources=self._sources,
metadata_providers=self._metadata_providers,
)
except asyncio.CancelledError:
# Graceful shutdown cancelled us mid-flight. Spawn the release
# as an independent Task tracked in _pending_releases — that

View file

@ -59,6 +59,18 @@ def test_discriminator_picks_fs_source():
assert isinstance(cfg.sources[0], FSSourceConfig)
assert cfg.sources[0].root == Path("/tmp/docs")
assert cfg.sources[0].delete_orphans is True
assert cfg.sources[0].metadata_provider is None
def test_source_carries_metadata_provider_name():
cfg = IngesterConfig.model_validate(
{
"sources": [
{"type": "fs", "root": "/tmp/docs", "metadata_provider": "example"}
]
}
)
assert cfg.sources[0].metadata_provider == "example"
def test_discriminator_picks_http_source():

View file

@ -4,35 +4,42 @@ from haiku.rag.ingester import metadata as metadata_module
from haiku.rag.ingester.metadata import (
ENTRY_POINT_GROUP,
MetadataProvider,
build_providers,
load_metadata_providers,
)
class _Provider:
async def __call__(self, source_id: str, uri: str) -> dict:
return {"source": source_id}
class _FakeEntryPoint:
def __init__(self, name, factory):
self.name = name
self._factory = factory
self.loaded = False
def load(self):
self.loaded = True
return self._factory
def test_load_keys_factories_by_entry_point_name(monkeypatch):
def factory():
return None
def test_load_keys_entry_points_by_name_without_loading(monkeypatch):
ep = _FakeEntryPoint("example-provider", _Provider)
captured: dict = {}
def fake_entry_points(*, group):
captured["group"] = group
return [_FakeEntryPoint("example-provider", factory)]
return [ep]
monkeypatch.setattr(metadata_module, "entry_points", fake_entry_points)
providers = load_metadata_providers()
discovered = load_metadata_providers()
assert captured["group"] == ENTRY_POINT_GROUP
assert providers == {"example-provider": factory}
assert discovered == {"example-provider": ep}
assert ep.loaded is False
def test_load_is_empty_when_none_registered(monkeypatch):
@ -49,3 +56,39 @@ async def test_callable_object_satisfies_protocol():
provider = Provider()
assert isinstance(provider, MetadataProvider)
assert await provider("src", "u") == {"classification": "secret"}
def test_build_providers_instantiates_named_factories():
providers = build_providers(
[("docs", "example-provider"), ("wiki", None)],
{"example-provider": _FakeEntryPoint("example-provider", _Provider)},
)
assert set(providers) == {"docs"}
assert isinstance(providers["docs"], _Provider)
def test_build_providers_skips_sources_without_a_provider():
discovered = {"example-provider": _FakeEntryPoint("example-provider", _Provider)}
assert build_providers([("docs", None)], discovered) == {}
def test_build_providers_raises_on_unknown_provider_name():
discovered = {"example-provider": _FakeEntryPoint("example-provider", _Provider)}
with pytest.raises(ValueError, match="unknown metadata provider 'missing'"):
build_providers([("docs", "missing")], discovered)
def test_build_providers_does_not_load_unreferenced_entry_points():
def _explode():
raise ImportError("optional dependency missing")
discovered = {
"used": _FakeEntryPoint("used", _Provider),
"unused": _FakeEntryPoint("unused", _explode),
}
providers = build_providers([("docs", "used")], discovered)
assert isinstance(providers["docs"], _Provider)
assert discovered["unused"].loaded is False

View file

@ -59,7 +59,7 @@ async def test_upsert_calls_create_document_from_source_and_returns_metadata():
assert result.content_hash == "abcd"
assert result.deleted is False
client.create_document_from_source.assert_awaited_once_with(
"https://example.com/a.pdf", sources=None, source_id="src"
"https://example.com/a.pdf", sources=None, source_id="src", metadata=None
)
@ -77,10 +77,114 @@ async def test_upsert_threads_configured_sources_to_client():
await run_job(client, _job(), sources=[configured])
client.create_document_from_source.assert_awaited_once_with(
"https://example.com/a.pdf", sources=[configured], source_id="src"
"https://example.com/a.pdf",
sources=[configured],
source_id="src",
metadata=None,
)
class _MetadataProvider:
"""Provider double returning scripted metadata, or raising."""
def __init__(self, metadata: dict | None = None, *, error: Exception | None = None):
self._metadata = metadata or {}
self._error = error
async def __call__(self, source_id: str, uri: str) -> dict:
if self._error is not None:
raise self._error
return {**self._metadata, "source": source_id}
@pytest.mark.asyncio
async def test_provider_metadata_passed_to_client():
client = _mock_client()
client.create_document_from_source.return_value = Document(
id="d", content="x", uri="u", metadata={}
)
providers = {"src": _MetadataProvider({"classification": "secret"})}
await run_job(client, _job(), metadata_providers=providers)
client.create_document_from_source.assert_awaited_once_with(
"https://example.com/a.pdf",
sources=None,
source_id="src",
metadata={"classification": "secret", "source": "src"},
)
@pytest.mark.asyncio
async def test_provider_cannot_override_system_keys():
"""Reserved source-derived keys are stripped from provider output so the
metadata-only refresh path can't let a provider overwrite md5 /
source_revision / content_type (which would corrupt sync_state)."""
client = _mock_client()
client.create_document_from_source.return_value = Document(
id="d", content="x", uri="u", metadata={}
)
providers = {
"src": _MetadataProvider(
{
"md5": "spoof",
"source_revision": "spoof",
"content_type": "text/spoof",
"classification": "secret",
}
)
}
await run_job(client, _job(), metadata_providers=providers)
client.create_document_from_source.assert_awaited_once_with(
"https://example.com/a.pdf",
sources=None,
source_id="src",
metadata={"classification": "secret", "source": "src"},
)
@pytest.mark.asyncio
async def test_no_provider_for_source_passes_no_metadata():
"""A provider registered for a different source must not apply here."""
client = _mock_client()
client.create_document_from_source.return_value = Document(
id="d", content="x", uri="u", metadata={}
)
providers = {"other": _MetadataProvider({"classification": "secret"})}
await run_job(client, _job(), metadata_providers=providers)
client.create_document_from_source.assert_awaited_once_with(
"https://example.com/a.pdf", sources=None, source_id="src", metadata=None
)
@pytest.mark.asyncio
async def test_provider_error_is_classified_and_blocks_ingest():
client = _mock_client()
providers = {"src": _MetadataProvider(error=httpx.ConnectError("provider down"))}
with pytest.raises(TransientError):
await run_job(client, _job(), metadata_providers=providers)
client.create_document_from_source.assert_not_awaited()
@pytest.mark.asyncio
async def test_provider_not_called_for_delete():
client = _mock_client()
client.get_document_by_uri.return_value = Document(id="doc-9", content="", uri="u")
provider = _MetadataProvider(error=AssertionError("must not run on DELETE"))
result = await run_job(
client, _job(op=JobOp.DELETE), metadata_providers={"src": provider}
)
assert result.deleted is True
@pytest.mark.asyncio
async def test_delete_calls_delete_document_when_present():
client = _mock_client()

View file

@ -607,7 +607,7 @@ async def test_breaker_isolates_sources(client, jobs, sync):
"""An open breaker pauses only the failing source. Workers keep draining
a healthy source's jobs while the failing source's jobs stay queued."""
def _route(uri, *, sources=None, source_id=None):
def _route(uri, *, sources=None, source_id=None, metadata=None):
if source_id == "bad":
raise TransientError("downstream down")
return Document(