write fetched bodies off the event loop

This commit is contained in:
Yiorgis Gozadinos 2026-06-22 10:37:48 +03:00
parent 183d595494
commit d3a1011baf
No known key found for this signature in database
2 changed files with 80 additions and 8 deletions

View file

@ -84,6 +84,19 @@ async def _prepare_document_from_docling(
) )
def _write_fetch_body_sync(body: bytes, suffix: str) -> Path:
with tempfile.NamedTemporaryFile(
mode="wb", suffix=suffix, delete=False
) as temp_file:
temp_file.write(body)
temp_file.flush()
return Path(temp_file.name)
async def _write_fetch_body(body: bytes, suffix: str) -> Path:
return await asyncio.to_thread(_write_fetch_body_sync, body, suffix)
def parent_uri_filter(parent_uri: str) -> str: def parent_uri_filter(parent_uri: str) -> str:
"""SQL `WHERE` clause matching documents whose ``metadata.parent_uri`` """SQL `WHERE` clause matching documents whose ``metadata.parent_uri``
equals ``parent_uri``. ``metadata`` is stored as a JSON string produced by equals ``parent_uri``. ``metadata`` is stored as a JSON string produced by
@ -426,12 +439,7 @@ async def _ingest_fetch_result(
target_path = result.disk_path target_path = result.disk_path
cleanup_path: Path | None = None cleanup_path: Path | None = None
else: else:
with tempfile.NamedTemporaryFile( target_path = await _write_fetch_body(result.body, file_extension)
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 cleanup_path = target_path
try: try:

View file

@ -1,13 +1,20 @@
import json import json
import tempfile import tempfile
import threading
from pathlib import Path from pathlib import Path
from unittest.mock import AsyncMock, patch from unittest.mock import AsyncMock, patch
import httpx import httpx
import pytest import pytest
from docling_core.types.doc.document import DoclingDocument
from docling_core.types.doc.labels import DocItemLabel
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.client.documents import DocumentImport from haiku.rag.client.documents import (
DocumentImport,
_prepare_document_from_docling,
_write_fetch_body,
)
from haiku.rag.config import Config from haiku.rag.config import Config
from haiku.rag.store.compression import decompress_json from haiku.rag.store.compression import decompress_json
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
@ -19,6 +26,63 @@ def vcr_cassette_dir():
return str(Path(__file__).parent / "cassettes" / "test_client") return str(Path(__file__).parent / "cassettes" / "test_client")
@pytest.mark.asyncio
async def test_prepare_document_from_docling_runs_off_event_loop_thread(monkeypatch):
import haiku.rag.client.documents as documents
event_loop_thread = threading.current_thread()
called_from: list[threading.Thread] = []
docling_doc = DoclingDocument(name="thread-check")
docling_doc.add_text(label=DocItemLabel.TEXT, text="Threaded content")
document = Document(content="")
original = documents._prepare_document_from_docling_sync
def spy(doc, docling):
called_from.append(threading.current_thread())
return original(doc, docling)
monkeypatch.setattr(documents, "_prepare_document_from_docling_sync", spy)
content = await _prepare_document_from_docling(document, docling_doc)
assert content == "Threaded content"
assert document.content == "Threaded content"
assert document.docling_document is not None
assert called_from, "Document.set_docling was never called"
assert called_from[0] is not event_loop_thread, (
"Document.set_docling ran on the event-loop thread; document prep must "
"be dispatched via asyncio.to_thread"
)
@pytest.mark.asyncio
async def test_write_fetch_body_runs_off_event_loop_thread(monkeypatch):
import haiku.rag.client.documents as documents
event_loop_thread = threading.current_thread()
called_from: list[threading.Thread] = []
original = documents._write_fetch_body_sync
def spy(body, suffix):
called_from.append(threading.current_thread())
return original(body, suffix)
monkeypatch.setattr(documents, "_write_fetch_body_sync", spy)
path = await _write_fetch_body(b"payload", ".bin")
try:
assert path.read_bytes() == b"payload"
finally:
path.unlink(missing_ok=True)
assert called_from, "_write_fetch_body_sync was never called"
assert called_from[0] is not event_loop_thread, (
"_write_fetch_body_sync ran on the event-loop thread; fetched body "
"writes must be dispatched via asyncio.to_thread"
)
@pytest.mark.vcr() @pytest.mark.vcr()
async def test_client_document_crud(qa_corpus: list[dict[str, str]], temp_db_path): async def test_client_document_crud(qa_corpus: list[dict[str, str]], temp_db_path):
"""Test HaikuRAG CRUD operations for documents.""" """Test HaikuRAG CRUD operations for documents."""