ingester: add source-adapter scaffolding + FS source

This commit is contained in:
Yiorgis Gozadinos 2026-05-21 15:58:48 +03:00
parent 2d199037b2
commit d0a730ef60
No known key found for this signature in database
9 changed files with 462 additions and 62 deletions

View file

@ -0,0 +1,19 @@
from haiku.rag.ingester.sources.base import (
FetchResult,
RevisionSnapshot,
Source,
SourceEvent,
SourceEventKind,
)
from haiku.rag.ingester.sources.filter import FileFilter
from haiku.rag.ingester.sources.fs import FSSource
__all__ = [
"FetchResult",
"FileFilter",
"FSSource",
"RevisionSnapshot",
"Source",
"SourceEvent",
"SourceEventKind",
]

View file

@ -0,0 +1,54 @@
from collections.abc import AsyncIterator, Mapping
from datetime import datetime
from enum import StrEnum
from typing import Protocol, runtime_checkable
from pydantic import BaseModel, Field
# uri -> revision. Captures what revisions of which URIs we had last seen
# for a given source. Passed to discover() so the source can yield only
# UPSERT / UNCHANGED / DELETE deltas instead of a full re-scan.
RevisionSnapshot = Mapping[str, str]
class SourceEventKind(StrEnum):
UPSERT = "upsert"
DELETE = "delete"
# Emitted for resources whose revision matches the snapshot. The poller
# uses these to bump last_seen_at without enqueueing work.
UNCHANGED = "unchanged"
class SourceEvent(BaseModel):
source_id: str
uri: str
kind: SourceEventKind
# Backend's own change indicator (mtime for FS, ETag for HTTP/S3, etc.).
# Opaque to consumers — only compared, never parsed. None for DELETE.
revision: str | None = None
discovered_at: datetime
class FetchResult(BaseModel):
uri: str
body: bytes
content_type: str
# MD5 of body. Stored in document metadata as the dedup key — lets the
# pipeline short-circuit when bytes are identical but the revision differs
# (e.g. S3 multipart re-upload landing a new ETag on the same content).
content_hash: str
revision: str | None = None
extra_metadata: dict[str, str] = Field(default_factory=dict)
@runtime_checkable
class Source(Protocol):
source_id: str
def supports(self, uri: str) -> bool: ...
async def fetch(self, uri: str) -> FetchResult: ...
def discover(
self, since: RevisionSnapshot | None = None
) -> AsyncIterator[SourceEvent]: ...

View file

@ -0,0 +1,51 @@
import pathspec
from watchfiles import Change, DefaultFilter
def _default_supported_extensions() -> list[str]:
from haiku.rag.converters.docling_local import DoclingLocalConverter
from haiku.rag.converters.text_utils import TextFileHandler
return DoclingLocalConverter.docling_extensions + TextFileHandler.text_extensions
class FileFilter(DefaultFilter):
def __init__(
self,
*,
ignore_patterns: list[str] | None = None,
include_patterns: list[str] | None = None,
supported_extensions: list[str] | None = None,
) -> None:
if supported_extensions is None:
supported_extensions = _default_supported_extensions()
self.extensions = tuple(supported_extensions)
self.ignore_spec = (
pathspec.PathSpec.from_lines("gitwildmatch", ignore_patterns)
if ignore_patterns
else None
)
self.include_spec = (
pathspec.PathSpec.from_lines("gitwildmatch", include_patterns)
if include_patterns
else None
)
super().__init__()
def __call__(self, change: Change, path: str) -> bool:
if not self.include_file(path):
return False
return super().__call__(change, path)
def include_file(self, path: str) -> bool:
if not path.endswith(self.extensions):
return False
if self.include_spec and not self.include_spec.match_file(path):
return False
if self.ignore_spec and self.ignore_spec.match_file(path):
return False
return True

View file

@ -0,0 +1,120 @@
import hashlib
import mimetypes
from collections.abc import AsyncIterator
from datetime import UTC, datetime
from pathlib import Path
from urllib.parse import unquote, urlparse
from haiku.rag.ingester.sources.base import (
FetchResult,
RevisionSnapshot,
SourceEvent,
SourceEventKind,
)
from haiku.rag.ingester.sources.filter import (
FileFilter,
_default_supported_extensions,
)
def _uri_to_path(uri: str) -> Path:
parsed = urlparse(uri)
if parsed.scheme in ("", "file"):
path = parsed.path if parsed.scheme == "file" else uri
return Path(unquote(path))
raise ValueError(f"Unsupported URI scheme for FSSource: {uri}")
class FSSource:
def __init__(
self,
*,
root: Path,
ignore_patterns: list[str] | None = None,
include_patterns: list[str] | None = None,
supported_extensions: list[str] | None = None,
) -> None:
# Resolve so symlinks and relative paths collapse to one canonical
# source_id. The queue uses source_id as a foreign key — two paths
# for the same root would mean duplicate sync_state rows.
self.root = Path(root).resolve()
self.source_id = f"fs:{self.root}"
self.supported_extensions = (
list(supported_extensions)
if supported_extensions is not None
else _default_supported_extensions()
)
self.filter = FileFilter(
ignore_patterns=ignore_patterns,
include_patterns=include_patterns,
supported_extensions=self.supported_extensions,
)
def supports(self, uri: str) -> bool:
scheme = urlparse(uri).scheme
if scheme not in ("", "file"):
return False
try:
_uri_to_path(uri)
except ValueError:
return False
return True
async def fetch(self, uri: str) -> FetchResult:
path = _uri_to_path(uri)
body = path.read_bytes()
content_type, _ = mimetypes.guess_type(path.name)
if content_type is None:
content_type = "application/octet-stream"
# mtime_ns rather than st_mtime: nanosecond integer avoids float
# precision collisions on rapid edits.
revision = str(path.stat().st_mtime_ns)
return FetchResult(
uri=path.as_uri(),
body=body,
content_type=content_type,
content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(),
revision=revision,
)
async def discover(
self, since: RevisionSnapshot | None = None
) -> AsyncIterator[SourceEvent]:
snapshot: dict[str, str] = dict(since) if since else {}
now = datetime.now(UTC)
seen: set[str] = set()
for path in sorted(self.root.rglob("*")):
if not path.is_file():
continue
if not self.filter.include_file(str(path)):
continue
uri = path.as_uri()
revision = str(path.stat().st_mtime_ns)
seen.add(uri)
previous = snapshot.get(uri)
kind = (
SourceEventKind.UNCHANGED
if previous == revision
else SourceEventKind.UPSERT
)
yield SourceEvent(
source_id=self.source_id,
uri=uri,
kind=kind,
revision=revision,
discovered_at=now,
)
# Anything in the snapshot we didn't encounter during the walk is
# gone from the source. Emit DELETE so the poller can clean up.
for uri in snapshot:
if uri in seen:
continue
yield SourceEvent(
source_id=self.source_id,
uri=uri,
kind=SourceEventKind.DELETE,
revision=None,
discovered_at=now,
)

View file

@ -1,78 +1,19 @@
import asyncio
import logging
from pathlib import Path
from typing import TYPE_CHECKING
from urllib.parse import urlparse
import pathspec
from watchfiles import Change, DefaultFilter, awatch
from watchfiles import Change, awatch
from haiku.rag.client import HaikuRAG
from haiku.rag.config import AppConfig, Config, S3MonitorEntry
from haiku.rag.ingester.sources.filter import FileFilter
from haiku.rag.store.models.document import Document
from haiku.rag.utils import escape_sql_string
if TYPE_CHECKING:
pass
logger = logging.getLogger(__name__)
class FileFilter(DefaultFilter):
def __init__(
self,
*,
ignore_patterns: list[str] | None = None,
include_patterns: list[str] | None = None,
supported_extensions: list[str] | None = None,
) -> None:
if supported_extensions is None:
# Default to docling-local extensions if not provided
from haiku.rag.converters.docling_local import DoclingLocalConverter
from haiku.rag.converters.text_utils import TextFileHandler
supported_extensions = (
DoclingLocalConverter.docling_extensions
+ TextFileHandler.text_extensions
)
self.extensions = tuple(supported_extensions)
self.ignore_spec = (
pathspec.PathSpec.from_lines("gitwildmatch", ignore_patterns)
if ignore_patterns
else None
)
self.include_spec = (
pathspec.PathSpec.from_lines("gitwildmatch", include_patterns)
if include_patterns
else None
)
super().__init__()
def __call__(self, change: Change, path: str) -> bool:
if not self.include_file(path):
return False
# Apply default watchfiles filter
return super().__call__(change, path)
def include_file(self, path: str) -> bool:
"""Check if a file should be included based on filters."""
# Check extension filter
if not path.endswith(self.extensions):
return False
# Apply include patterns if specified (whitelist mode)
if self.include_spec:
if not self.include_spec.match_file(path):
return False
# Apply ignore patterns (blacklist mode)
if self.ignore_spec:
if self.ignore_spec.match_file(path):
return False
return True
__all__ = ["FileFilter", "FileWatcher", "S3Watcher"]
class FileWatcher:

