Cover MCP, ingester and converter error paths

Add tests for the MCP tools' degradation contracts, malformed WebDAV
multistatus bodies, dry-run poller sweeps including the circuit-open and
discover-failure paths, FS source scheme and symlink handling, docling-serve
zip parsing, and the remaining embedding and reranker helpers. Parametrize
_strip_etag.

Drop the misplaced pragma on the analyze handler, which sat on the return and
left the except uncovered. Add one on the FS symlink OSError guard, which
resolve(strict=False) absorbs for every real link.
This commit is contained in:
Yiorgis Gozadinos 2026-07-26 19:14:26 +03:00
parent dddaf0f84c
commit 1e8e5e9f6f
No known key found for this signature in database
13 changed files with 850 additions and 22 deletions

View file

@ -143,7 +143,9 @@ class FSSource:
if path.is_symlink():
try:
resolved = path.resolve(strict=False)
except OSError:
except OSError: # pragma: no cover - strict=False absorbs
# symlink cycles and missing targets, so no real link
# reaches this; kept as a guard against platform drift.
continue
if not resolved.is_relative_to(self.root):
continue

View file

@ -244,6 +244,6 @@ def create_mcp_server(
result = await rag.analyze(question, filter=filter, images=images)
return result.answer
except Exception as e:
return f"Error running analysis capability: {e!s}" # pragma: no cover
return f"Error running analysis capability: {e!s}"
return mcp

View file

@ -330,3 +330,41 @@ async def test_fs_source_fetch_reads_off_event_loop_thread(fs_root: Path):
"FSSource._read_body ran on the event-loop thread; the read+hash must "
"be dispatched via asyncio.to_thread"
)
@pytest.mark.asyncio
async def test_fetch_rejects_foreign_scheme(tmp_path):
"""`supports()` short-circuits on scheme, but fetch/head resolve directly,
so the unsupported-scheme path must be handled there too."""
src = FSSource(root=tmp_path, supported_extensions=[".md"], source_id="local")
with pytest.raises(UnsupportedSourceError):
await src.fetch("s3://bucket/key.md")
assert await src.head("s3://bucket/key.md") is None
@pytest.mark.asyncio
async def test_fetch_falls_back_to_octet_stream_for_unknown_extension(tmp_path):
target = tmp_path / "data.unknownext"
target.write_bytes(b"payload")
src = FSSource(
root=tmp_path, supported_extensions=[".unknownext"], source_id="local"
)
result = await src.fetch(target.as_uri())
assert result.content_type == "application/octet-stream"
assert result.body == b"payload"
@pytest.mark.asyncio
async def test_discover_skips_symlink_to_missing_in_root_target(tmp_path):
"""A broken symlink inside the root resolves to a path that is not a file."""
(tmp_path / "real.md").write_text("real")
(tmp_path / "broken.md").symlink_to(tmp_path / "absent.md")
src = FSSource(root=tmp_path, supported_extensions=[".md"], source_id="local")
events = [e async for e in src.discover()]
assert {e.uri for e in events} == {(tmp_path / "real.md").as_uri()}

View file

@ -422,3 +422,14 @@ async def test_fetch_skips_head_when_no_max_size():
)
await src.fetch("https://example.com/a.md")
assert calls == ["GET"]
@pytest.mark.asyncio
async def test_aclose_closes_the_http_client():
src = HTTPSource(
source_id="urls",
urls=[],
transport=httpx.MockTransport(lambda r: httpx.Response(200)),
)
await src.aclose()
assert src._http.is_closed

View file

@ -602,3 +602,159 @@ async def test_fs_poller_enqueues_initial_files(tmp_path, jobs, sync):
queued = await jobs.list_jobs(source_id="local")
assert {Path(j.uri).name for j in queued} == {"a.md", "b.md"}
assert all(j.status is JobStatus.QUEUED for j in queued)
# --- _dry_run_once ---
@pytest.mark.asyncio
async def test_dry_run_collects_changes_without_writing(fs_config, jobs, sync):
source = _StubSource(
"src",
[
[
_event("file:///a.md"),
_event("file:///b.md", kind=SourceEventKind.UNCHANGED),
_event("file:///c.md", kind=SourceEventKind.DELETE),
]
],
)
poller = _periodic(source, fs_config, jobs, sync)
ok, summary, changes = await poller._dry_run_once()
assert ok is True
assert summary.upsert_count == 1
assert summary.unchanged_count == 1
assert summary.delete_count == 1
assert {c.op for c in changes} == {JobOp.UPSERT, JobOp.DELETE}
# A dry run must not touch the queue.
assert await jobs.list_jobs(source_id="src") == []
@pytest.mark.asyncio
async def test_dry_run_ignores_deletes_when_delete_orphans_false(
fs_config, jobs, sync, tmp_path
):
config = FSSourceConfig(
type="fs",
id="src",
root=tmp_path,
delete_orphans=False,
poll_interval_s=0.05,
)
source = _StubSource("src", [[_event("file:///c.md", kind=SourceEventKind.DELETE)]])
poller = _periodic(source, config, jobs, sync)
ok, summary, changes = await poller._dry_run_once()
assert ok is True
assert summary.delete_count == 0
assert summary.ignored_delete_count == 1
assert changes == []
@pytest.mark.asyncio
async def test_dry_run_skipped_when_circuit_open(fs_config, jobs, sync):
class _Clock:
now = 0.0
def __call__(self):
return self.now
breaker = CircuitBreaker(
CircuitBreakerConfig(failure_threshold=1, cooldown_s=30.0),
now_fn=_Clock(),
)
source = _StubSource("src", [])
source.fail_with = RuntimeError("upstream down")
poller = _periodic(source, fs_config, jobs, sync, breaker=breaker)
assert await poller._sweep_once() is False
assert breaker.is_open is True
before = source.discover_calls
ok, summary, changes = await poller._dry_run_once()
assert ok is False
assert changes == []
assert source.discover_calls == before
assert poller.last_skip_reason == "circuit_open"
@pytest.mark.asyncio
async def test_dry_run_records_failure_when_discover_raises(fs_config, jobs, sync):
source = _StubSource("src", [])
source.fail_with = RuntimeError("upstream down")
poller = _periodic(source, fs_config, jobs, sync)
ok, summary, changes = await poller._dry_run_once()
assert ok is False
assert changes == []
assert poller._breaker.consecutive_failures == 1
@pytest.mark.asyncio
async def test_dry_run_skipped_when_queue_has_pending_work(fs_config, jobs, sync):
source = _StubSource("src", [[_event("file:///a.md")]])
poller = _periodic(source, fs_config, jobs, sync)
await jobs.enqueue("src", "file:///pending.md", JobOp.UPSERT)
ok, _summary, changes = await poller._dry_run_once()
assert ok is False
assert changes == []
assert poller.last_skip_reason == "pending_work"
@pytest.mark.asyncio
async def test_watch_deleted_skipped_when_delete_orphans_false(tmp_path, jobs, sync):
from watchfiles import Change
from haiku.rag.ingester.pollers.fs import FSPoller
from haiku.rag.ingester.sources.fs import FSSource
cfg = FSSourceConfig(
type="fs",
id="local",
root=tmp_path,
delete_orphans=False,
poll_interval_s=60.0,
)
poller = FSPoller(
source=FSSource(root=tmp_path, supported_extensions=[".md"], source_id="local"),
config=cfg,
job_repo=jobs,
sync_repo=sync,
)
await poller._handle_watch_change(Change.deleted, tmp_path / "gone.md")
assert await jobs.list_jobs(source_id="local") == []
@pytest.mark.asyncio
async def test_dry_run_manifest_reports_failed_sources(tmp_path, jobs, sync):
"""A source whose discover() raises is named in the failed list while the
manifest still carries the sources that succeeded."""
manager = PollerManager(
configs=[FSSourceConfig(type="fs", id="ok", root=tmp_path)],
job_repo=jobs,
sync_repo=sync,
)
broken = _StubSource("broken", [])
broken.fail_with = RuntimeError("upstream down")
manager._pollers.append(
_periodic(
broken,
FSSourceConfig(type="fs", id="broken", root=tmp_path, poll_interval_s=60.0),
jobs,
sync,
)
)
manifest, failed = await manager.dry_run_manifest()
assert failed == ["broken"]
assert {s.source_id for s in manifest.sources} == {"ok", "broken"}

View file

@ -7,26 +7,29 @@ from haiku.rag.ingester.sources.base import FileTooLargeError, SourceEventKind
from haiku.rag.ingester.sources.webdav import WebDAVSource, _strip_etag
def test_strip_etag_strong_quoted():
assert _strip_etag('"abc123"') == "abc123"
def test_strip_etag_weak_marker():
assert _strip_etag('W/"abc123"') == "abc123"
def test_strip_etag_unquoted():
assert _strip_etag("abc123") == "abc123"
def test_strip_etag_whitespace():
assert _strip_etag(' W/"abc" ') == "abc"
def test_strip_etag_empty_returns_none():
assert _strip_etag("") is None
assert _strip_etag('""') is None
assert _strip_etag(None) is None
@pytest.mark.parametrize(
"raw,expected",
[
('"abc123"', "abc123"),
('W/"abc123"', "abc123"),
("abc123", "abc123"),
(' W/"abc" ', "abc"),
("", None),
('""', None),
(None, None),
],
ids=[
"strong_quoted",
"weak_marker",
"unquoted",
"whitespace",
"empty",
"empty_quotes",
"none",
],
)
def test_strip_etag(raw, expected):
assert _strip_etag(raw) == expected
def _transport(handler) -> httpx.MockTransport:
@ -628,3 +631,131 @@ async def test_fetch_skips_head_when_no_max_size():
)
await src.fetch("https://nc.example.com/dav/a.txt")
assert calls == ["GET"]
# Malformed multistatus bodies: a <response> that can't be decoded is dropped
# rather than aborting the whole listing.
def _raw_multistatus(*response_blocks: str) -> bytes:
body = ['<?xml version="1.0" encoding="utf-8"?>', '<d:multistatus xmlns:d="DAV:">']
body.extend(response_blocks)
body.append("</d:multistatus>")
return "\n".join(body).encode()
_NO_HREF = """ <d:response>
<d:propstat>
<d:status>HTTP/1.1 200 OK</d:status>
<d:prop><d:getetag>"r"</d:getetag></d:prop>
</d:propstat>
</d:response>"""
_EMPTY_HREF = """ <d:response>
<d:href></d:href>
<d:propstat>
<d:status>HTTP/1.1 200 OK</d:status>
<d:prop><d:getetag>"r"</d:getetag></d:prop>
</d:propstat>
</d:response>"""
_NO_STATUS = """ <d:response>
<d:href>/dav/a.md</d:href>
<d:propstat>
<d:prop><d:getetag>"r"</d:getetag></d:prop>
</d:propstat>
</d:response>"""
_NOT_FOUND_STATUS = """ <d:response>
<d:href>/dav/a.md</d:href>
<d:propstat>
<d:status>HTTP/1.1 404 Not Found</d:status>
<d:prop><d:getetag>"r"</d:getetag></d:prop>
</d:propstat>
</d:response>"""
_STATUS_WITHOUT_PROP = """ <d:response>
<d:href>/dav/a.md</d:href>
<d:propstat>
<d:status>HTTP/1.1 200 OK</d:status>
</d:propstat>
</d:response>"""
@pytest.mark.asyncio
@pytest.mark.parametrize(
"block",
[_NO_HREF, _EMPTY_HREF, _NO_STATUS, _NOT_FOUND_STATUS],
ids=["no_href", "empty_href", "propstat_without_status", "propstat_404"],
)
async def test_head_returns_none_for_undecodable_response(block):
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(207, content=_raw_multistatus(block))
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(handler),
)
assert await src.head("https://nc.example.com/dav/a.md") is None
@pytest.mark.asyncio
async def test_head_returns_none_for_empty_multistatus():
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(207, content=_raw_multistatus())
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(handler),
)
assert await src.head("https://nc.example.com/dav/a.md") is None
@pytest.mark.asyncio
async def test_entry_with_status_but_no_prop_has_no_revision():
"""A 200 propstat carrying no <prop> still yields an entry, without a revision."""
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(207, content=_raw_multistatus(_STATUS_WITHOUT_PROP))
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(handler),
)
assert await src.head("https://nc.example.com/dav/a.md") is None
@pytest.mark.asyncio
async def test_discover_skips_base_url_reported_as_file():
"""Broken servers list the base URL itself as a non-collection; it and any
href outside the base are skipped."""
body = _multistatus(
{"href": "/dav/", "etag": '"base"'},
{"href": "/outside/x.md", "etag": '"out"'},
{"href": "/dav/keep.md", "etag": '"keep"'},
)
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(207, content=body)
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(handler),
)
events = [event async for event in src.discover()]
assert {e.uri for e in events} == {"https://nc.example.com/dav/keep.md"}
@pytest.mark.asyncio
async def test_aclose_closes_the_http_client():
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(lambda r: httpx.Response(200)),
)
await src.aclose()
assert src._http.is_closed

View file

@ -798,3 +798,78 @@ async def test_serve_chunker_accepts_picture_laden_docling(doclaynet_first_page_
chunks = await serve_chunker.chunk(doc)
assert len(chunks) > 0, "docling-serve chunker returned 0 chunks"
class TestDoclingServeChunkerRefResolution:
"""_resolve_label_from_document and the dict-shaped doc_items branch."""
@pytest.fixture
def chunker(self):
config = AppConfig()
config.providers.docling_serve.base_url = "http://localhost:5001"
config.processing.chunk_size = 256
config.processing.chunking_tokenizer = "Qwen/Qwen3-Embedding-0.6B"
return DoclingServeChunker(config)
@pytest.fixture
def document(self):
from docling_core.types.doc.document import DoclingDocument
return DoclingDocument.model_validate(
{
"name": "doc",
"texts": [
{
"self_ref": "#/texts/0",
"text": "body",
"orig": "body",
"label": "paragraph",
}
],
"tables": [],
"pictures": [],
"groups": [],
"body": {"self_ref": "#/body", "children": []},
"furniture": {"self_ref": "#/furniture", "children": []},
}
)
@pytest.mark.parametrize(
"ref",
["not-a-ref", "#/texts/999", "#/nope/0"],
ids=["unparseable", "index_out_of_range", "unknown_collection"],
)
def test_unresolvable_ref_yields_no_label(self, document, ref):
from haiku.rag.chunkers.docling_serve import _resolve_label_from_document
assert _resolve_label_from_document(ref, document) is None
def test_resolvable_ref_yields_label(self, document):
from haiku.rag.chunkers.docling_serve import _resolve_label_from_document
assert _resolve_label_from_document("#/texts/0", document) == "paragraph"
@pytest.mark.asyncio
async def test_chunk_of_none_returns_empty(self, chunker):
assert await chunker.chunk(None) == []
@pytest.mark.asyncio
async def test_dict_shaped_doc_items_are_decoded(self, chunker, document):
"""docling-serve returns refs as strings today; the dict shape is
accepted in case the API changes."""
async def fake_chunk_api(_document):
return [
{
"raw_text": "body",
"doc_items": [{"self_ref": "#/texts/0", "label": "paragraph"}],
}
]
chunker._call_chunk_api = fake_chunk_api # type: ignore[method-assign]
chunks = await chunker.chunk(document)
assert len(chunks) == 1
assert chunks[0].metadata["doc_item_refs"] == ["#/texts/0"]
assert chunks[0].metadata["labels"] == ["paragraph"]

View file

@ -405,6 +405,46 @@ class TestDoclingLocalConverter:
assert isinstance(doc, DoclingDocument)
assert doc.name == "test"
@pytest.mark.asyncio
async def test_convert_file_reads_unknown_extension_as_text(
self, converter, tmp_path
):
"""An extension in neither the docling nor the text set is read as text."""
source = tmp_path / "notes.xyz"
source.write_text("Plain body for an unknown extension.")
doc = await converter.convert_file(source)
assert isinstance(doc, DoclingDocument)
assert "Plain body for an unknown extension." in doc.export_to_markdown()
@pytest.mark.asyncio
async def test_convert_file_raises_for_undecodable_file(self, converter, tmp_path):
source = tmp_path / "binary.xyz"
source.write_bytes(b"\xff\xfe\x00\x01 not utf-8")
with pytest.raises(ValueError, match="Failed to parse file"):
await converter.convert_file(source)
@pytest.mark.asyncio
async def test_convert_text_wraps_conversion_failure(self, converter, monkeypatch):
def boom(*_args, **_kwargs):
raise RuntimeError("docling exploded")
monkeypatch.setattr(converter, "_sync_convert_docling_text", boom)
with pytest.raises(ValueError, match="Failed to convert text"):
await converter.convert_text("# Test", name="test.md")
@pytest.mark.asyncio
async def test_convert_text_falls_back_when_format_not_inferable(self, converter):
"""docling raises ConversionError for an extension it has no backend
for; the simple-document fallback keeps the text."""
doc = await converter.convert_text("just some prose", name="mystery.zzz")
assert isinstance(doc, DoclingDocument)
assert "just some prose" in doc.export_to_markdown()
@pytest.mark.asyncio
async def test_convert_code_file(self, converter):
"""Test that code files are wrapped in code blocks."""
@ -1522,3 +1562,121 @@ class TestDoclingServeConverterIntegration:
assert str(sample.image.uri).startswith("data:image/"), (
"Rehydrated picture URI should be a data: URI, not a bare artifact filename"
)
class TestDoclingServeZipParsing:
"""_parse_zip_to_docling decodes the target_type=zip payload. These drive
its branches directly no docling-serve instance involved."""
@pytest.fixture
def converter(self):
config = AppConfig()
config.processing.converter = "docling-serve"
conv = get_converter(config)
assert isinstance(conv, DoclingServeConverter)
return conv
@staticmethod
def _zip(entries: dict[str, bytes]) -> bytes:
import io
import zipfile
buf = io.BytesIO()
with zipfile.ZipFile(buf, mode="w") as zf:
for name, blob in entries.items():
zf.writestr(name, blob)
return buf.getvalue()
@staticmethod
def _doc_json(**extra) -> dict:
base = {
"name": "document",
"texts": [],
"tables": [],
"pictures": [],
"groups": [],
"body": {"self_ref": "#/body", "children": []},
"furniture": {"self_ref": "#/furniture", "children": []},
}
base.update(extra)
return base
def test_raises_without_top_level_json(self, converter):
blob = self._zip({"artifacts/image.png": b"png"})
with pytest.raises(ValueError, match="no top-level JSON document"):
converter._parse_zip_to_docling(blob, "doc.pdf")
def test_picture_without_image_is_left_alone(self, converter):
import json as _json
doc_json = self._doc_json(
pictures=[
{
"self_ref": "#/pictures/0",
"label": "picture",
"image": None,
"prov": [],
}
]
)
blob = self._zip({"document.json": _json.dumps(doc_json).encode()})
doc = converter._parse_zip_to_docling(blob, "doc.pdf")
assert doc.pictures[0].image is None
def test_data_uri_image_is_passed_through(self, converter):
import json as _json
data_uri = "data:image/png;base64,aGVsbG8="
doc_json = self._doc_json(
pictures=[
{
"self_ref": "#/pictures/0",
"label": "picture",
"image": {
"mimetype": "image/png",
"dpi": 72,
"size": {"width": 1, "height": 1},
"uri": data_uri,
},
"prov": [],
}
]
)
blob = self._zip({"document.json": _json.dumps(doc_json).encode()})
doc = converter._parse_zip_to_docling(blob, "doc.pdf")
assert str(doc.pictures[0].image.uri) == data_uri
def test_non_dict_page_entry_is_skipped_while_inlining(self, converter):
"""A page entry that isn't an object must not blow up the image-inlining
loop with an AttributeError; it falls through to schema validation."""
import json as _json
from pydantic import ValidationError
doc_json = self._doc_json(pages={"1": "not-a-page-object"})
blob = self._zip({"document.json": _json.dumps(doc_json).encode()})
with pytest.raises(ValidationError):
converter._parse_zip_to_docling(blob, "doc.pdf")
@pytest.mark.asyncio
async def test_convert_text_rejects_unsupported_format(self, converter):
with pytest.raises(ValueError, match="Unsupported format"):
await converter.convert_text("body", format="pdf")
@pytest.mark.asyncio
async def test_convert_text_plain_builds_document_locally(self, converter):
"""format="plain" never reaches the network."""
converter.client.submit_and_poll_zip = AsyncMock( # ty: ignore[invalid-assignment]
side_effect=AssertionError("must not call docling-serve")
)
doc = await converter.convert_text("just text", format="plain")
assert isinstance(doc, DoclingDocument)
assert "just text" in doc.export_to_markdown()

View file

@ -481,3 +481,29 @@ def test_from_config_wires_retry_and_breaker():
assert client._max_attempts == 7
assert client._breaker_config.failure_threshold == 9
assert client._breaker_config.cooldown_s == 90.0
@pytest.mark.asyncio
async def test_submit_without_task_id_raises():
"""A 200 that carries no task_id is a protocol violation, not a silent pass."""
import httpx
from haiku.rag.providers.docling_serve import DoclingServeClient
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, json={})
client = DoclingServeClient(base_urls="http://docling:5001")
files = {"files": ("doc.pdf", b"pdf", "application/octet-stream")}
async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as http:
with pytest.raises(ValueError, match="did not return a task_id"):
await client._submit_and_wait(
http,
"http://docling:5001",
"/v1/convert/source/async",
files,
{},
{},
"doc.pdf",
)

View file

@ -148,3 +148,47 @@ def test_vllm_embedder_does_not_double_append_v1():
base_url = embedder._base_url.rstrip("/") # type: ignore[attr-defined] # ty: ignore[unresolved-attribute]
assert base_url.endswith("/v1")
assert not base_url.endswith("/v1/v1")
def test_vector_dim_property_reports_configured_dimension():
from haiku.rag.embeddings import EmbedderWrapper
assert EmbedderWrapper(embedder=None, vector_dim=512).vector_dim == 512
@pytest.mark.parametrize(
"provider,env_var",
[("voyageai", "VOYAGE_API_KEY"), ("cohere", "CO_API_KEY")],
)
def test_saas_providers_build_offline(monkeypatch, provider, env_var):
"""Construction only wires the SDK; no request is made."""
monkeypatch.setenv(env_var, "test-key")
config = AppConfig(
embeddings=EmbeddingsConfig(
model=EmbeddingModelConfig(
provider=provider, name="some-model", vector_dim=1024
),
),
)
embedder = get_embedder(config)
assert embedder.vector_dim == 1024
def test_cohere_floats_rejects_missing_embeddings():
from types import SimpleNamespace
from haiku.rag.embeddings.cohere import _floats
result = SimpleNamespace(embeddings=SimpleNamespace(float_=None))
with pytest.raises(ValueError, match="no float embeddings"):
_floats(result)
def test_voyageai_to_pil_rejects_unsupported_type():
from haiku.rag.embeddings.voyageai import _to_pil
with pytest.raises(TypeError, match="Unsupported image type"):
_to_pil("not an image") # ty: ignore[invalid-argument-type]

View file

@ -332,3 +332,152 @@ class TestMCPImageInput:
result = await ask(question="q")
assert result == "answer"
assert captured["images"] is None
class TestMCPFileAndUrlIngestion:
@pytest.mark.asyncio
async def test_add_document_from_file(self, temp_db_path, tmp_path):
async with HaikuRAG(temp_db_path, create=True):
pass
source = tmp_path / "note.txt"
source.write_text("Ingested from a file path.")
mcp = create_mcp_server(temp_db_path, read_only=False)
add_file = await _get_tool(mcp, "add_document_from_file")
doc_id = await add_file(file_path=str(source), title="File Doc")
assert doc_id is not None
get_doc = await _get_tool(mcp, "get_document")
doc = await get_doc(document_id=doc_id)
assert doc.title == "File Doc"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"tool_name,kwargs",
[
("add_document_from_file", {"file_path": "/tmp/x.txt"}),
("add_document_from_url", {"url": "https://example.com/x.txt"}),
],
)
@pytest.mark.parametrize(
"results,expected",
[
(
[Document(id="first", content="a"), Document(id="second", content="b")],
"first",
),
([], None),
],
ids=["directory_reports_first_id", "empty_directory_reports_none"],
)
async def test_add_tools_handle_multi_document_sources(
self, mcp_db, monkeypatch, tool_name, kwargs, results, expected
):
"""A source resolving to several documents reports the first id."""
async def fake_from_source(self, source, title=None, metadata=None, **kw):
return results
monkeypatch.setattr(HaikuRAG, "create_document_from_source", fake_from_source)
mcp = create_mcp_server(mcp_db, read_only=False)
add = await _get_tool(mcp, tool_name)
assert await add(**kwargs) == expected
@pytest.mark.asyncio
async def test_add_document_from_url(self, mcp_db, monkeypatch):
async def fake_from_source(self, source, title=None, metadata=None, **kwargs):
assert source == "https://example.com/doc.txt"
return Document(id="url-doc", content="fetched")
monkeypatch.setattr(HaikuRAG, "create_document_from_source", fake_from_source)
mcp = create_mcp_server(mcp_db, read_only=False)
add_url = await _get_tool(mcp, "add_document_from_url")
assert await add_url(url="https://example.com/doc.txt") == "url-doc"
class TestMCPToolsDegradeOnError:
"""Every tool swallows client failures and returns its empty value rather
than propagating an exception to the MCP transport."""
@pytest.mark.asyncio
@pytest.mark.parametrize(
"client_method,tool_name,kwargs,expected",
[
(
"create_document_from_source",
"add_document_from_file",
{"file_path": "/tmp/x.txt"},
None,
),
(
"create_document_from_source",
"add_document_from_url",
{"url": "https://example.com/x"},
None,
),
("create_document", "add_document_from_text", {"content": "x"}, None),
("delete_document", "delete_document", {"document_id": "x"}, False),
("search", "search_documents", {"query": "x"}, []),
("get_document_by_id", "get_document", {"document_id": "x"}, None),
("list_documents", "list_documents", {}, []),
],
)
async def test_tool_returns_empty_value_when_client_raises(
self, mcp_db, monkeypatch, client_method, tool_name, kwargs, expected
):
async def boom(self, *args, **kw):
raise RuntimeError("client exploded")
monkeypatch.setattr(HaikuRAG, client_method, boom)
mcp = create_mcp_server(mcp_db, read_only=False)
tool = await _get_tool(mcp, tool_name)
assert await tool(**kwargs) == expected
@pytest.mark.asyncio
async def test_list_documents_returns_empty_for_invalid_filter(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=True)
list_docs = await _get_tool(mcp, "list_documents")
assert await list_docs(filter="no_such_column = 1") == []
@pytest.mark.asyncio
async def test_analyze_reports_the_error(self, mcp_db, monkeypatch):
async def boom(self, question, filter=None, images=None):
raise RuntimeError("sandbox exploded")
monkeypatch.setattr(HaikuRAG, "analyze", boom)
mcp = create_mcp_server(mcp_db, read_only=True)
analyze = await _get_tool(mcp, "analyze")
assert "sandbox exploded" in await analyze(question="q")
@pytest.mark.asyncio
async def test_ask_question_appends_citations_when_requested(
self, mcp_db, monkeypatch
):
from haiku.rag.store.models.citation import Citation
citation = Citation(
chunk_id="c1",
document_id="d1",
content="cited text",
document_uri="test://ai-overview",
document_title="AI Overview",
)
async def fake_ask(self, question, filter=None, images=None):
return ("the answer", [citation])
monkeypatch.setattr(HaikuRAG, "ask", fake_ask)
mcp = create_mcp_server(mcp_db, read_only=True)
ask = await _get_tool(mcp, "ask_question")
with_cite = await ask(question="q", cite=True)
assert with_cite.startswith("the answer")
assert "AI Overview" in with_cite
assert await ask(question="q", cite=False) == "the answer"

