From 37099988a681bc9063b7ecb99477a342bd435956 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Wed, 29 Apr 2026 13:21:26 +0300 Subject: [PATCH] poll S3 prefixes with S3Watcher and wire into serve --- haiku_rag_slim/haiku/rag/app.py | 16 +- haiku_rag_slim/haiku/rag/monitor.py | 107 +++++++- tests/test_s3_monitor.py | 392 ++++++++++++++++++++++++++++ 3 files changed, 513 insertions(+), 2 deletions(-) create mode 100644 tests/test_s3_monitor.py diff --git a/haiku_rag_slim/haiku/rag/app.py b/haiku_rag_slim/haiku/rag/app.py index 65281277..a6ecd55f 100644 --- a/haiku_rag_slim/haiku/rag/app.py +++ b/haiku_rag_slim/haiku/rag/app.py @@ -21,7 +21,7 @@ from rich.syntax import Syntax from haiku.rag.client import HaikuRAG, RebuildMode from haiku.rag.config import AppConfig, Config from haiku.rag.mcp import create_mcp_server -from haiku.rag.monitor import FileWatcher +from haiku.rag.monitor import FileWatcher, S3Watcher from haiku.rag.store.models.document import Document if TYPE_CHECKING: @@ -821,6 +821,20 @@ class HaikuRAGApp: # pragma: no cover monitor_task = asyncio.create_task(monitor.observe()) tasks.append(monitor_task) + if self.config.monitor.s3: + from haiku.rag.converters import get_converter + + supported_extensions = get_converter( + self.config + ).supported_extensions + for entry in self.config.monitor.s3: + s3_watcher = S3Watcher( + client=client, + entry=entry, + supported_extensions=supported_extensions, + ) + tasks.append(asyncio.create_task(s3_watcher.observe())) + # Start MCP server if enabled if enable_mcp: server = create_mcp_server( diff --git a/haiku_rag_slim/haiku/rag/monitor.py b/haiku_rag_slim/haiku/rag/monitor.py index a8ca8bdb..7cabdfd4 100644 --- a/haiku_rag_slim/haiku/rag/monitor.py +++ b/haiku_rag_slim/haiku/rag/monitor.py @@ -2,13 +2,15 @@ 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 haiku.rag.client import HaikuRAG -from haiku.rag.config import AppConfig, Config +from haiku.rag.config import AppConfig, Config, S3MonitorEntry from haiku.rag.store.models.document import Document +from haiku.rag.utils import escape_sql_string if TYPE_CHECKING: pass @@ -215,3 +217,106 @@ class FileWatcher: logger.info(f"Deleted document {existing_doc.id} for {file}") except Exception as e: logger.error(f"Failed to delete document for {file}: {e}") + + +class S3Watcher: + """Polls an S3 prefix and keeps documents in sync with the index. + + Uses ListObjectsV2 ETags as the cheap-skip key. When a key's listing + ETag differs from the stored `metadata["etag"]`, delegates to + `client.create_document_from_source` which performs the full + HeadObject + GetObject + MD5 compare two-stage detection. + """ + + def __init__( + self, + client: HaikuRAG, + entry: S3MonitorEntry, + supported_extensions: list[str], + ) -> None: + from haiku.rag.s3 import make_s3_session + + parsed = urlparse(entry.uri) + if not parsed.netloc: + raise ValueError(f"Invalid S3 monitor URI: {entry.uri}") + + self.client = client + self.entry = entry + self.bucket = parsed.netloc + self.prefix = parsed.path.lstrip("/") + self.uri_prefix = f"s3://{self.bucket}/{self.prefix}" + self._make_s3_session = make_s3_session + self.filter = FileFilter( + ignore_patterns=entry.ignore_patterns or None, + include_patterns=entry.include_patterns or None, + supported_extensions=supported_extensions, + ) + + async def observe(self) -> None: + logger.info( + f"Watching S3 {self.entry.uri} (poll_interval={self.entry.poll_interval}s)" + ) + await self.refresh() + while True: + await asyncio.sleep(self.entry.poll_interval) + try: + await self.refresh() + except Exception as e: + logger.error(f"S3 watcher refresh failed for {self.entry.uri}: {e}") + + async def refresh(self) -> None: + uris_seen: dict[str, str] = {} + session, client_kwargs = self._make_s3_session(self.entry.storage_options) + + async with session.client("s3", **client_kwargs) as s3: + paginator = s3.get_paginator("list_objects_v2") + async for page in paginator.paginate( + Bucket=self.bucket, Prefix=self.prefix + ): + for obj in page.get("Contents", []): + key = obj["Key"] + if not self.filter.include_file(key): + continue + uri = f"s3://{self.bucket}/{key}" + uris_seen[uri] = obj["ETag"].strip('"') + + existing_etags = await self._existing_etags_under_prefix() + + for uri, etag in uris_seen.items(): + if existing_etags.get(uri) == etag: + continue + await self._upsert_object(uri) + + if self.entry.delete_orphans: + await self._delete_orphans(set(uris_seen.keys()), existing_etags) + + async def _existing_etags_under_prefix(self) -> dict[str, str]: + safe_prefix = escape_sql_string(self.uri_prefix) + docs = await self.client.list_documents(filter=f"uri LIKE '{safe_prefix}%'") + return { + doc.uri: (doc.metadata or {}).get("etag", "") for doc in docs if doc.uri + } + + async def _upsert_object(self, uri: str) -> Document | None: + try: + result = await self.client.create_document_from_source( + uri, storage_options=self.entry.storage_options + ) + doc = result if isinstance(result, Document) else result[0] + logger.info(f"Upserted document {doc.id} from {uri}") + return doc + except Exception as e: + logger.error(f"Failed to upsert document from {uri}: {e}") + return None + + async def _delete_orphans( + self, uris_seen: set[str], existing_etags: dict[str, str] + ) -> None: + for uri in existing_etags.keys() - uris_seen: + try: + doc = await self.client.get_document_by_uri(uri) + if doc and doc.id: + await self.client.delete_document(doc.id) + logger.info(f"Deleted orphaned document {doc.id} for {uri}") + except Exception as e: + logger.error(f"Failed to delete orphan {uri}: {e}") diff --git a/tests/test_s3_monitor.py b/tests/test_s3_monitor.py new file mode 100644 index 00000000..0f946b1f --- /dev/null +++ b/tests/test_s3_monitor.py @@ -0,0 +1,392 @@ +import asyncio +import sys +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from haiku.rag.client import HaikuRAG +from haiku.rag.config import AppConfig, MonitorConfig, S3MonitorEntry +from haiku.rag.store.models.document import Document + + +@pytest.fixture +def fake_aioboto3(monkeypatch): + fake = MagicMock() + monkeypatch.setitem(sys.modules, "aioboto3", fake) + return fake + + +@pytest.fixture +def s3_paginator(): + """Return (paginator_mock, set_pages) — `set_pages([page, ...])` rewires the iterator.""" + pages: list[dict] = [] + + async def _paginate(**_kwargs): + for page in pages: + yield page + + paginator = MagicMock() + paginator.paginate.side_effect = lambda **kw: _paginate(**kw) + + def set_pages(new_pages): + pages.clear() + pages.extend(new_pages) + + return paginator, set_pages + + +@pytest.fixture +def fake_s3_client(fake_aioboto3, s3_paginator): + paginator, set_pages = s3_paginator + s3_client = MagicMock() + s3_client.get_paginator.return_value = paginator + + client_ctx = AsyncMock() + client_ctx.__aenter__.return_value = s3_client + client_ctx.__aexit__.return_value = None + + session = MagicMock() + session.client.return_value = client_ctx + fake_aioboto3.Session.return_value = session + + return s3_client, set_pages + + +def _entry(**kwargs): + base = { + "uri": "s3://my-bucket/incoming/", + "poll_interval": 60, + "delete_orphans": False, + "ignore_patterns": [], + "include_patterns": [], + "storage_options": {}, + } + base.update(kwargs) + return S3MonitorEntry(**base) + + +def _doc(uri: str, etag: str, doc_id: str | None = None) -> Document: + return Document( + id=doc_id or uri, + content="...", + uri=uri, + metadata={"etag": etag, "md5": "deadbeef"}, + ) + + +@pytest.mark.asyncio +async def test_s3_watcher_refresh_upserts_new_objects(fake_s3_client): + s3, set_pages = fake_s3_client + set_pages( + [ + { + "Contents": [ + {"Key": "incoming/a.txt", "ETag": '"abc"'}, + {"Key": "incoming/b.txt", "ETag": '"def"'}, + ] + } + ] + ) + + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [] + rag.create_document_from_source.return_value = Document( + id="x", content="...", uri="s3://my-bucket/incoming/a.txt" + ) + + watcher = S3Watcher( + client=rag, entry=_entry(), supported_extensions=[".txt", ".md", ".pdf"] + ) + await watcher.refresh() + + assert rag.create_document_from_source.await_count == 2 + called_uris = {c.args[0] for c in rag.create_document_from_source.await_args_list} + assert called_uris == { + "s3://my-bucket/incoming/a.txt", + "s3://my-bucket/incoming/b.txt", + } + + +@pytest.mark.asyncio +async def test_s3_watcher_skips_unchanged_etag(fake_s3_client): + s3, set_pages = fake_s3_client + set_pages([{"Contents": [{"Key": "incoming/a.txt", "ETag": '"abc"'}]}]) + + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [_doc("s3://my-bucket/incoming/a.txt", "abc")] + + watcher = S3Watcher(client=rag, entry=_entry(), supported_extensions=[".txt"]) + await watcher.refresh() + + rag.create_document_from_source.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_s3_watcher_upserts_when_etag_differs(fake_s3_client): + s3, set_pages = fake_s3_client + set_pages([{"Contents": [{"Key": "incoming/a.txt", "ETag": '"new"'}]}]) + + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [_doc("s3://my-bucket/incoming/a.txt", "old")] + rag.create_document_from_source.return_value = Document( + id="x", content="...", uri="s3://my-bucket/incoming/a.txt" + ) + + watcher = S3Watcher(client=rag, entry=_entry(), supported_extensions=[".txt"]) + await watcher.refresh() + + rag.create_document_from_source.assert_awaited_once_with( + "s3://my-bucket/incoming/a.txt", storage_options={} + ) + + +@pytest.mark.asyncio +async def test_s3_watcher_strips_etag_quotes(fake_s3_client): + s3, set_pages = fake_s3_client + set_pages([{"Contents": [{"Key": "incoming/a.txt", "ETag": '"abc"'}]}]) + + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [ + _doc("s3://my-bucket/incoming/a.txt", "abc") # already stripped in storage + ] + + watcher = S3Watcher(client=rag, entry=_entry(), supported_extensions=[".txt"]) + await watcher.refresh() + + rag.create_document_from_source.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_s3_watcher_deletes_orphans_when_enabled(fake_s3_client): + s3, set_pages = fake_s3_client + set_pages([{"Contents": [{"Key": "incoming/a.txt", "ETag": '"abc"'}]}]) + + from haiku.rag.monitor import S3Watcher + + a_doc = _doc("s3://my-bucket/incoming/a.txt", "abc", doc_id="a-id") + orphan = _doc("s3://my-bucket/incoming/old.txt", "stale", doc_id="orphan-id") + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [a_doc, orphan] + rag.get_document_by_uri.return_value = orphan + + watcher = S3Watcher( + client=rag, + entry=_entry(delete_orphans=True), + supported_extensions=[".txt"], + ) + await watcher.refresh() + + rag.delete_document.assert_awaited_once_with("orphan-id") + + +@pytest.mark.asyncio +async def test_s3_watcher_does_not_delete_orphans_when_disabled(fake_s3_client): + s3, set_pages = fake_s3_client + set_pages([{"Contents": []}]) + + from haiku.rag.monitor import S3Watcher + + orphan = _doc("s3://my-bucket/incoming/old.txt", "stale", doc_id="orphan-id") + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [orphan] + + watcher = S3Watcher( + client=rag, + entry=_entry(delete_orphans=False), + supported_extensions=[".txt"], + ) + await watcher.refresh() + + rag.delete_document.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_s3_watcher_orphan_scope_is_per_entry(fake_s3_client): + """A doc under a different bucket prefix must not be touched.""" + s3, set_pages = fake_s3_client + set_pages([{"Contents": []}]) + + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [] # filter scopes to my-bucket + + watcher = S3Watcher( + client=rag, + entry=_entry(delete_orphans=True), + supported_extensions=[".txt"], + ) + await watcher.refresh() + + rag.list_documents.assert_awaited_once() + filter_kwarg = rag.list_documents.await_args.kwargs["filter"] + assert filter_kwarg == "uri LIKE 's3://my-bucket/incoming/%'" + + +@pytest.mark.asyncio +async def test_s3_watcher_applies_include_and_ignore_patterns(fake_s3_client): + s3, set_pages = fake_s3_client + set_pages( + [ + { + "Contents": [ + {"Key": "incoming/keep.md", "ETag": '"1"'}, + {"Key": "incoming/draft.md", "ETag": '"2"'}, + {"Key": "incoming/skip.txt", "ETag": '"3"'}, + ] + } + ] + ) + + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [] + rag.create_document_from_source.return_value = Document( + id="x", content="...", uri="s3://my-bucket/incoming/keep.md" + ) + + watcher = S3Watcher( + client=rag, + entry=_entry(include_patterns=["*.md"], ignore_patterns=["draft*"]), + supported_extensions=[".md", ".txt"], + ) + await watcher.refresh() + + assert rag.create_document_from_source.await_count == 1 + assert ( + rag.create_document_from_source.await_args.args[0] + == "s3://my-bucket/incoming/keep.md" + ) + + +@pytest.mark.asyncio +async def test_s3_watcher_observe_survives_transient_list_failure(fake_s3_client): + """First refresh succeeds; second refresh raises; loop survives and recovers.""" + s3, set_pages = fake_s3_client + + pages_initial = [{"Contents": [{"Key": "incoming/a.txt", "ETag": '"abc"'}]}] + pages_after = [{"Contents": [{"Key": "incoming/a.txt", "ETag": '"abc"'}]}] + + set_pages(pages_initial) + paginate_calls = {"n": 0} + + async def paginate_side_effect(**_kwargs): + paginate_calls["n"] += 1 + if paginate_calls["n"] == 2: + raise RuntimeError("transient list failure") + for page in pages_after if paginate_calls["n"] > 1 else pages_initial: + yield page + + s3.get_paginator.return_value.paginate.side_effect = paginate_side_effect + + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + rag.list_documents.return_value = [] + rag.create_document_from_source.return_value = Document( + id="x", content="...", uri="s3://my-bucket/incoming/a.txt" + ) + + watcher = S3Watcher( + client=rag, + entry=_entry(poll_interval=0), + supported_extensions=[".txt"], + ) + task = asyncio.create_task(watcher.observe()) + + # Let three iterations run: initial refresh, transient failure, recovery. + for _ in range(20): + await asyncio.sleep(0) + if paginate_calls["n"] >= 3: + break + + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + + assert paginate_calls["n"] >= 3 # loop kept going past the failure + + +@pytest.mark.asyncio +async def test_s3_watcher_invalid_uri_rejected(): + from haiku.rag.monitor import S3Watcher + + rag = AsyncMock(spec=HaikuRAG) + with pytest.raises(ValueError, match="Invalid S3 monitor URI"): + S3Watcher( + client=rag, + entry=S3MonitorEntry(uri="s3://"), + supported_extensions=[".txt"], + ) + + +@pytest.mark.asyncio +async def test_serve_starts_one_s3_task_per_entry(monkeypatch, fake_s3_client): + """`serve` wires one S3Watcher task per MonitorConfig.s3 entry.""" + from haiku.rag import app as app_module + + captured_tasks: list = [] + original_create_task = asyncio.create_task + + def tracking_create_task(coro, *args, **kwargs): + captured_tasks.append(coro) + return original_create_task(coro, *args, **kwargs) + + monkeypatch.setattr(app_module.asyncio, "create_task", tracking_create_task) + + config = AppConfig( + monitor=MonitorConfig( + s3=[ + S3MonitorEntry(uri="s3://bucket-a/x/"), + S3MonitorEntry(uri="s3://bucket-b/y/"), + ] + ) + ) + + fw_observe_calls = {"n": 0} + + async def fake_fw_observe(self): + fw_observe_calls["n"] += 1 + + monkeypatch.setattr(app_module.FileWatcher, "observe", fake_fw_observe) + + s3_observe_calls = {"n": 0} + + async def fake_s3_observe(self): + s3_observe_calls["n"] += 1 + + monkeypatch.setattr(app_module.S3Watcher, "observe", fake_s3_observe) + + # Provide a dummy supported_extensions to skip docling import. + class _Conv: + supported_extensions = [".txt"] + + monkeypatch.setattr("haiku.rag.converters.get_converter", lambda cfg: _Conv()) + + import tempfile + from pathlib import Path + + with tempfile.TemporaryDirectory() as tmp: + db_path = Path(tmp) / "db.lancedb" + app = app_module.HaikuRAGApp(db_path=db_path, config=config) + async with HaikuRAG(db_path, config=config, create=True): + pass # create the database + + await app.serve(enable_monitor=True, enable_mcp=False) + + # Both S3 entries should have triggered observe(), plus the FileWatcher. + assert fw_observe_calls["n"] == 1 + assert s3_observe_calls["n"] == 2