route create_document_from_source through source adapters

This commit is contained in:
Yiorgis Gozadinos 2026-05-21 16:56:27 +03:00
parent b9637cd625
commit 717fc3320b
No known key found for this signature in database
9 changed files with 295 additions and 379 deletions

View file

@ -1,15 +1,15 @@
import hashlib
import mimetypes
import tempfile import tempfile
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from urllib.parse import urlparse 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.client.titles import resolve_title
from haiku.rag.converters import get_converter 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.chunk import Chunk
from haiku.rag.store.models.document import Document from haiku.rag.store.models.document import Document
from haiku.rag.store.models.document_item import extract_items 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. Handles versioning/rollback on failure.
""" """
# Ensure all chunks have embeddings before storing
chunks = await ensure_chunks_embedded(client._config, chunks) chunks = await ensure_chunks_embedded(client._config, chunks)
# Snapshot table versions for versioned rollback (if supported)
versions = await client.store.current_table_versions() versions = await client.store.current_table_versions()
# Create the document
created_doc = await client.document_repository.create(document) created_doc = await client.document_repository.create(document)
try: try:
assert created_doc.id is not None, ( assert created_doc.id is not None, (
"Document ID should not be None after creation" "Document ID should not be None after creation"
) )
# Set document_id and order for all chunks
for order, chunk in enumerate(chunks): for order, chunk in enumerate(chunks):
chunk.document_id = created_doc.id chunk.document_id = created_doc.id
chunk.order = order chunk.order = order
# Batch create all chunks in a single operation
await client.chunk_repository.create(chunks) await client.chunk_repository.create(chunks)
# Extract and store document items for context expansion
items = extract_items(created_doc.id, docling_document) items = extract_items(created_doc.id, docling_document)
await client.document_item_repository.create_items(created_doc.id, items) 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: if client._config.storage.auto_vacuum:
client._schedule_vacuum() client._schedule_vacuum()
return created_doc return created_doc
except Exception: except Exception:
# Roll back to the captured versions and re-raise
await client.store.restore_table_versions(versions) await client.store.restore_table_versions(versions)
raise raise
@ -92,7 +84,6 @@ async def _update_document_with_chunks(
versions = await client.store.current_table_versions() versions = await client.store.current_table_versions()
# Delete existing chunks before writing new ones
await client.chunk_repository.delete_by_document_id(document.id) await client.chunk_repository.delete_by_document_id(document.id)
try: try:
@ -137,14 +128,11 @@ async def create_document(
""" """
from haiku.rag.embeddings import embed_chunks from haiku.rag.embeddings import embed_chunks
# Convert → Chunk → Embed using primitives
converter = get_converter(client._config) converter = get_converter(client._config)
docling_document = await converter.convert_text(content, format=format) docling_document = await converter.convert_text(content, format=format)
chunks = await client.chunk(docling_document) chunks = await client.chunk(docling_document)
embedded_chunks = await embed_chunks(chunks, client._config) 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() stored_content = docling_document.export_to_markdown()
if title is None: if title is None:
@ -191,6 +179,112 @@ async def import_document(
return await _store_document_with_chunks(client, document, chunks, docling_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( async def create_document_from_source(
client: "HaikuRAG", client: "HaikuRAG",
source: str | Path, 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 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 (which is normally ``file://`` for local files or the URL for remote
sources). This is useful when callers want to persist documents under a sources). Not supported for directory sources, which produce one document
logical identifier (e.g. an ArXiv ID) rather than the on-disk path. Not per file.
supported for directory sources, which produce one document per file.
Returns a single Document for files/URLs, a list for directories. 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) source_str = str(source)
parsed_url = urlparse(source_str) 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(): # Directory case: recurse with the existing FS filter and produce one
if uri is not None: # document per file. Remote schemes (http/s3) never hit this branch.
raise ValueError( if parsed_url.scheme in ("", "file"):
"uri override is not supported for directory sources; each file " local_path = (
"produces its own document with its own auto-derived URI." Path(parsed_url.path)
) if parsed_url.scheme == "file"
from haiku.rag.monitor import FileFilter else (Path(source) if isinstance(source, str) else source)
documents = []
filter = FileFilter(
ignore_patterns=client._config.monitor.ignore_patterns or None,
include_patterns=client._config.monitor.include_patterns or None,
) )
for path in source_path.rglob("*"): if local_path.is_dir():
if path.is_file() and filter.include_file(str(path)): if uri is not None:
doc = await _create_document_from_file( raise ValueError(
client, path, title=None, metadata=metadata "uri override is not supported for directory sources; each file "
"produces its own document with its own auto-derived URI."
) )
documents.append(doc) from haiku.rag.ingester.sources.filter import FileFilter
return documents
return await _create_document_from_file( documents: list[Document] = []
client, source_path, title=title, metadata=metadata, uri=uri filter = FileFilter(
) ignore_patterns=client._config.monitor.ignore_patterns or None,
include_patterns=client._config.monitor.include_patterns or None,
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
) )
return await _update_document_with_chunks( for child in local_path.rglob("*"):
client, existing_doc, embedded_chunks, docling_document 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: else:
if title is None: stored_uri = source_str
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('"')
existing_doc = await client.get_document_by_uri(stored_uri) 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} # Cheap revision-based short-circuit: only worth a HEAD when we have a
if merged_metadata != existing_doc.metadata: # stored revision to compare against. S3 doc metadata persists "etag";
existing_doc.metadata = merged_metadata # FS/HTTP currently don't, so this branch is effectively S3-only today.
updated = True stored_revision = (
(existing_doc.metadata or {}).get("etag") if existing_doc else None
if updated: )
return await client.document_repository.update(existing_doc) if existing_doc and stored_revision:
return existing_doc current_revision = await fetcher.head(source_str)
if current_revision == stored_revision:
content_type, _ = mimetypes.guess_type(key) return await _refresh_doc_metadata(
if not content_type: client,
content_type = "application/octet-stream" existing_doc,
title=title,
file_extension = get_extension_from_content_type_or_url(url, content_type) user_metadata=metadata,
if file_extension not in supported_extensions: source_metadata=None,
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
) )
return await _update_document_with_chunks(
client, existing_doc, embedded_chunks, docling_document result = await fetcher.fetch(source_str)
)
else: # MD5 short-circuit: the bytes are unchanged even if the revision wasn't.
if title is None: # Refresh the source-derived metadata (etag may have rolled) but skip
title = await resolve_title( # convert/embed/store entirely.
client._config, docling_document, stored_content if existing_doc and existing_doc.metadata.get("md5") == result.content_hash:
) source_meta: dict = {
document = Document( "contentType": result.content_type,
content=stored_content, "md5": result.content_hash,
uri=stored_uri, **result.extra_metadata,
}
return await _refresh_doc_metadata(
client,
existing_doc,
title=title, title=title,
metadata=metadata, user_metadata=metadata,
) source_metadata=source_meta,
document.set_docling(docling_document)
return await _store_document_with_chunks(
client, document, embedded_chunks, docling_document
) )
return await _ingest_fetch_result(
client,
result,
title=title,
user_metadata=metadata,
stored_uri=stored_uri,
existing_doc=existing_doc,
)
async def update_document( async def update_document(
client: "HaikuRAG", client: "HaikuRAG",
@ -606,7 +437,6 @@ async def update_document(
""" """
from haiku.rag.embeddings import embed_chunks 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: if content is not None and docling_document is not None:
raise ValueError( raise ValueError(
"content and docling_document are mutually exclusive. " "content and docling_document are mutually exclusive. "
@ -622,11 +452,9 @@ async def update_document(
if metadata is not None: if metadata is not None:
existing_doc.metadata = metadata existing_doc.metadata = metadata
# Only metadata/title update - no rechunking needed
if content is None and chunks is None and docling_document is None: if content is None and chunks is None and docling_document is None:
return await client.document_repository.update(existing_doc) return await client.document_repository.update(existing_doc)
# Custom chunks provided - use them as-is
if chunks is not None: if chunks is not None:
if docling_document is not None: if docling_document is not None:
existing_doc.content = docling_document.export_to_markdown() existing_doc.content = docling_document.export_to_markdown()
@ -638,7 +466,6 @@ async def update_document(
client, existing_doc, chunks, docling_document client, existing_doc, chunks, docling_document
) )
# DoclingDocument provided without chunks - chunk and embed
if docling_document is not None: if docling_document is not None:
existing_doc.content = docling_document.export_to_markdown() existing_doc.content = docling_document.export_to_markdown()
existing_doc.set_docling(docling_document) existing_doc.set_docling(docling_document)
@ -649,7 +476,6 @@ async def update_document(
client, existing_doc, embedded_chunks, docling_document client, existing_doc, embedded_chunks, docling_document
) )
# Content provided without chunks - convert, chunk, and embed
assert content is not None assert content is not None
existing_doc.content = content existing_doc.content = content
converter = get_converter(client._config) converter = get_converter(client._config)

View file

@ -1,6 +1,7 @@
from collections.abc import AsyncIterator, Mapping from collections.abc import AsyncIterator, Mapping
from datetime import datetime from datetime import datetime
from enum import StrEnum from enum import StrEnum
from pathlib import Path
from typing import Protocol, runtime_checkable from typing import Protocol, runtime_checkable
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
@ -39,6 +40,11 @@ class FetchResult(BaseModel):
content_hash: str content_hash: str
revision: str | None = None revision: str | None = None
extra_metadata: dict[str, str] = Field(default_factory=dict) 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 @runtime_checkable
@ -47,6 +53,14 @@ class Source(Protocol):
def supports(self, uri: str) -> bool: ... 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: ... async def fetch(self, uri: str) -> FetchResult: ...
def discover( def discover(

View file

@ -60,8 +60,17 @@ class FSSource:
return False return False
return True 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: 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() body = path.read_bytes()
content_type, _ = mimetypes.guess_type(path.name) content_type, _ = mimetypes.guess_type(path.name)
if content_type is None: if content_type is None:
@ -75,6 +84,7 @@ class FSSource:
content_type=content_type, content_type=content_type,
content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(), content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(),
revision=revision, revision=revision,
disk_path=path,
) )
async def discover( async def discover(

View file

@ -48,6 +48,13 @@ class HTTPSource:
def _client(self) -> httpx.AsyncClient: def _client(self) -> httpx.AsyncClient:
return httpx.AsyncClient(headers=self.headers, transport=self._transport) 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 def fetch(self, uri: str) -> FetchResult:
async with self._client() as http: async with self._client() as http:
response = await http.get(uri) response = await http.get(uri)

View file

@ -23,6 +23,14 @@ def _parse_s3_uri(uri: str) -> tuple[str, str]:
return parsed.netloc, parsed.path.lstrip("/") 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: class S3Source:
def __init__( def __init__(
self, self,
@ -54,12 +62,22 @@ class S3Source:
def supports(self, uri: str) -> bool: def supports(self, uri: str) -> bool:
return uri.startswith(self.uri_prefix) 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: async def fetch(self, uri: str) -> FetchResult:
import obstore # type: ignore[import-not-found] import obstore # type: ignore[import-not-found]
from haiku.rag.s3 import make_s3_store 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) store = make_s3_store(bucket, self.storage_options)
head = await obstore.head_async(store, key) head = await obstore.head_async(store, key)

View file

@ -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.content_type == "text/markdown"
assert result.revision == str(target.stat().st_mtime_ns) 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 @pytest.mark.asyncio

View file

@ -29,6 +29,13 @@ def test_source_id_is_user_provided():
assert HTTPSource(source_id="arxiv").source_id == "arxiv" 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 @pytest.mark.asyncio
async def test_fetch_returns_bytes_and_md5_and_etag(): async def test_fetch_returns_bytes_and_md5_and_etag():
body = b"hello world" body = b"hello world"

View file

@ -75,6 +75,23 @@ def test_invalid_uri_raises():
S3Source(uri="s3:///key") 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 @pytest.mark.asyncio
async def test_fetch_returns_bytes_md5_etag(fake_obstore_io): async def test_fetch_returns_bytes_md5_etag(fake_obstore_io):
head_async, get_async = fake_obstore_io head_async, get_async = fake_obstore_io

View file

@ -59,6 +59,9 @@ def test_source_protocol_runtime_checkable():
def supports(self, uri: str) -> bool: def supports(self, uri: str) -> bool:
return True return True
async def head(self, uri: str):
return None
async def fetch(self, uri: str): async def fetch(self, uri: str):
raise NotImplementedError raise NotImplementedError