From 2a6d72171d1d8fae156e42695f8ce383e78e5bd5 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Fri, 4 Sep 2026 12:03:06 +0300 Subject: [PATCH] Make MCP failures errors on the wire FastMCP masks unexpected exceptions and logs the traceback server-side, so paths and provider URLs never cross the transport. Expected failures raise ToolError with a message: unknown document, unknown collection, invalid filter, invalid base64, and agent failures naming only the exception type. No tool returns an empty value or an error string on failure any more. A filter is validated on its own before the read that would use it, with a filtered count on one of the selected databases: that is the query engine rejecting the filter and nothing else, so its message (columns and the statement) can be forwarded, while a ValueError raised later in the read stays masked and no database outside the selection is opened. Refs #599 --- CHANGELOG.md | 5 + docs/mcp.md | 9 + haiku_rag_slim/haiku/rag/mcp.py | 117 ++++++++----- tests/test_mcp.py | 294 ++++++++++++++++++++------------ 4 files changed, 271 insertions(+), 154 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index bc98984f..7940d86c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -22,6 +22,11 @@ - `processing.conversion_options.picture_description.model` defaults to `enable_thinking: false`, and the field now reaches the VLM: docling's picture-description request carries `reasoning_effort` in `params`. +- MCP tools raise on failure; an empty result no longer doubles as an error. + Unknown document, unknown collection, invalid filter and invalid base64 + carry a message; `ask_question` and `analyze` failures name the exception + type. Anything else is masked (`mask_error_details=True`) and logged + server-side. - `haiku-rag mcp` covers the configured `lancedb.databases` set. `sources` on `search_documents`, `search_documents_by_image`, `ask_question` and `analyze`; `source` on `get_document`; an unknown name is a tool error. diff --git a/docs/mcp.md b/docs/mcp.md index fd500d07..d63f5880 100644 --- a/docs/mcp.md +++ b/docs/mcp.md @@ -109,6 +109,15 @@ uri LIKE '%.pdf' title = 'Q3 report' ``` +### Errors + +A failure is an MCP error, never an empty result. Expected failures carry a +message: a document id that matches nothing, a collection the server does not +cover, a filter the query engine rejects (with its message), invalid base64, +and an `ask_question` or `analyze` failure naming only the exception type. +Anything else reaches the client as `Error calling tool 'name'` and its +traceback goes to the server log. + ### Instructions The server publishes `instructions` describing the knowledge base: what it diff --git a/haiku_rag_slim/haiku/rag/mcp.py b/haiku_rag_slim/haiku/rag/mcp.py index e94b7624..5c285952 100644 --- a/haiku_rag_slim/haiku/rag/mcp.py +++ b/haiku_rag_slim/haiku/rag/mcp.py @@ -1,4 +1,5 @@ import asyncio +import logging from collections.abc import AsyncIterator from contextlib import AsyncExitStack, asynccontextmanager from importlib import metadata @@ -21,6 +22,8 @@ from haiku.rag.utils import format_citations if TYPE_CHECKING: from haiku.rag.client.scope import DatabaseScope +logger = logging.getLogger(__name__) + _FILTER_COLUMNS = ", ".join(DocumentMetaRecord.model_fields) Filter = Annotated[ @@ -47,7 +50,12 @@ def _read_only(title: str) -> ToolAnnotations: def _decode_image(image_base64: str) -> bytes: import base64 - return base64.b64decode(image_base64, validate=True) + try: + return base64.b64decode(image_base64, validate=True) + except ValueError as e: + # binascii.Error for characters outside the alphabet or bad padding, + # ValueError itself for non-ASCII input. + raise ToolError("Invalid base64 image") from e def _decode_images(images_base64: list[str] | None) -> list[bytes] | None: @@ -56,6 +64,28 @@ def _decode_images(images_base64: list[str] | None) -> list[bytes] | None: return [_decode_image(b64) for b64 in images_base64] +async def _check_filter( + rag: HaikuRAG, filter: str | None, sources: list[str] | None = None +) -> None: + """Evaluate a filter on its own before the read that would use it. + + A filtered count on one selected database runs the same predicate on the + same table and nothing else, so a ValueError here is the query engine + rejecting the filter; its message names columns and the statement, never + a location. A ValueError raised later in the read stays masked. Only the + selection is touched: every database shares the schema, so one suffices. + """ + if filter is None: + return + selected = await rag.clients_covering(sources) + if not selected: + return + try: + await selected[0].count_documents(filter=filter) + except ValueError as e: + raise ToolError(f"Invalid filter {filter!r}: {e}") from e + + def _instructions(scope: "DatabaseScope", config: AppConfig) -> str: """What the server is for, naming no tools: the client has every tool's description from the listing.""" @@ -136,11 +166,14 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: finally: client = None + # Masking keeps paths and provider URLs out of an unexpected error's text; + # the traceback goes to the server log. A ToolError reaches the client as is. mcp = FastMCP( "haiku-rag", instructions=_instructions(scope, config), version=metadata.version("haiku.rag-slim"), lifespan=lifespan, + mask_error_details=True, ) @mcp.tool(annotations=_read_only("Search documents")) @@ -167,8 +200,9 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: include_images: Attach the bytes of pictures in the results as base64 PNG under `image_data`. False for a smaller response. """ + rag = await _client() try: - rag = await _client() + await _check_filter(rag, filter, sources) return await rag.search( query, limit=limit, @@ -178,8 +212,6 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: ) except UnknownDatabaseError as e: raise ToolError(str(e)) from e - except Exception: - return [] # Image-as-query tool, only registered when the configured embedder # supports image embeddings. Probed at server-build time when no Store is @@ -211,9 +243,10 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: include_images: Attach the bytes of pictures in the results as base64 PNG under `image_data`. False for a smaller response. """ + raw = _decode_image(image_base64) + rag = await _client() try: - raw = _decode_image(image_base64) - rag = await _client() + await _check_filter(rag, filter, sources) return await rag.search( raw, limit=limit, @@ -223,13 +256,9 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: ) except UnknownDatabaseError as e: raise ToolError(str(e)) from e - except Exception: - return [] @mcp.tool(annotations=_read_only("Get document")) - async def get_document( - document_id: str, source: str | None = None - ) -> Document | None: + async def get_document(document_id: str, source: str | None = None) -> Document: """Read one document whole, in reading order. Use this after a search when a passage is not enough. Returns the @@ -241,13 +270,14 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: source: The collection holding it. Without one every collection is asked. """ + rag = await _client() try: - rag = await _client() - return await rag.get_document_by_id(document_id, source) + document = await rag.get_document_by_id(document_id, source) except UnknownDatabaseError as e: raise ToolError(str(e)) from e - except Exception: - return None + if document is None: + raise ToolError(f"No document with id {document_id!r}") + return document @mcp.tool(annotations=_read_only("List documents")) async def list_documents( @@ -265,23 +295,20 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: limit: How many documents to return. offset: How many documents to skip, for paging. """ - try: - rag = await _client() - documents = await rag.list_documents(limit, offset, filter) - - return [ - DocumentInfo( - id=doc.id, - title=doc.title or "Untitled", - uri=doc.uri or "", - created=doc.created_at.strftime("%Y-%m-%d"), - source=doc.source, - metadata=doc.metadata, - ) - for doc in documents - ] - except Exception: - return [] + rag = await _client() + await _check_filter(rag, filter) + documents = await rag.list_documents(limit, offset, filter) + return [ + DocumentInfo( + id=doc.id, + title=doc.title or "Untitled", + uri=doc.uri or "", + created=doc.created_at.strftime("%Y-%m-%d"), + source=doc.source, + metadata=doc.metadata, + ) + for doc in documents + ] @mcp.tool(annotations=_read_only("Ask a question")) async def ask_question( @@ -303,19 +330,20 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: images_base64: Images to attach to the question, PNG or JPEG bytes as base64. Needs a vision-capable model on the server. """ + images = _decode_images(images_base64) + rag = await _client() try: - images = _decode_images(images_base64) - rag = await _client() answer, citations = await rag.ask(question, images=images, sources=sources) - if cite and citations: - answer += "\n\n" + format_citations( - citations, include_source=rag.covers_multiple - ) - return answer except UnknownDatabaseError as e: raise ToolError(str(e)) from e except Exception as e: - return f"Error answering question: {e!s}" + logger.exception("ask_question failed") + raise ToolError(f"ask_question failed: {type(e).__name__}") from e + if cite and citations: + answer += "\n\n" + format_citations( + citations, include_source=rag.covers_multiple + ) + return answer @mcp.tool(annotations=_read_only("Analyze documents")) async def analyze( @@ -336,16 +364,17 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP: images_base64: Images to attach to the question, PNG or JPEG bytes as base64. Needs a vision-capable model on the server. """ + images = _decode_images(images_base64) + rag = await _client() try: - images = _decode_images(images_base64) - rag = await _client() result = await rag.analyze( question, filter=filter, images=images, sources=sources ) - return result.answer except UnknownDatabaseError as e: raise ToolError(str(e)) from e except Exception as e: - return f"Error running analysis capability: {e!s}" + logger.exception("analyze failed") + raise ToolError(f"analyze failed: {type(e).__name__}") from e + return result.answer return mcp diff --git a/tests/test_mcp.py b/tests/test_mcp.py index 82799789..b21b24cd 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -1,6 +1,8 @@ +import logging from types import SimpleNamespace import pytest +from fastmcp.exceptions import ToolError from haiku.rag.client import HaikuRAG from haiku.rag.mcp import _covering as _mcp_covering @@ -87,6 +89,14 @@ async def _get_tool(mcp, name): return tool.fn +async def _call(mcp, name, **kwargs): + """Call a tool over the wire, returning the result whether or not it errored.""" + from fastmcp import Client + + async with Client(mcp) as client: + return await client.call_tool(name, kwargs, raise_on_error=False) + + class TestMCPReadTools: @pytest.mark.asyncio async def test_search_documents(self, mcp_db): @@ -184,14 +194,6 @@ class TestMCPReadTools: assert "docling_document" not in serialized assert "docling_version" not in serialized - @pytest.mark.asyncio - async def test_get_document_not_found(self, mcp_db): - mcp = create_mcp_server(mcp_db) - get_doc = await _get_tool(mcp, "get_document") - - result = await get_doc(document_id="nonexistent-id") - assert result is None - @pytest.mark.asyncio async def test_list_documents(self, mcp_db): mcp = create_mcp_server(mcp_db) @@ -233,6 +235,36 @@ class TestMCPReadTools: ] assert overview["metadata"] == {"author": "Ada"} + @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", + source="alpha", + ) + + async def fake_ask(self, question, filter=None, images=None, sources=None): + return ("the answer", [citation]) + + monkeypatch.setattr(HaikuRAG, "ask", fake_ask) + mcp = create_mcp_server(mcp_db) + 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 + # One database: its name adds nothing. + assert "alpha" not in with_cite + + assert await ask(question="q", cite=False) == "the answer" + @pytest.mark.filterwarnings("ignore:Found propagated trace context:RuntimeWarning") class TestMCPDescribesItself: @@ -367,14 +399,28 @@ class TestMCPCoversTheConfiguredSet: async def test_an_unknown_database_is_an_error_not_an_empty_result( self, two_dbs, multimodal_embedder, tool_name, kwargs ): - from fastmcp.exceptions import ToolError - mcp = _covering_all(two_dbs) tool = await _get_tool(mcp, tool_name) with pytest.raises(ToolError, match="nope"): await tool(**kwargs) + @pytest.mark.asyncio + async def test_a_filtered_search_touches_only_the_selected_databases(self, two_dbs): + """alpha is gone; a filtered search selecting beta must not notice.""" + import shutil + + shutil.rmtree(two_dbs.lancedb.databases["alpha"]) + mcp = _covering_all(two_dbs) + search = await _get_tool(mcp, "search_documents") + + results = await search( + query="cats", filter="uri LIKE '%beta%'", sources=["beta"] + ) + assert results + assert {r.source for r in results} == {"beta"} + assert await search(query="cats", filter="uri LIKE '%beta%'", sources=[]) == [] + @pytest.mark.asyncio async def test_the_listing_covers_every_database(self, two_dbs): mcp = _covering_all(two_dbs) @@ -495,13 +541,13 @@ class TestMCPImageQuery: results = await search_by_image( image_base64=base64.b64encode(png).decode("ascii"), filter="uri LIKE 'x%'", - sources=["alpha"], + sources=[], ) assert results == [] assert seen["query"] == png assert seen["filter"] == "uri LIKE 'x%'" - assert seen["sources"] == ["alpha"] + assert seen["sources"] == [] @pytest.mark.asyncio async def test_image_query_rejects_characters_outside_the_alphabet( @@ -519,36 +565,10 @@ class TestMCPImageQuery: mcp = create_mcp_server(mcp_db) search_by_image = await _get_tool(mcp, "search_documents_by_image") - assert await search_by_image(image_base64="AAAA!!!!") == [] + with pytest.raises(ToolError): + await search_by_image(image_base64="AAAA!!!!") assert not searched - @pytest.mark.asyncio - async def test_image_query_returns_empty_on_invalid_base64( - self, mcp_db, multimodal_embedder - ): - """Garbage base64 from the caller is swallowed, returning an empty - list rather than crashing the MCP server.""" - mcp = create_mcp_server(mcp_db) - search_by_image = await _get_tool(mcp, "search_documents_by_image") - - # Not valid base64 (contains non-base64 chars) — the strict decoder - # in search_documents_by_image rejects it. - results = await search_by_image(image_base64="!!! not base64 !!!") - assert results == [] - - @pytest.mark.asyncio - async def test_image_query_returns_empty_when_the_search_raises( - self, mcp_db, multimodal_embedder, monkeypatch - ): - async def boom(self, *args, **kw): - raise RuntimeError("client exploded") - - monkeypatch.setattr(HaikuRAG, "search", boom) - mcp = create_mcp_server(mcp_db) - search_by_image = await _get_tool(mcp, "search_documents_by_image") - - assert await search_by_image(image_base64="AAAA") == [] - class TestMCPImageInput: @pytest.mark.asyncio @@ -590,14 +610,6 @@ class TestMCPImageInput: assert result == "answer" assert captured["images"] == [jpeg] - @pytest.mark.asyncio - async def test_ask_question_rejects_invalid_base64(self, mcp_db): - mcp = create_mcp_server(mcp_db) - ask = await _get_tool(mcp, "ask_question") - - result = await ask(question="q", images_base64=["!!! not base64 !!!"]) - assert "Error" in result - @pytest.mark.asyncio async def test_ask_question_without_images_passes_none(self, mcp_db, monkeypatch): captured = {} @@ -615,78 +627,140 @@ class TestMCPImageInput: assert captured["images"] is None -class TestMCPToolsDegradeOnError: - """Every tool swallows client failures and returns its empty value rather - than propagating an exception to the MCP transport.""" +@pytest.mark.filterwarnings("ignore:Found propagated trace context:RuntimeWarning") +class TestMCPErrorContract: + """A failure is an error on the wire, never an empty result. Expected + failures say what went wrong; anything else is masked and logged on the + server.""" + + @pytest.mark.asyncio + async def test_an_unknown_document_is_an_error(self, mcp_db): + result = await _call( + create_mcp_server(mcp_db), "get_document", document_id="nonexistent-id" + ) + + assert result.is_error + assert "nonexistent-id" in result.content[0].text @pytest.mark.asyncio @pytest.mark.parametrize( - "client_method,tool_name,kwargs,expected", - [ - ("search", "search_documents", {"query": "x"}, []), - ("get_document_by_id", "get_document", {"document_id": "x"}, None), - ("list_documents", "list_documents", {}, []), - ], + "tool_name,kwargs", + [("search_documents", {"query": "x"}), ("list_documents", {})], ) - async def test_tool_returns_empty_value_when_client_raises( - self, mcp_db, monkeypatch, client_method, tool_name, kwargs, expected + async def test_an_invalid_filter_is_an_error_naming_the_filter( + self, mcp_db, tool_name, kwargs ): - async def boom(self, *args, **kw): - raise RuntimeError("client exploded") - - monkeypatch.setattr(HaikuRAG, client_method, boom) - mcp = create_mcp_server(mcp_db) - 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) - 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, sources=None): - raise RuntimeError("sandbox exploded") - - monkeypatch.setattr(HaikuRAG, "analyze", boom) - mcp = create_mcp_server(mcp_db) - 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", - source="alpha", + result = await _call( + create_mcp_server(mcp_db), tool_name, filter="no_such_column = 1", **kwargs ) - async def fake_ask(self, question, filter=None, images=None, sources=None): - return ("the answer", [citation]) + assert result.is_error + assert "no_such_column = 1" in result.content[0].text - monkeypatch.setattr(HaikuRAG, "ask", fake_ask) - mcp = create_mcp_server(mcp_db) - ask = await _get_tool(mcp, "ask_question") + @pytest.mark.asyncio + @pytest.mark.parametrize("filter", [None, "title = 'AI Overview'"]) + async def test_a_value_error_from_the_read_is_not_an_invalid_filter( + self, mcp_db, monkeypatch, filter + ): + """Only the filter check translates ValueError; one raised by the read + itself, with or without a valid filter, stays masked.""" - with_cite = await ask(question="q", cite=True) - assert with_cite.startswith("the answer") - assert "AI Overview" in with_cite - # One database: its name adds nothing. - assert "alpha" not in with_cite + async def boom(self, *args, **kw): + raise ValueError("boom at /secret/path") - assert await ask(question="q", cite=False) == "the answer" + monkeypatch.setattr(HaikuRAG, "search", boom) + result = await _call( + create_mcp_server(mcp_db), "search_documents", query="x", filter=filter + ) + + assert result.is_error + assert "filter" not in result.content[0].text + assert "/secret/path" not in result.content[0].text + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "payload", ["!!! not base64 !!!", "é"], ids=["outside_alphabet", "non_ascii"] + ) + @pytest.mark.parametrize( + "tool_name,image_param,many", + [ + ("search_documents_by_image", "image_base64", False), + ("ask_question", "images_base64", True), + ("analyze", "images_base64", True), + ], + ) + async def test_invalid_base64_is_an_error( + self, mcp_db, multimodal_embedder, tool_name, image_param, many, payload + ): + kwargs: dict[str, object] = {"question": "q"} if many else {} + kwargs[image_param] = [payload] if many else payload + + result = await _call(create_mcp_server(mcp_db), tool_name, **kwargs) + + assert result.is_error + assert "base64" in result.content[0].text + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "client_method,tool_name", + [("ask", "ask_question"), ("analyze", "analyze")], + ) + async def test_an_agent_failure_names_only_its_type( + self, mcp_db, monkeypatch, caplog, client_method, tool_name + ): + async def boom(self, question, filter=None, images=None, sources=None): + raise RuntimeError("boom at /secret/path") + + monkeypatch.setattr(HaikuRAG, client_method, boom) + with caplog.at_level(logging.ERROR, logger="haiku.rag.mcp"): + result = await _call(create_mcp_server(mcp_db), tool_name, question="q") + + assert result.is_error + assert "RuntimeError" in result.content[0].text + assert "/secret/path" not in result.content[0].text + assert any( + r.exc_info and "boom at /secret/path" in str(r.exc_info[1]) + for r in caplog.records + ) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "client_method,tool_name,kwargs", + [ + ("search", "search_documents", {"query": "x"}), + ("search", "search_documents_by_image", {"image_base64": "AAAA"}), + ("get_document_by_id", "get_document", {"document_id": "x"}), + ("list_documents", "list_documents", {}), + ], + ) + async def test_an_unexpected_failure_is_masked_and_logged( + self, + mcp_db, + multimodal_embedder, + monkeypatch, + caplog, + client_method, + tool_name, + kwargs, + ): + async def boom(self, *args, **kw): + raise RuntimeError("boom at /secret/path") + + monkeypatch.setattr(HaikuRAG, client_method, boom) + # fastmcp's logger does not propagate, so listen to it directly. + fastmcp_logger = logging.getLogger("fastmcp") + fastmcp_logger.addHandler(caplog.handler) + try: + result = await _call(create_mcp_server(mcp_db), tool_name, **kwargs) + finally: + fastmcp_logger.removeHandler(caplog.handler) + + assert result.is_error + assert "/secret/path" not in result.content[0].text + assert any( + r.exc_info and "boom at /secret/path" in str(r.exc_info[1]) + for r in caplog.records + ) class TestMCPClientLifetime: