From 717fc3320be99c38e4459b82a91cc6c118cf671d Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 21 May 2026 16:56:27 +0300 Subject: [PATCH] route create_document_from_source through source adapters --- haiku_rag_slim/haiku/rag/client/documents.py | 580 ++++++------------ .../haiku/rag/ingester/sources/base.py | 14 + .../haiku/rag/ingester/sources/fs.py | 12 +- .../haiku/rag/ingester/sources/http.py | 7 + .../haiku/rag/ingester/sources/s3.py | 20 +- tests/ingester/test_fs_source.py | 14 + tests/ingester/test_http_source.py | 7 + tests/ingester/test_s3_source.py | 17 + tests/ingester/test_sources_base.py | 3 + 9 files changed, 295 insertions(+), 379 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/client/documents.py b/haiku_rag_slim/haiku/rag/client/documents.py index 605c8950..4140ac50 100644 --- a/haiku_rag_slim/haiku/rag/client/documents.py +++ b/haiku_rag_slim/haiku/rag/client/documents.py @@ -1,15 +1,15 @@ -import hashlib -import mimetypes import tempfile from pathlib import Path from typing import TYPE_CHECKING from urllib.parse import urlparse -import httpx - -from haiku.rag.client.processing import ensure_chunks_embedded +from haiku.rag.client.processing import ( + ensure_chunks_embedded, + get_extension_from_content_type_or_url, +) from haiku.rag.client.titles import resolve_title from haiku.rag.converters import get_converter +from haiku.rag.ingester.sources import FetchResult, resolve_fetcher from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.document import Document from haiku.rag.store.models.document_item import extract_items @@ -30,38 +30,30 @@ async def _store_document_with_chunks( Handles versioning/rollback on failure. """ - # Ensure all chunks have embeddings before storing chunks = await ensure_chunks_embedded(client._config, chunks) - # Snapshot table versions for versioned rollback (if supported) versions = await client.store.current_table_versions() - # Create the document created_doc = await client.document_repository.create(document) try: assert created_doc.id is not None, ( "Document ID should not be None after creation" ) - # Set document_id and order for all chunks for order, chunk in enumerate(chunks): chunk.document_id = created_doc.id chunk.order = order - # Batch create all chunks in a single operation await client.chunk_repository.create(chunks) - # Extract and store document items for context expansion items = extract_items(created_doc.id, docling_document) await client.document_item_repository.create_items(created_doc.id, items) - # Vacuum old versions in background (non-blocking) if auto_vacuum enabled if client._config.storage.auto_vacuum: client._schedule_vacuum() return created_doc except Exception: - # Roll back to the captured versions and re-raise await client.store.restore_table_versions(versions) raise @@ -92,7 +84,6 @@ async def _update_document_with_chunks( versions = await client.store.current_table_versions() - # Delete existing chunks before writing new ones await client.chunk_repository.delete_by_document_id(document.id) try: @@ -137,14 +128,11 @@ async def create_document( """ from haiku.rag.embeddings import embed_chunks - # Convert → Chunk → Embed using primitives converter = get_converter(client._config) docling_document = await converter.convert_text(content, format=format) chunks = await client.chunk(docling_document) embedded_chunks = await embed_chunks(chunks, client._config) - # Store markdown export as content for better display/readability. - # The original is preserved in docling_document. stored_content = docling_document.export_to_markdown() if title is None: @@ -191,6 +179,112 @@ async def import_document( return await _store_document_with_chunks(client, document, chunks, docling_document) +async def _refresh_doc_metadata( + client: "HaikuRAG", + doc: Document, + *, + title: str | None, + user_metadata: dict, + source_metadata: dict | None, +) -> Document: + """Update a document's title + metadata without re-chunking. Used by the + cheap revision and MD5 short-circuits in create_document_from_source.""" + updated = False + if title is not None and title != doc.title: + doc.title = title + updated = True + + merged = {**(doc.metadata or {}), **user_metadata} + if source_metadata: + merged.update(source_metadata) + if merged != doc.metadata: + doc.metadata = merged + updated = True + + if updated: + return await client.document_repository.update(doc) + return doc + + +async def _ingest_fetch_result( + client: "HaikuRAG", + result: FetchResult, + *, + title: str | None, + user_metadata: dict, + stored_uri: str, + existing_doc: Document | None, +) -> Document: + """Convert / chunk / embed / store a fetched document. Replaces an + existing document if one is supplied.""" + from haiku.rag.embeddings import embed_chunks + + converter = get_converter(client._config) + file_extension = get_extension_from_content_type_or_url( + result.uri, result.content_type + ) + if file_extension not in converter.supported_extensions: + raise ValueError( + f"Unsupported content type/extension: {result.content_type}/{file_extension}" + ) + + source_metadata: dict = { + "contentType": result.content_type, + "md5": result.content_hash, + **result.extra_metadata, + } + + if result.disk_path is not None: + target_path = result.disk_path + cleanup_path: Path | None = None + else: + with tempfile.NamedTemporaryFile( + mode="wb", suffix=file_extension, delete=False + ) as temp_file: + temp_file.write(result.body) + temp_file.flush() + target_path = Path(temp_file.name) + cleanup_path = target_path + + try: + docling_document = await client.convert(target_path, source_uri=result.uri) + chunks = await client.chunk(docling_document) + embedded_chunks = await embed_chunks(chunks, client._config) + finally: + if cleanup_path is not None: + cleanup_path.unlink(missing_ok=True) + + stored_content = docling_document.export_to_markdown() + final_metadata = {**user_metadata, **source_metadata} + + if existing_doc: + existing_doc.content = stored_content + existing_doc.metadata = final_metadata + existing_doc.set_docling(docling_document) + if title is not None: + existing_doc.title = title + elif existing_doc.title is None: + existing_doc.title = await resolve_title( + client._config, docling_document, stored_content + ) + return await _update_document_with_chunks( + client, existing_doc, embedded_chunks, docling_document + ) + + if title is None: + title = await resolve_title(client._config, docling_document, stored_content) + document = Document( + content=stored_content, + uri=stored_uri, + title=title, + metadata=final_metadata, + ) + document.set_docling(docling_document) + return await _store_document_with_chunks( + client, document, embedded_chunks, docling_document + ) + + async def create_document_from_source( client: "HaikuRAG", source: str | Path, @@ -208,9 +302,8 @@ async def create_document_from_source( If ``uri`` is provided, it overrides the URI auto-derived from the source (which is normally ``file://`` for local files or the URL for remote - sources). This is useful when callers want to persist documents under a - logical identifier (e.g. an ArXiv ID) rather than the on-disk path. Not - supported for directory sources, which produce one document per file. + sources). Not supported for directory sources, which produce one document + per file. Returns a single Document for files/URLs, a list for directories. """ @@ -218,372 +311,110 @@ async def create_document_from_source( source_str = str(source) parsed_url = urlparse(source_str) - if parsed_url.scheme in ("http", "https"): - return await _create_or_update_document_from_url( - client, source_str, title=title, metadata=metadata, uri=uri - ) - elif parsed_url.scheme == "s3": - return await _create_or_update_document_from_s3( - client, - source_str, - storage_options=storage_options, - title=title, - metadata=metadata, - uri=uri, - ) - elif parsed_url.scheme == "file": - source_path = Path(parsed_url.path) - else: - source_path = Path(source) if isinstance(source, str) else source - if source_path.is_dir(): - if uri is not None: - raise ValueError( - "uri override is not supported for directory sources; each file " - "produces its own document with its own auto-derived URI." - ) - from haiku.rag.monitor import FileFilter - - documents = [] - filter = FileFilter( - ignore_patterns=client._config.monitor.ignore_patterns or None, - include_patterns=client._config.monitor.include_patterns or None, + # Directory case: recurse with the existing FS filter and produce one + # document per file. Remote schemes (http/s3) never hit this branch. + if parsed_url.scheme in ("", "file"): + local_path = ( + Path(parsed_url.path) + if parsed_url.scheme == "file" + else (Path(source) if isinstance(source, str) else source) ) - for path in source_path.rglob("*"): - if path.is_file() and filter.include_file(str(path)): - doc = await _create_document_from_file( - client, path, title=None, metadata=metadata + if local_path.is_dir(): + if uri is not None: + raise ValueError( + "uri override is not supported for directory sources; each file " + "produces its own document with its own auto-derived URI." ) - documents.append(doc) - return documents + from haiku.rag.ingester.sources.filter import FileFilter - return await _create_document_from_file( - client, source_path, title=title, metadata=metadata, uri=uri - ) - - -async def _create_document_from_file( - client: "HaikuRAG", - source_path: Path, - title: str | None = None, - metadata: dict | None = None, - uri: str | None = None, -) -> Document: - """Create or update a document from a single file path. - - ``uri`` overrides the auto-derived ``file://`` URI; it's used as the - canonical document identifier for lookup and storage. - """ - from haiku.rag.embeddings import embed_chunks - - metadata = metadata or {} - - converter = get_converter(client._config) - if source_path.suffix.lower() not in converter.supported_extensions: - raise ValueError(f"Unsupported file extension: {source_path.suffix}") - - if not source_path.exists(): - raise ValueError(f"File does not exist: {source_path}") - - if uri is None: - uri = source_path.absolute().as_uri() - md5_hash = hashlib.md5(source_path.read_bytes(), usedforsecurity=False).hexdigest() - - content_type, _ = mimetypes.guess_type(str(source_path)) - if not content_type: - content_type = "application/octet-stream" - metadata.update({"contentType": content_type, "md5": md5_hash}) - - # Check if document already exists - existing_doc = await client.get_document_by_uri(uri) - if existing_doc and existing_doc.metadata.get("md5") == md5_hash: - # MD5 unchanged; update title/metadata if provided - updated = False - if title is not None and title != existing_doc.title: - existing_doc.title = title - updated = True - - merged_metadata = {**(existing_doc.metadata or {}), **metadata} - if merged_metadata != existing_doc.metadata: - existing_doc.metadata = merged_metadata - updated = True - - if updated: - return await client.document_repository.update(existing_doc) - return existing_doc - - # Convert → Chunk → Embed - docling_document = await client.convert(source_path) - chunks = await client.chunk(docling_document) - embedded_chunks = await embed_chunks(chunks, client._config) - - stored_content = docling_document.export_to_markdown() - - if existing_doc: - # Update existing document and rechunk - existing_doc.content = stored_content - existing_doc.metadata = metadata - existing_doc.set_docling(docling_document) - if title is not None: - existing_doc.title = title - elif existing_doc.title is None: - existing_doc.title = await resolve_title( - client._config, docling_document, stored_content + documents: list[Document] = [] + filter = FileFilter( + ignore_patterns=client._config.monitor.ignore_patterns or None, + include_patterns=client._config.monitor.include_patterns or None, ) - return await _update_document_with_chunks( - client, existing_doc, embedded_chunks, docling_document + 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 + ) + assert isinstance(doc, Document) + documents.append(doc) + return documents + + if not local_path.exists(): + raise ValueError(f"File does not exist: {local_path}") + + # Match the old _create_document_from_file behaviour: fail fast on + # unsupported extension before reading any bytes. + converter = get_converter(client._config) + if local_path.suffix.lower() not in converter.supported_extensions: + raise ValueError(f"Unsupported file extension: {local_path.suffix}") + + # Single resource — resolve the right Source adapter for this URI. + fetcher = resolve_fetcher(source_str, storage_options=storage_options) + + # The stored URI is what we look up + persist by. For an explicit uri + # override, use it as-is. Otherwise canonicalize local paths to file:// + # and leave remote URIs alone. + if uri is not None: + stored_uri = uri + elif parsed_url.scheme in ("", "file"): + stored_uri = ( + (Path(parsed_url.path) if parsed_url.scheme == "file" else Path(source_str)) + .absolute() + .as_uri() ) else: - if title is None: - title = await resolve_title( - client._config, docling_document, stored_content - ) - document = Document( - content=stored_content, - uri=uri, - title=title, - metadata=metadata, - ) - document.set_docling(docling_document) - return await _store_document_with_chunks( - client, document, embedded_chunks, docling_document - ) - - -async def _create_or_update_document_from_url( - client: "HaikuRAG", - url: str, - title: str | None = None, - metadata: dict | None = None, - uri: str | None = None, -) -> Document: - """Create or update a document from a URL by downloading and parsing the content. - - ``uri`` overrides the URL as the stored document identifier. - """ - from haiku.rag.client.processing import get_extension_from_content_type_or_url - from haiku.rag.embeddings import embed_chunks - - metadata = metadata or {} - stored_uri = uri if uri is not None else url - - converter = get_converter(client._config) - supported_extensions = converter.supported_extensions - - async with httpx.AsyncClient() as http: - response = await http.get(url) - response.raise_for_status() - - md5_hash = hashlib.md5(response.content).hexdigest() - - content_type = response.headers.get("content-type", "").lower() - - # Check if document already exists - existing_doc = await client.get_document_by_uri(stored_uri) - if existing_doc and existing_doc.metadata.get("md5") == md5_hash: - updated = False - if title is not None and title != existing_doc.title: - existing_doc.title = title - updated = True - - metadata.update({"contentType": content_type, "md5": md5_hash}) - merged_metadata = {**(existing_doc.metadata or {}), **metadata} - if merged_metadata != existing_doc.metadata: - existing_doc.metadata = merged_metadata - updated = True - - if updated: - return await client.document_repository.update(existing_doc) - return existing_doc - - file_extension = get_extension_from_content_type_or_url(url, content_type) - - if file_extension not in supported_extensions: - raise ValueError( - f"Unsupported content type/extension: {content_type}/{file_extension}" - ) - - with tempfile.NamedTemporaryFile( - mode="wb", suffix=file_extension, delete=False - ) as temp_file: - temp_file.write(response.content) - temp_file.flush() - temp_path = Path(temp_file.name) - - try: - docling_document = await client.convert(temp_path, source_uri=url) - chunks = await client.chunk(docling_document) - embedded_chunks = await embed_chunks(chunks, client._config) - finally: - temp_path.unlink(missing_ok=True) - - metadata.update({"contentType": content_type, "md5": md5_hash}) - - stored_content = docling_document.export_to_markdown() - - if existing_doc: - existing_doc.content = stored_content - existing_doc.metadata = metadata - existing_doc.set_docling(docling_document) - if title is not None: - existing_doc.title = title - elif existing_doc.title is None: - existing_doc.title = await resolve_title( - client._config, docling_document, stored_content - ) - return await _update_document_with_chunks( - client, existing_doc, embedded_chunks, docling_document - ) - else: - if title is None: - title = await resolve_title( - client._config, docling_document, stored_content - ) - document = Document( - content=stored_content, - uri=stored_uri, - title=title, - metadata=metadata, - ) - document.set_docling(docling_document) - return await _store_document_with_chunks( - client, document, embedded_chunks, docling_document - ) - - -async def _create_or_update_document_from_s3( - client: "HaikuRAG", - url: str, - *, - storage_options: dict[str, str] | None = None, - title: str | None = None, - metadata: dict | None = None, - uri: str | None = None, -) -> Document: - """Create or update a document from an s3:// URL. - - Two-stage change detection: - - HEAD ETag matches stored metadata["etag"] → skip GET and re-chunk. - - ETag differs but content MD5 matches → refresh etag only, no re-chunk. - - Otherwise download, convert, chunk, embed. - - ``uri`` overrides the s3:// URL as the stored document identifier. - """ - import obstore # type: ignore[import-not-found] - - from haiku.rag.client.processing import get_extension_from_content_type_or_url - from haiku.rag.embeddings import embed_chunks - from haiku.rag.s3 import make_s3_store - - metadata = metadata or {} - stored_uri = uri if uri is not None else url - - parsed = urlparse(url) - bucket = parsed.netloc - key = parsed.path.lstrip("/") - if not bucket or not key: - raise ValueError(f"Invalid S3 URI: {url}") - - converter = get_converter(client._config) - supported_extensions = converter.supported_extensions - - store = make_s3_store(bucket, storage_options) - - head = await obstore.head_async(store, key) - etag = (head.get("e_tag") or "").strip('"') + stored_uri = source_str existing_doc = await client.get_document_by_uri(stored_uri) - if existing_doc and existing_doc.metadata.get("etag") == etag: - updated = False - if title is not None and title != existing_doc.title: - existing_doc.title = title - updated = True - merged_metadata = {**(existing_doc.metadata or {}), **metadata} - if merged_metadata != existing_doc.metadata: - existing_doc.metadata = merged_metadata - updated = True - - if updated: - return await client.document_repository.update(existing_doc) - return existing_doc - - content_type, _ = mimetypes.guess_type(key) - if not content_type: - content_type = "application/octet-stream" - - file_extension = get_extension_from_content_type_or_url(url, content_type) - if file_extension not in supported_extensions: - raise ValueError( - f"Unsupported content type/extension: {content_type}/{file_extension}" - ) - - get_resp = await obstore.get_async(store, key) - body = await get_resp.bytes_async() - - md5_hash = hashlib.md5(body, usedforsecurity=False).hexdigest() - - with tempfile.NamedTemporaryFile( - mode="wb", suffix=file_extension, delete=False - ) as temp_file: - temp_file.write(body) - temp_file.flush() - temp_path = Path(temp_file.name) - - metadata.update({"contentType": content_type, "md5": md5_hash, "etag": etag}) - - if existing_doc and existing_doc.metadata.get("md5") == md5_hash: - temp_path.unlink(missing_ok=True) - merged_metadata = {**(existing_doc.metadata or {}), **metadata} - updated = False - if merged_metadata != existing_doc.metadata: - existing_doc.metadata = merged_metadata - updated = True - if title is not None and title != existing_doc.title: - existing_doc.title = title - updated = True - if updated: - return await client.document_repository.update(existing_doc) - return existing_doc - - try: - docling_document = await client.convert(temp_path, source_uri=url) - chunks = await client.chunk(docling_document) - embedded_chunks = await embed_chunks(chunks, client._config) - finally: - temp_path.unlink(missing_ok=True) - - stored_content = docling_document.export_to_markdown() - - if existing_doc: - existing_doc.content = stored_content - existing_doc.metadata = metadata - existing_doc.set_docling(docling_document) - if title is not None: - existing_doc.title = title - elif existing_doc.title is None: - existing_doc.title = await resolve_title( - client._config, docling_document, stored_content + # Cheap revision-based short-circuit: only worth a HEAD when we have a + # stored revision to compare against. S3 doc metadata persists "etag"; + # FS/HTTP currently don't, so this branch is effectively S3-only today. + stored_revision = ( + (existing_doc.metadata or {}).get("etag") if existing_doc else None + ) + if existing_doc and stored_revision: + current_revision = await fetcher.head(source_str) + if current_revision == stored_revision: + return await _refresh_doc_metadata( + client, + existing_doc, + title=title, + user_metadata=metadata, + source_metadata=None, ) - return await _update_document_with_chunks( - client, existing_doc, embedded_chunks, docling_document - ) - else: - if title is None: - title = await resolve_title( - client._config, docling_document, stored_content - ) - document = Document( - content=stored_content, - uri=stored_uri, + + result = await fetcher.fetch(source_str) + + # MD5 short-circuit: the bytes are unchanged even if the revision wasn't. + # Refresh the source-derived metadata (etag may have rolled) but skip + # convert/embed/store entirely. + if existing_doc and existing_doc.metadata.get("md5") == result.content_hash: + source_meta: dict = { + "contentType": result.content_type, + "md5": result.content_hash, + **result.extra_metadata, + } + return await _refresh_doc_metadata( + client, + existing_doc, title=title, - metadata=metadata, - ) - document.set_docling(docling_document) - return await _store_document_with_chunks( - client, document, embedded_chunks, docling_document + user_metadata=metadata, + source_metadata=source_meta, ) + return await _ingest_fetch_result( + client, + result, + title=title, + user_metadata=metadata, + stored_uri=stored_uri, + existing_doc=existing_doc, + ) + async def update_document( client: "HaikuRAG", @@ -606,7 +437,6 @@ async def update_document( """ from haiku.rag.embeddings import embed_chunks - # Validate: content and docling_document are mutually exclusive if content is not None and docling_document is not None: raise ValueError( "content and docling_document are mutually exclusive. " @@ -622,11 +452,9 @@ async def update_document( if metadata is not None: existing_doc.metadata = metadata - # Only metadata/title update - no rechunking needed if content is None and chunks is None and docling_document is None: return await client.document_repository.update(existing_doc) - # Custom chunks provided - use them as-is if chunks is not None: if docling_document is not None: existing_doc.content = docling_document.export_to_markdown() @@ -638,7 +466,6 @@ async def update_document( client, existing_doc, chunks, docling_document ) - # DoclingDocument provided without chunks - chunk and embed if docling_document is not None: existing_doc.content = docling_document.export_to_markdown() existing_doc.set_docling(docling_document) @@ -649,7 +476,6 @@ async def update_document( client, existing_doc, embedded_chunks, docling_document ) - # Content provided without chunks - convert, chunk, and embed assert content is not None existing_doc.content = content converter = get_converter(client._config) diff --git a/haiku_rag_slim/haiku/rag/ingester/sources/base.py b/haiku_rag_slim/haiku/rag/ingester/sources/base.py index 5440376e..021515d7 100644 --- a/haiku_rag_slim/haiku/rag/ingester/sources/base.py +++ b/haiku_rag_slim/haiku/rag/ingester/sources/base.py @@ -1,6 +1,7 @@ from collections.abc import AsyncIterator, Mapping from datetime import datetime from enum import StrEnum +from pathlib import Path from typing import Protocol, runtime_checkable from pydantic import BaseModel, Field @@ -39,6 +40,11 @@ class FetchResult(BaseModel): content_hash: str revision: str | None = None extra_metadata: dict[str, str] = Field(default_factory=dict) + # When the body is already on disk (FSSource), points at the original + # file so the pipeline can hand it to docling without copying through a + # tempfile. Remote sources (HTTP/S3) leave this None — their bytes only + # exist in memory. + disk_path: Path | None = None @runtime_checkable @@ -47,6 +53,14 @@ class Source(Protocol): def supports(self, uri: str) -> bool: ... + async def head(self, uri: str) -> str | None: + """Return the current revision for `uri` cheaply, if the backend + supports it. Returning None means "I have no cheap revision lookup — + you'll have to fetch()". The pipeline uses this to short-circuit + re-ingest when the stored revision is unchanged. + """ + ... + async def fetch(self, uri: str) -> FetchResult: ... def discover( diff --git a/haiku_rag_slim/haiku/rag/ingester/sources/fs.py b/haiku_rag_slim/haiku/rag/ingester/sources/fs.py index b321dee7..1835187b 100644 --- a/haiku_rag_slim/haiku/rag/ingester/sources/fs.py +++ b/haiku_rag_slim/haiku/rag/ingester/sources/fs.py @@ -60,8 +60,17 @@ class FSSource: return False return True + async def head(self, uri: str) -> str | None: + path = _uri_to_path(uri).absolute() + if not path.exists(): + return None + return str(path.stat().st_mtime_ns) + async def fetch(self, uri: str) -> FetchResult: - path = _uri_to_path(uri) + # Absolute path is needed for as_uri() and matches the old + # _create_document_from_file behavior (which keyed docs on the + # absolute file:// URI). + path = _uri_to_path(uri).absolute() body = path.read_bytes() content_type, _ = mimetypes.guess_type(path.name) if content_type is None: @@ -75,6 +84,7 @@ class FSSource: content_type=content_type, content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(), revision=revision, + disk_path=path, ) async def discover( diff --git a/haiku_rag_slim/haiku/rag/ingester/sources/http.py b/haiku_rag_slim/haiku/rag/ingester/sources/http.py index d59c3c0a..d37052d2 100644 --- a/haiku_rag_slim/haiku/rag/ingester/sources/http.py +++ b/haiku_rag_slim/haiku/rag/ingester/sources/http.py @@ -48,6 +48,13 @@ class HTTPSource: def _client(self) -> httpx.AsyncClient: return httpx.AsyncClient(headers=self.headers, transport=self._transport) + async def head(self, uri: str) -> str | None: + # HTTP doesn't get a cheap revision lookup in v1: the existing + # ingestion flow always GETs and the dedup uses MD5. A HEAD-first + # optimization could land later without changing this contract — + # callers just need to tolerate the extra HEAD. + return None + async def fetch(self, uri: str) -> FetchResult: async with self._client() as http: response = await http.get(uri) diff --git a/haiku_rag_slim/haiku/rag/ingester/sources/s3.py b/haiku_rag_slim/haiku/rag/ingester/sources/s3.py index 17328b0c..87ad7c35 100644 --- a/haiku_rag_slim/haiku/rag/ingester/sources/s3.py +++ b/haiku_rag_slim/haiku/rag/ingester/sources/s3.py @@ -23,6 +23,14 @@ def _parse_s3_uri(uri: str) -> tuple[str, str]: return parsed.netloc, parsed.path.lstrip("/") +def _parse_s3_object_uri(uri: str) -> tuple[str, str]: + """Like _parse_s3_uri but rejects bucket-only URIs (no key).""" + bucket, key = _parse_s3_uri(uri) + if not key: + raise ValueError(f"Invalid S3 URI: {uri}") + return bucket, key + + class S3Source: def __init__( self, @@ -54,12 +62,22 @@ class S3Source: def supports(self, uri: str) -> bool: return uri.startswith(self.uri_prefix) + async def head(self, uri: str) -> str | None: + import obstore # type: ignore[import-not-found] + + from haiku.rag.s3 import make_s3_store + + bucket, key = _parse_s3_object_uri(uri) + store = make_s3_store(bucket, self.storage_options) + head = await obstore.head_async(store, key) + return (head.get("e_tag") or "").strip('"').strip() or None + async def fetch(self, uri: str) -> FetchResult: import obstore # type: ignore[import-not-found] from haiku.rag.s3 import make_s3_store - bucket, key = _parse_s3_uri(uri) + bucket, key = _parse_s3_object_uri(uri) store = make_s3_store(bucket, self.storage_options) head = await obstore.head_async(store, key) diff --git a/tests/ingester/test_fs_source.py b/tests/ingester/test_fs_source.py index 7c28d691..07ad2bb2 100644 --- a/tests/ingester/test_fs_source.py +++ b/tests/ingester/test_fs_source.py @@ -47,6 +47,20 @@ async def test_fs_source_fetch_returns_bytes_and_md5(fs_root: Path): ) assert result.content_type == "text/markdown" assert result.revision == str(target.stat().st_mtime_ns) + assert result.disk_path == target + + +@pytest.mark.asyncio +async def test_fs_source_head_returns_mtime(fs_root: Path): + src = FSSource(root=fs_root) + target = fs_root / "a.md" + assert await src.head(target.as_uri()) == str(target.stat().st_mtime_ns) + + +@pytest.mark.asyncio +async def test_fs_source_head_returns_none_for_missing_file(fs_root: Path): + src = FSSource(root=fs_root) + assert await src.head((fs_root / "missing.md").as_uri()) is None @pytest.mark.asyncio diff --git a/tests/ingester/test_http_source.py b/tests/ingester/test_http_source.py index c3a96f31..088fcbf3 100644 --- a/tests/ingester/test_http_source.py +++ b/tests/ingester/test_http_source.py @@ -29,6 +29,13 @@ def test_source_id_is_user_provided(): assert HTTPSource(source_id="arxiv").source_id == "arxiv" +@pytest.mark.asyncio +async def test_head_returns_none(): + # HTTP has no cheap revision lookup in v1 — the pipeline always GETs. + src = HTTPSource(source_id="default") + assert await src.head("https://example.com/a.md") is None + + @pytest.mark.asyncio async def test_fetch_returns_bytes_and_md5_and_etag(): body = b"hello world" diff --git a/tests/ingester/test_s3_source.py b/tests/ingester/test_s3_source.py index cbf79fa2..cb7cb32a 100644 --- a/tests/ingester/test_s3_source.py +++ b/tests/ingester/test_s3_source.py @@ -75,6 +75,23 @@ def test_invalid_uri_raises(): S3Source(uri="s3:///key") +@pytest.mark.asyncio +async def test_head_returns_etag(fake_obstore_io): + head_async, _ = fake_obstore_io + head_async.return_value = {"e_tag": '"abc123"'} + src = S3Source(uri="s3://bucket/") + assert await src.head("s3://bucket/file.txt") == "abc123" + head_async.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_head_returns_none_when_no_etag(fake_obstore_io): + head_async, _ = fake_obstore_io + head_async.return_value = {} + src = S3Source(uri="s3://bucket/") + assert await src.head("s3://bucket/file.txt") is None + + @pytest.mark.asyncio async def test_fetch_returns_bytes_md5_etag(fake_obstore_io): head_async, get_async = fake_obstore_io diff --git a/tests/ingester/test_sources_base.py b/tests/ingester/test_sources_base.py index 34f82516..b2ffe995 100644 --- a/tests/ingester/test_sources_base.py +++ b/tests/ingester/test_sources_base.py @@ -59,6 +59,9 @@ def test_source_protocol_runtime_checkable(): def supports(self, uri: str) -> bool: return True + async def head(self, uri: str): + return None + async def fetch(self, uri: str): raise NotImplementedError