Add metadata-provider discovery for the ingester
This commit is contained in:
parent
61f27085f9
commit
5722260857
2 changed files with 74 additions and 0 deletions
23
haiku_rag_slim/haiku/rag/ingester/metadata.py
Normal file
23
haiku_rag_slim/haiku/rag/ingester/metadata.py
Normal 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)}
|
||||
51
tests/ingester/test_metadata.py
Normal file
51
tests/ingester/test_metadata.py
Normal 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"}
|
||||
Loading…
Reference in a new issue