From d3a1011baf032fca6379c6c6b10d5d9b55162c56 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Mon, 22 Jun 2026 10:37:48 +0300 Subject: [PATCH] write fetched bodies off the event loop --- haiku_rag_slim/haiku/rag/client/documents.py | 22 ++++--- tests/test_client.py | 66 +++++++++++++++++++- 2 files changed, 80 insertions(+), 8 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/client/documents.py b/haiku_rag_slim/haiku/rag/client/documents.py index 76aa2757..7e20181f 100644 --- a/haiku_rag_slim/haiku/rag/client/documents.py +++ b/haiku_rag_slim/haiku/rag/client/documents.py @@ -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: """SQL `WHERE` clause matching documents whose ``metadata.parent_uri`` equals ``parent_uri``. ``metadata`` is stored as a JSON string produced by @@ -426,13 +439,8 @@ async def _ingest_fetch_result( 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 + target_path = await _write_fetch_body(result.body, file_extension) + cleanup_path = target_path try: with logfire.span("document.convert", uri=result.uri): diff --git a/tests/test_client.py b/tests/test_client.py index b2ebc8a2..857a1bff 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -1,13 +1,20 @@ import json import tempfile +import threading from pathlib import Path from unittest.mock import AsyncMock, patch import httpx 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.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.store.compression import decompress_json from haiku.rag.store.models.chunk import Chunk @@ -19,6 +26,63 @@ def vcr_cassette_dir(): 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() async def test_client_document_crud(qa_corpus: list[dict[str, str]], temp_db_path): """Test HaikuRAG CRUD operations for documents."""