Add metadata-provider discovery for the ingester

This commit is contained in:
Yiorgis Gozadinos 2026-06-15 08:05:45 +03:00
parent 61f27085f9
commit 5722260857
No known key found for this signature in database
2 changed files with 74 additions and 0 deletions

View file

@ -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)}

View file

@ -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"}