diff --git a/haiku_rag_slim/haiku/rag/ingester/metadata.py b/haiku_rag_slim/haiku/rag/ingester/metadata.py new file mode 100644 index 00000000..1a41cc7b --- /dev/null +++ b/haiku_rag_slim/haiku/rag/ingester/metadata.py @@ -0,0 +1,23 @@ +from collections.abc import Callable +from importlib.metadata import entry_points +from typing import Protocol, runtime_checkable + +ENTRY_POINT_GROUP = "haiku.rag.metadata_providers" + + +@runtime_checkable +class MetadataProvider(Protocol): + """Computes per-document metadata for the ingester. A package registers a + zero-arg factory under the ``haiku.rag.metadata_providers`` entry-point + group; the factory returns an instance whose ``__call__`` the ingester + invokes per job with the document's source id and uri.""" + + async def __call__(self, source_id: str, uri: str) -> dict: ... + + +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)} diff --git a/tests/ingester/test_metadata.py b/tests/ingester/test_metadata.py new file mode 100644 index 00000000..e6018683 --- /dev/null +++ b/tests/ingester/test_metadata.py @@ -0,0 +1,51 @@ +import pytest + +from haiku.rag.ingester import metadata as metadata_module +from haiku.rag.ingester.metadata import ( + ENTRY_POINT_GROUP, + MetadataProvider, + load_metadata_providers, +) + + +class _FakeEntryPoint: + def __init__(self, name, factory): + self.name = name + self._factory = factory + + def load(self): + return self._factory + + +def test_load_keys_factories_by_entry_point_name(monkeypatch): + def factory(): + return None + + captured: dict = {} + + def fake_entry_points(*, group): + captured["group"] = group + return [_FakeEntryPoint("example-provider", factory)] + + monkeypatch.setattr(metadata_module, "entry_points", fake_entry_points) + + providers = load_metadata_providers() + + assert captured["group"] == ENTRY_POINT_GROUP + assert providers == {"example-provider": factory} + + +def test_load_is_empty_when_none_registered(monkeypatch): + monkeypatch.setattr(metadata_module, "entry_points", lambda *, group: []) + assert load_metadata_providers() == {} + + +@pytest.mark.asyncio +async def test_callable_object_satisfies_protocol(): + class Provider: + async def __call__(self, source_id: str, uri: str) -> dict: + return {"classification": "secret"} + + provider = Provider() + assert isinstance(provider, MetadataProvider) + assert await provider("src", "u") == {"classification": "secret"}