route create_document_from_source through source adapters
This commit is contained in:
parent
b9637cd625
commit
717fc3320b
9 changed files with 295 additions and 379 deletions
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue