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
This commit is contained in:
parent
88cd2b1d6c
commit
2a6d72171d
4 changed files with 271 additions and 154 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Reference in a new issue