write fetched bodies off the event loop
This commit is contained in:
parent
183d595494
commit
d3a1011baf
2 changed files with 80 additions and 8 deletions
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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."""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue