Pass fetched FetchResult to ingester metadata providers
This commit is contained in:
parent
3f5141abd0
commit
1e65d36766
14 changed files with 436 additions and 61 deletions
|
|
@ -9,6 +9,7 @@
|
|||
|
||||
- Bump `docling>=2.102.2,<3.0.0` and `docling-core>=2.82.0,<3.0.0`; the `<3.0.0` cap holds the DoclingDocument schema at 1.10.0.
|
||||
- Relax `opencv-python-headless` to `>=4.6.0.66,<5.0.0.0` (was `>=4.13.0.92`) to match `docling-ibm-models`' declared range.
|
||||
- `haiku.rag.metadata_providers` callables take a third argument, the fetched `FetchResult`: `__call__(source_id, uri, result)`. The provider runs after fetch instead of before; on revision-unchanged sweeps it is skipped and existing provider metadata is preserved.
|
||||
|
||||
### Fixed
|
||||
|
||||
|
|
|
|||
|
|
@ -188,8 +188,9 @@ apply.
|
|||
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.
|
||||
group; when the document is fetched for ingestion, the ingester calls it
|
||||
with `(source_id, uri, result)`, where `result` is the source's
|
||||
`FetchResult`, and merges the returned dict into the document's metadata.
|
||||
|
||||
```yaml
|
||||
- type: webdav
|
||||
|
|
@ -205,13 +206,18 @@ so a class is its own factory:
|
|||
# example_pkg/__init__.py
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from haiku.rag.ingester.sources import FetchResult
|
||||
|
||||
|
||||
class Provider:
|
||||
async def __call__(self, source_id: str, uri: str) -> dict:
|
||||
async def __call__(
|
||||
self, source_id: str, uri: str, result: FetchResult
|
||||
) -> dict:
|
||||
path = urlparse(uri).path
|
||||
return {
|
||||
"collection": source_id,
|
||||
"folder": path.rsplit("/", 1)[0] or "/",
|
||||
"bytes": str(len(result.body)),
|
||||
}
|
||||
```
|
||||
|
||||
|
|
@ -222,11 +228,15 @@ 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).
|
||||
across calls. When a document's source revision is unchanged, the
|
||||
ingester keeps the existing cheap HEAD short-circuit and preserves the
|
||||
stored provider metadata; the provider runs again when the document is
|
||||
fetched for a new or changed revision. 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).
|
||||
|
||||
### Custom sources
|
||||
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ if TYPE_CHECKING:
|
|||
from PIL import Image as PILImage
|
||||
|
||||
from haiku.rag.embeddings import EmbedderWrapper
|
||||
from haiku.rag.ingester.metadata import MetadataProvider
|
||||
from haiku.rag.ingester.sources.base import Source
|
||||
from haiku.rag.reranking.base import RerankerBase
|
||||
from haiku.rag.sandbox import AnalysisResult
|
||||
|
|
@ -279,6 +280,7 @@ class HaikuRAG:
|
|||
storage_options: dict[str, str] | None = None,
|
||||
sources: "list[Source] | None" = None,
|
||||
source_id: str | None = None,
|
||||
metadata_provider: "MetadataProvider | None" = None,
|
||||
) -> Document | list[Document]:
|
||||
from haiku.rag.client.documents import create_document_from_source
|
||||
|
||||
|
|
@ -291,6 +293,7 @@ class HaikuRAG:
|
|||
storage_options=storage_options,
|
||||
sources=sources,
|
||||
source_id=source_id,
|
||||
metadata_provider=metadata_provider,
|
||||
)
|
||||
|
||||
async def update_document(
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ if TYPE_CHECKING:
|
|||
from docling_core.types.doc.document import DoclingDocument
|
||||
|
||||
from haiku.rag.client import HaikuRAG
|
||||
from haiku.rag.ingester.metadata import MetadataProvider
|
||||
from haiku.rag.ingester.sources.base import FetchResult, Source
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
|
@ -51,6 +52,12 @@ class DocumentImport:
|
|||
# and is skipped.
|
||||
MAX_ATTACHMENT_DEPTH = 3
|
||||
|
||||
# Keys the source pipeline owns (content_type/md5/source_revision, which 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 is merged into the document.
|
||||
_RESERVED_METADATA_KEYS = frozenset({"content_type", "md5", "source_revision"})
|
||||
|
||||
|
||||
def parent_uri_filter(parent_uri: str) -> str:
|
||||
"""SQL `WHERE` clause matching documents whose ``metadata.parent_uri``
|
||||
|
|
@ -332,6 +339,26 @@ async def _refresh_doc_metadata(
|
|||
return doc
|
||||
|
||||
|
||||
async def _provider_metadata(
|
||||
provider: "MetadataProvider | None",
|
||||
source_id: str,
|
||||
uri: str,
|
||||
result: "FetchResult",
|
||||
) -> dict:
|
||||
if provider is None:
|
||||
return {}
|
||||
# Hand the provider an isolated copy: mutating the live FetchResult
|
||||
# (e.g. result.content_hash or result.extra_metadata) would feed the
|
||||
# MD5 short-circuit and source_meta, bypassing the reserved-key filter
|
||||
# that only guards the returned dict.
|
||||
provider_result = result.model_copy(deep=True)
|
||||
return {
|
||||
k: v
|
||||
for k, v in (await provider(source_id, uri, provider_result)).items()
|
||||
if k not in _RESERVED_METADATA_KEYS
|
||||
}
|
||||
|
||||
|
||||
async def _ingest_fetch_result(
|
||||
client: "HaikuRAG",
|
||||
result: "FetchResult",
|
||||
|
|
@ -575,6 +602,7 @@ async def create_document_from_source(
|
|||
storage_options: dict[str, str] | None = None,
|
||||
sources: "list[Source] | None" = None,
|
||||
source_id: str | None = None,
|
||||
metadata_provider: "MetadataProvider | None" = None,
|
||||
) -> Document | list[Document]:
|
||||
"""Create or update document(s) from a file path, directory, or URL.
|
||||
|
||||
|
|
@ -621,7 +649,13 @@ async def create_document_from_source(
|
|||
for child in local_path.rglob("*"):
|
||||
if child.is_file() and filter.include_file(str(child)):
|
||||
doc = await create_document_from_source(
|
||||
client, child, title=None, metadata=metadata
|
||||
client,
|
||||
child,
|
||||
title=None,
|
||||
metadata=metadata,
|
||||
sources=sources,
|
||||
source_id=source_id,
|
||||
metadata_provider=metadata_provider,
|
||||
)
|
||||
assert isinstance(doc, Document)
|
||||
documents.append(doc)
|
||||
|
|
@ -692,6 +726,11 @@ async def create_document_from_source(
|
|||
fetch_span.set_attribute("bytes", len(result.body))
|
||||
fetch_span.set_attribute("content_hash", result.content_hash)
|
||||
|
||||
provider_metadata = await _provider_metadata(
|
||||
metadata_provider, source_id or fetcher.source_id, source_str, result
|
||||
)
|
||||
user_metadata = {**metadata, **provider_metadata}
|
||||
|
||||
# MD5 short-circuit: the bytes are unchanged even if the revision wasn't.
|
||||
# Refresh the source-derived metadata (revision may have rolled) but skip
|
||||
# convert/embed/store entirely.
|
||||
|
|
@ -707,7 +746,7 @@ async def create_document_from_source(
|
|||
client,
|
||||
existing_doc,
|
||||
title=title,
|
||||
user_metadata=metadata,
|
||||
user_metadata=user_metadata,
|
||||
source_metadata=source_meta,
|
||||
)
|
||||
|
||||
|
|
@ -715,7 +754,7 @@ async def create_document_from_source(
|
|||
client,
|
||||
result,
|
||||
title=title,
|
||||
user_metadata=metadata,
|
||||
user_metadata=user_metadata,
|
||||
stored_uri=stored_uri,
|
||||
existing_doc=existing_doc,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
from collections.abc import Callable, Iterable, Mapping
|
||||
from importlib.metadata import entry_points
|
||||
from typing import Protocol, runtime_checkable
|
||||
from typing import TYPE_CHECKING, Protocol, runtime_checkable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from haiku.rag.ingester.sources.base import FetchResult
|
||||
|
||||
ENTRY_POINT_GROUP = "haiku.rag.metadata_providers"
|
||||
|
||||
|
|
@ -10,9 +13,11 @@ 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."""
|
||||
invokes per job with the document's source id, uri, and fetched content."""
|
||||
|
||||
async def __call__(self, source_id: str, uri: str) -> dict: ...
|
||||
async def __call__(
|
||||
self, source_id: str, uri: str, result: "FetchResult"
|
||||
) -> dict: ...
|
||||
|
||||
|
||||
MetadataProviderFactory = Callable[[], MetadataProvider]
|
||||
|
|
|
|||
|
|
@ -20,13 +20,6 @@ if TYPE_CHECKING:
|
|||
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."""
|
||||
|
|
@ -134,19 +127,11 @@ 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,
|
||||
metadata_provider=(metadata_providers or {}).get(job.source_id),
|
||||
)
|
||||
# Directory ingestion returns list[Document] — workers ingest single
|
||||
# resources, so a list here is a programming error in the caller.
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -7,10 +7,11 @@ from haiku.rag.ingester.metadata import (
|
|||
build_providers,
|
||||
load_metadata_providers,
|
||||
)
|
||||
from haiku.rag.ingester.sources.base import FetchResult
|
||||
|
||||
|
||||
class _Provider:
|
||||
async def __call__(self, source_id: str, uri: str) -> dict:
|
||||
async def __call__(self, source_id: str, uri: str, result: FetchResult) -> dict:
|
||||
return {"source": source_id}
|
||||
|
||||
|
||||
|
|
@ -50,12 +51,18 @@ def test_load_is_empty_when_none_registered(monkeypatch):
|
|||
@pytest.mark.asyncio
|
||||
async def test_callable_object_satisfies_protocol():
|
||||
class Provider:
|
||||
async def __call__(self, source_id: str, uri: str) -> dict:
|
||||
async def __call__(self, source_id: str, uri: str, result: FetchResult) -> dict:
|
||||
return {"classification": "secret"}
|
||||
|
||||
provider = Provider()
|
||||
result = FetchResult(
|
||||
uri="u",
|
||||
body=b"x",
|
||||
content_type="text/plain",
|
||||
content_hash="9dd4e461268c8034f5c8564e155c67a6",
|
||||
)
|
||||
assert isinstance(provider, MetadataProvider)
|
||||
assert await provider("src", "u") == {"classification": "secret"}
|
||||
assert await provider("src", "u", result) == {"classification": "secret"}
|
||||
|
||||
|
||||
def test_build_providers_instantiates_named_factories():
|
||||
|
|
|
|||
|
|
@ -59,7 +59,10 @@ 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", metadata=None
|
||||
"https://example.com/a.pdf",
|
||||
sources=None,
|
||||
source_id="src",
|
||||
metadata_provider=None,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -80,7 +83,7 @@ async def test_upsert_threads_configured_sources_to_client():
|
|||
"https://example.com/a.pdf",
|
||||
sources=[configured],
|
||||
source_id="src",
|
||||
metadata=None,
|
||||
metadata_provider=None,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -91,19 +94,20 @@ class _MetadataProvider:
|
|||
self._metadata = metadata or {}
|
||||
self._error = error
|
||||
|
||||
async def __call__(self, source_id: str, uri: str) -> dict:
|
||||
async def __call__(self, source_id: str, uri: str, result: FetchResult) -> 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():
|
||||
async def test_provider_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"})}
|
||||
provider = _MetadataProvider({"classification": "secret"})
|
||||
providers = {"src": provider}
|
||||
|
||||
await run_job(client, _job(), metadata_providers=providers)
|
||||
|
||||
|
|
@ -111,37 +115,29 @@ async def test_provider_metadata_passed_to_client():
|
|||
"https://example.com/a.pdf",
|
||||
sources=None,
|
||||
source_id="src",
|
||||
metadata={"classification": "secret", "source": "src"},
|
||||
metadata_provider=provider,
|
||||
)
|
||||
|
||||
|
||||
@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)."""
|
||||
async def test_provider_not_called_in_pipeline():
|
||||
"""Provider execution happens inside create_document_from_source after
|
||||
fetch, where FetchResult exists."""
|
||||
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",
|
||||
}
|
||||
)
|
||||
}
|
||||
provider = _MetadataProvider(
|
||||
error=AssertionError("pipeline must not call provider")
|
||||
)
|
||||
|
||||
await run_job(client, _job(), metadata_providers=providers)
|
||||
await run_job(client, _job(), metadata_providers={"src": provider})
|
||||
|
||||
client.create_document_from_source.assert_awaited_once_with(
|
||||
"https://example.com/a.pdf",
|
||||
sources=None,
|
||||
source_id="src",
|
||||
metadata={"classification": "secret", "source": "src"},
|
||||
metadata_provider=provider,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -157,19 +153,23 @@ async def test_no_provider_for_source_passes_no_metadata():
|
|||
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
|
||||
"https://example.com/a.pdf",
|
||||
sources=None,
|
||||
source_id="src",
|
||||
metadata_provider=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_error_is_classified_and_blocks_ingest():
|
||||
async def test_provider_error_from_client_is_classified_and_blocks_ingest():
|
||||
client = _mock_client()
|
||||
providers = {"src": _MetadataProvider(error=httpx.ConnectError("provider down"))}
|
||||
client.create_document_from_source.side_effect = httpx.ConnectError("provider down")
|
||||
providers = {"src": _MetadataProvider()}
|
||||
|
||||
with pytest.raises(TransientError):
|
||||
await run_job(client, _job(), metadata_providers=providers)
|
||||
|
||||
client.create_document_from_source.assert_not_awaited()
|
||||
client.create_document_from_source.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -4,12 +4,13 @@ as unchanged. Catches the FS-specific bug where revision was lost in the
|
|||
pipeline and every periodic sweep re-enqueued every file forever.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from haiku.rag.client import HaikuRAG
|
||||
from haiku.rag.ingester.sources.base import SourceEventKind
|
||||
from haiku.rag.ingester.sources.base import FetchResult, SourceEventKind
|
||||
from haiku.rag.ingester.sources.fs import FSSource
|
||||
|
||||
|
||||
|
|
@ -107,3 +108,159 @@ async def test_fs_head_short_circuit_skips_fetch_for_unchanged_revision(
|
|||
|
||||
assert second.id == first.id
|
||||
assert fetch_calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.vcr()
|
||||
async def test_provider_backed_unchanged_revision_keeps_head_short_circuit(
|
||||
temp_db_path, tmp_path, monkeypatch
|
||||
):
|
||||
"""Provider-backed sources keep the cheap HEAD path when the revision is
|
||||
unchanged. Existing provider metadata persists until the content changes."""
|
||||
file_path = tmp_path / "doc.md"
|
||||
file_path.write_text("hello")
|
||||
content_hash = hashlib.md5(b"hello", usedforsecurity=False).hexdigest()
|
||||
revision = str(file_path.stat().st_mtime_ns)
|
||||
stored_uri = file_path.absolute().as_uri()
|
||||
|
||||
seen: dict = {}
|
||||
|
||||
class Provider:
|
||||
async def __call__(self, source_id: str, uri: str, result: FetchResult) -> dict:
|
||||
seen["called"] = True
|
||||
raise AssertionError("provider must not run on unchanged HEAD")
|
||||
|
||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||
first = await client.create_document(
|
||||
"hello",
|
||||
uri=stored_uri,
|
||||
metadata={
|
||||
"md5": content_hash,
|
||||
"source_revision": revision,
|
||||
"content_type": "text/markdown",
|
||||
"classification": "secret",
|
||||
},
|
||||
)
|
||||
|
||||
fetch_calls: list[str] = []
|
||||
original_fetch = FSSource.fetch
|
||||
|
||||
async def _track_fetch(self: FSSource, uri: str): # type: ignore[no-untyped-def]
|
||||
fetch_calls.append(uri)
|
||||
return await original_fetch(self, uri)
|
||||
|
||||
monkeypatch.setattr(FSSource, "fetch", _track_fetch)
|
||||
|
||||
second = await client.create_document_from_source(
|
||||
file_path, metadata_provider=Provider()
|
||||
)
|
||||
|
||||
assert second.id == first.id
|
||||
assert fetch_calls == []
|
||||
assert seen == {}
|
||||
assert second.metadata["classification"] == "secret"
|
||||
assert second.metadata["md5"] == first.metadata["md5"]
|
||||
assert second.metadata["source_revision"] == first.metadata["source_revision"]
|
||||
assert second.metadata["content_type"] == first.metadata["content_type"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.vcr()
|
||||
async def test_metadata_provider_applies_to_fresh_create(temp_db_path, tmp_path):
|
||||
"""Fresh ingests pass FetchResult to the provider, merge provider metadata,
|
||||
and still keep source-owned metadata authoritative."""
|
||||
file_path = tmp_path / "doc.md"
|
||||
file_path.write_text("hello")
|
||||
content_hash = hashlib.md5(b"hello", usedforsecurity=False).hexdigest()
|
||||
revision = str(file_path.stat().st_mtime_ns)
|
||||
|
||||
seen: dict = {}
|
||||
|
||||
class Provider:
|
||||
async def __call__(self, source_id: str, uri: str, result: FetchResult) -> dict:
|
||||
seen["source_id"] = source_id
|
||||
seen["uri"] = uri
|
||||
seen["body"] = result.body
|
||||
seen["disk_path"] = result.disk_path
|
||||
seen["content_type"] = result.content_type
|
||||
return {
|
||||
"classification": "secret",
|
||||
"md5": "spoof",
|
||||
"source_revision": "spoof",
|
||||
"content_type": "text/spoof",
|
||||
}
|
||||
|
||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||
doc = await client.create_document_from_source(
|
||||
file_path, metadata_provider=Provider()
|
||||
)
|
||||
|
||||
assert seen["body"] == b"hello"
|
||||
assert seen["disk_path"] == file_path
|
||||
assert seen["uri"] == str(file_path)
|
||||
assert seen["source_id"].startswith("fs:")
|
||||
assert doc.metadata["classification"] == "secret"
|
||||
assert doc.metadata["md5"] == content_hash
|
||||
assert doc.metadata["source_revision"] == revision
|
||||
assert doc.metadata["content_type"] == seen["content_type"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.vcr()
|
||||
async def test_provider_mutating_fetch_result_cannot_corrupt_source_metadata(
|
||||
temp_db_path, tmp_path
|
||||
):
|
||||
"""A provider only contributes metadata via its returned dict. Mutating the
|
||||
FetchResult it receives must not reach the md5 short-circuit or source
|
||||
metadata, so source-owned keys stay authoritative."""
|
||||
file_path = tmp_path / "doc.md"
|
||||
file_path.write_text("hello")
|
||||
content_hash = hashlib.md5(b"hello", usedforsecurity=False).hexdigest()
|
||||
|
||||
class Provider:
|
||||
async def __call__(self, source_id: str, uri: str, result: FetchResult) -> dict:
|
||||
result.content_hash = "spoof"
|
||||
result.extra_metadata["md5"] = "spoof"
|
||||
result.extra_metadata["injected"] = "x"
|
||||
return {"classification": "secret"}
|
||||
|
||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||
doc = await client.create_document_from_source(
|
||||
file_path, metadata_provider=Provider()
|
||||
)
|
||||
|
||||
assert doc.metadata["classification"] == "secret"
|
||||
assert doc.metadata["md5"] == content_hash
|
||||
assert "injected" not in doc.metadata
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.vcr()
|
||||
async def test_directory_ingest_threads_configured_source_to_provider(
|
||||
temp_db_path, tmp_path
|
||||
):
|
||||
"""Directory ingestion with a configured source passes that source's id and
|
||||
fetch context to each child, so the provider sees the configured source id
|
||||
rather than an ad-hoc fs: identity."""
|
||||
(tmp_path / "doc.md").write_text("hello")
|
||||
|
||||
seen_source_ids: list[str] = []
|
||||
|
||||
class Provider:
|
||||
async def __call__(self, source_id: str, uri: str, result: FetchResult) -> dict:
|
||||
seen_source_ids.append(source_id)
|
||||
return {"collection": source_id}
|
||||
|
||||
source = FSSource(root=tmp_path, source_id="docs")
|
||||
|
||||
async with HaikuRAG(temp_db_path, create=True) as client:
|
||||
docs = await client.create_document_from_source(
|
||||
tmp_path,
|
||||
sources=[source],
|
||||
source_id="docs",
|
||||
metadata_provider=Provider(),
|
||||
)
|
||||
|
||||
assert isinstance(docs, list)
|
||||
assert seen_source_ids == ["docs"]
|
||||
assert docs[0].metadata["collection"] == "docs"
|
||||
|
|
|
|||
|
|
@ -685,7 +685,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, metadata=None):
|
||||
def _route(uri, *, sources=None, source_id=None, metadata_provider=None):
|
||||
if source_id == "bad":
|
||||
raise TransientError("downstream down")
|
||||
return Document(
|
||||
|
|
|
|||
Loading…
Reference in a new issue