View file

View file

@ -0,0 +1,147 @@
import hashlib
from pathlib import Path
import pytest
from haiku.rag.ingester.sources.base import SourceEventKind
from haiku.rag.ingester.sources.fs import FSSource
@pytest.fixture
def fs_root(tmp_path: Path) -> Path:
(tmp_path / "a.md").write_text("alpha")
(tmp_path / "b.txt").write_text("beta")
sub = tmp_path / "sub"
sub.mkdir()
(sub / "c.md").write_text("gamma")
(tmp_path / "skip.log").write_text("noise")
return tmp_path
def test_fs_source_supports_file_uri(fs_root: Path):
src = FSSource(root=fs_root)
assert src.supports((fs_root / "a.md").as_uri())
assert src.supports(str(fs_root / "a.md"))
def test_fs_source_rejects_other_schemes(fs_root: Path):
src = FSSource(root=fs_root)
assert not src.supports("http://example.com/a.md")
assert not src.supports("s3://bucket/a.md")
def test_fs_source_source_id_is_canonical(fs_root: Path):
src = FSSource(root=fs_root)
assert src.source_id == f"fs:{fs_root.resolve()}"
@pytest.mark.asyncio
async def test_fs_source_fetch_returns_bytes_and_md5(fs_root: Path):
src = FSSource(root=fs_root)
target = fs_root / "a.md"
result = await src.fetch(target.as_uri())
assert result.uri == target.as_uri()
assert result.body == b"alpha"
assert (
result.content_hash == hashlib.md5(b"alpha", usedforsecurity=False).hexdigest()
)
assert result.content_type == "text/markdown"
assert result.revision == str(target.stat().st_mtime_ns)
@pytest.mark.asyncio
async def test_fs_source_fetch_accepts_bare_path(fs_root: Path):
src = FSSource(root=fs_root)
target = fs_root / "a.md"
result = await src.fetch(str(target))
assert result.uri == target.as_uri()
@pytest.mark.asyncio
async def test_fs_source_fetch_missing_file_raises(fs_root: Path):
src = FSSource(root=fs_root)
with pytest.raises(FileNotFoundError):
await src.fetch((fs_root / "missing.md").as_uri())
@pytest.mark.asyncio
async def test_fs_source_discover_initial_scan_yields_upsert(fs_root: Path):
src = FSSource(root=fs_root, supported_extensions=[".md", ".txt"])
events = [e async for e in src.discover(since=None)]
uris = {e.uri for e in events}
assert uris == {
(fs_root / "a.md").as_uri(),
(fs_root / "b.txt").as_uri(),
(fs_root / "sub" / "c.md").as_uri(),
}
assert all(e.kind is SourceEventKind.UPSERT for e in events)
assert all(e.source_id == src.source_id for e in events)
assert all(e.revision is not None for e in events)
@pytest.mark.asyncio
async def test_fs_source_discover_unchanged_against_snapshot(fs_root: Path):
src = FSSource(root=fs_root, supported_extensions=[".md", ".txt"])
initial = {e.uri: e.revision or "" async for e in src.discover(since=None)}
again = [e async for e in src.discover(since=initial)]
assert again
assert all(e.kind is SourceEventKind.UNCHANGED for e in again)
@pytest.mark.asyncio
async def test_fs_source_discover_changed_yields_upsert(fs_root: Path):
src = FSSource(root=fs_root, supported_extensions=[".md", ".txt"])
initial = {e.uri: e.revision or "" async for e in src.discover(since=None)}
stale = {uri: "0" for uri in initial}
events = [e async for e in src.discover(since=stale)]
assert {e.kind for e in events} == {SourceEventKind.UPSERT}
@pytest.mark.asyncio
async def test_fs_source_discover_emits_delete_for_missing(fs_root: Path):
src = FSSource(root=fs_root, supported_extensions=[".md", ".txt"])
snapshot = {(fs_root / "ghost.md").as_uri(): "999"}
events = [e async for e in src.discover(since=snapshot)]
deletes = [e for e in events if e.kind is SourceEventKind.DELETE]
assert len(deletes) == 1
assert deletes[0].uri == (fs_root / "ghost.md").as_uri()
assert deletes[0].revision is None
@pytest.mark.asyncio
async def test_fs_source_discover_respects_extension_filter(fs_root: Path):
src = FSSource(root=fs_root, supported_extensions=[".md"])
uris = {e.uri async for e in src.discover(since=None)}
assert (fs_root / "a.md").as_uri() in uris
assert (fs_root / "b.txt").as_uri() not in uris
@pytest.mark.asyncio
async def test_fs_source_discover_respects_ignore_patterns(fs_root: Path):
src = FSSource(
root=fs_root,
supported_extensions=[".md", ".txt"],
ignore_patterns=["**/sub/**"],
)
uris = {e.uri async for e in src.discover(since=None)}
assert (fs_root / "sub" / "c.md").as_uri() not in uris
assert (fs_root / "a.md").as_uri() in uris
@pytest.mark.asyncio
async def test_fs_source_discover_respects_include_patterns(fs_root: Path):
src = FSSource(
root=fs_root,
supported_extensions=[".md", ".txt"],
include_patterns=["**/*.md"],
)
uris = {e.uri async for e in src.discover(since=None)}
assert (fs_root / "b.txt").as_uri() not in uris
assert (fs_root / "a.md").as_uri() in uris
def test_filefilter_backward_compatible_reexport():
from haiku.rag.ingester.sources.filter import FileFilter as IngesterFileFilter
from haiku.rag.monitor import FileFilter as MonitorFileFilter
assert MonitorFileFilter is IngesterFileFilter