View file

@ -243,3 +243,14 @@ def test_concatenate_shifts_page_nos_and_unique_self_refs():
assert page_nos == [1, 2], f"expected b's page 1 to shift to page 2, got {page_nos}"
assert sorted(merged.pages.keys()) == [1, 2]
def test_iter_pdf_slices_rejects_unopenable_pdf(tmp_path):
"""pdfium refuses non-PDF bytes; the caller sees UnsupportedSourceError."""
from haiku.rag.client.exceptions import UnsupportedSourceError
junk = tmp_path / "not-really.pdf"
junk.write_bytes(b"this is not a pdf at all")
with pytest.raises(UnsupportedSourceError, match="cannot open PDF"):
list(iter_pdf_slices(junk, slice_size=1))

View file

@ -479,3 +479,30 @@ async def test_cross_encoder_reranker():
assert "0" in top_ids or "2" in top_ids
except ImportError:
pytest.skip("sentence-transformers not installed")
@pytest.mark.asyncio
async def test_cross_encoder_reranks_via_model_ranking(monkeypatch):
"""The rank() results map back onto the input chunks by corpus_id."""
from haiku.rag.reranking import cross_encoder as ce_module
class _StubCrossEncoder:
def __init__(self, model):
self.model = model
def rank(self, query, documents, top_k=10):
# Reverse order so the mapping back to chunks is observable.
return [
{"corpus_id": i, "score": 1.0 - (i / 10)}
for i in reversed(range(len(documents)))
][:top_k]
monkeypatch.setattr(ce_module, "CrossEncoder", _StubCrossEncoder)
reranker = ce_module.CrossEncoderReranker("stub/model")
reranked = await reranker.rerank("query", chunks, top_n=2)
assert len(reranked) == 2
last_index = len(chunks) - 1
assert reranked[0][0] is chunks[last_index]
assert reranked[0][1] == pytest.approx(1.0 - last_index / 10)