poll S3 prefixes with S3Watcher and wire into serve
This commit is contained in:
parent
d009da06d6
commit
37099988a6
3 changed files with 513 additions and 2 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
392
tests/test_s3_monitor.py
Normal file
392
tests/test_s3_monitor.py
Normal file
|
|
@ -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
|
||||
Loading…
Reference in a new issue