View file

@ -0,0 +1,68 @@
from datetime import UTC, datetime
from haiku.rag.ingester.sources.base import (
FetchResult,
Source,
SourceEvent,
SourceEventKind,
)
def test_source_event_kind_values():
assert SourceEventKind.UPSERT.value == "upsert"
assert SourceEventKind.DELETE.value == "delete"
assert SourceEventKind.UNCHANGED.value == "unchanged"
def test_source_event_round_trip():
event = SourceEvent(
source_id="fs:/tmp/docs",
uri="file:///tmp/docs/a.md",
kind=SourceEventKind.UPSERT,
revision="123456",
discovered_at=datetime(2026, 5, 20, 12, 0, 0, tzinfo=UTC),
)
raw = event.model_dump_json()
again = SourceEvent.model_validate_json(raw)
assert again == event
def test_fetch_result_round_trip():
result = FetchResult(
uri="file:///tmp/docs/a.md",
body=b"hello",
content_type="text/markdown",
content_hash="abcd1234",
revision="123456",
extra_metadata={"source": "fs"},
)
raw = result.model_dump_json()
again = FetchResult.model_validate_json(raw)
assert again == result
def test_fetch_result_defaults_extra_metadata_to_empty():
result = FetchResult(
uri="file:///tmp/docs/a.md",
body=b"x",
content_type="text/markdown",
content_hash="x",
revision=None,
)
assert result.extra_metadata == {}
def test_source_protocol_runtime_checkable():
class Dummy:
source_id = "dummy"
def supports(self, uri: str) -> bool:
return True
async def fetch(self, uri: str):
raise NotImplementedError
def discover(self, since=None):
raise NotImplementedError
assert isinstance(Dummy(), Source)