Return search results as agent text, images and structured content
search_documents and search_documents_by_image return a ToolResult: the format_for_agent rendering with rank, Document ID and Collection so the text alone drives the document tools; one ImageContent per distinct picture, labelled with its result; and the SearchResult list without image_data as structured content. format_for_agent gains an opt-in include_document_id, so the capabilities' rendering is unchanged. collect_pictures is the one place pictures are deduplicated and validated for both wire formats. Refs #599
This commit is contained in:
parent
15afb97a6e
commit
5653e876a8
7 changed files with 316 additions and 60 deletions
|
|
@ -25,6 +25,12 @@
|
||||||
- `processing.conversion_options.picture_description.model` defaults to
|
- `processing.conversion_options.picture_description.model` defaults to
|
||||||
`enable_thinking: false`, and the field now reaches the VLM: docling's
|
`enable_thinking: false`, and the field now reaches the VLM: docling's
|
||||||
picture-description request carries `reasoning_effort` in `params`.
|
picture-description request carries `reasoning_effort` in `params`.
|
||||||
|
- MCP `search_documents` and `search_documents_by_image` return the agent
|
||||||
|
rendering as text (rank, `Document ID`, `Collection` over several
|
||||||
|
databases, title, headings, passage), pictures as `ImageContent` blocks,
|
||||||
|
and the `SearchResult` list without `image_data` as structured content.
|
||||||
|
`SearchResult.format_for_agent(include_document_id=)`;
|
||||||
|
`collect_pictures` in `haiku.rag.tools.search`.
|
||||||
- MCP tools raise on failure; an empty result no longer doubles as an error.
|
- MCP tools raise on failure; an empty result no longer doubles as an error.
|
||||||
Unknown document, unknown collection, invalid filter and invalid base64
|
Unknown document, unknown collection, invalid filter and invalid base64
|
||||||
carry a message; `ask_question` and `analyze` failures name the exception
|
carry a message; `ask_question` and `analyze` failures name the exception
|
||||||
|
|
|
||||||
14
docs/mcp.md
14
docs/mcp.md
|
|
@ -84,11 +84,15 @@ repeating it.
|
||||||
| `ask_question` | always | `question`, `images_base64`, `sources` |
|
| `ask_question` | always | `question`, `images_base64`, `sources` |
|
||||||
| `analyze` | always | `question`, `filter`, `images_base64`, `sources` |
|
| `analyze` | always | `question`, `filter`, `images_base64`, `sources` |
|
||||||
|
|
||||||
`search_documents` runs hybrid search, vector and full-text, and returns
|
`search_documents` runs hybrid search, vector and full-text. Its text content
|
||||||
results best first. Scores are not comparable across queries or search types.
|
is the rendering the in-process agents read: results best first, each with its
|
||||||
Rank is the signal. `include_images` attaches picture bytes as base64 PNG under
|
rank, `Document ID`, `Collection` when the server covers several, the document
|
||||||
`image_data`. `search_documents_by_image` embeds the query image and searches
|
title, section headings and the passage. Pictures in the results follow as
|
||||||
by vector similarity alone.
|
image blocks, one per distinct picture, each preceded by a line naming its
|
||||||
|
result; `include_images: false` leaves them out. The structured content is the
|
||||||
|
`SearchResult` list without picture bytes. Scores are not comparable across
|
||||||
|
queries or search types, so rank is the signal. `search_documents_by_image`
|
||||||
|
embeds the query image and searches by vector similarity alone.
|
||||||
|
|
||||||
`get_document` returns a document whole, in reading order. For a long one,
|
`get_document` returns a document whole, in reading order. For a long one,
|
||||||
`get_document_outline` returns the heading tree with page numbers and
|
`get_document_outline` returns the heading tree with page numbers and
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,8 @@ from typing import TYPE_CHECKING, Annotated
|
||||||
|
|
||||||
from fastmcp import FastMCP
|
from fastmcp import FastMCP
|
||||||
from fastmcp.exceptions import ToolError
|
from fastmcp.exceptions import ToolError
|
||||||
from mcp.types import ToolAnnotations
|
from fastmcp.tools import ToolResult
|
||||||
|
from mcp.types import ContentBlock, ImageContent, TextContent, ToolAnnotations
|
||||||
from pydantic import Field
|
from pydantic import Field
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
@ -19,6 +20,7 @@ from haiku.rag.store.models import Document, SearchResult
|
||||||
from haiku.rag.store.models.document_item import DocumentItem
|
from haiku.rag.store.models.document_item import DocumentItem
|
||||||
from haiku.rag.store.schema import DocumentMetaRecord
|
from haiku.rag.store.schema import DocumentMetaRecord
|
||||||
from haiku.rag.tools.document import DocumentInfo, DocumentSection, OutlineNode
|
from haiku.rag.tools.document import DocumentInfo, DocumentSection, OutlineNode
|
||||||
|
from haiku.rag.tools.search import collect_pictures
|
||||||
from haiku.rag.utils import format_citations
|
from haiku.rag.utils import format_citations
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
|
@ -110,6 +112,52 @@ def _instructions(scope: "DatabaseScope", config: AppConfig) -> str:
|
||||||
return "\n".join(lines)
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def _search_result(results: list[SearchResult], covers_multiple: bool) -> ToolResult:
|
||||||
|
"""Results as the in-process agents read them, then each distinct picture
|
||||||
|
as an image block labelled with its result, and the results as structured
|
||||||
|
content without the picture bytes."""
|
||||||
|
import base64
|
||||||
|
|
||||||
|
total = len(results)
|
||||||
|
text = "\n\n".join(
|
||||||
|
result.format_for_agent(
|
||||||
|
rank=rank,
|
||||||
|
total=total,
|
||||||
|
include_collection=covers_multiple,
|
||||||
|
include_document_id=True,
|
||||||
|
)
|
||||||
|
for rank, result in enumerate(results, 1)
|
||||||
|
)
|
||||||
|
content: list[ContentBlock] = [
|
||||||
|
TextContent(type="text", text=text or "No results found.")
|
||||||
|
]
|
||||||
|
pictures, _ = collect_pictures(results)
|
||||||
|
for source, chunk_id, self_ref, picture in pictures:
|
||||||
|
collection = f" in {source}" if covers_multiple and source else ""
|
||||||
|
content.append(
|
||||||
|
TextContent(
|
||||||
|
type="text",
|
||||||
|
text=f"Picture {self_ref} of search result [{chunk_id}]{collection}",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
content.append(
|
||||||
|
ImageContent(
|
||||||
|
type="image",
|
||||||
|
data=base64.b64encode(picture.data).decode("ascii"),
|
||||||
|
mimeType="image/png",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return ToolResult(
|
||||||
|
content=content,
|
||||||
|
structured_content={
|
||||||
|
"result": [
|
||||||
|
result.model_dump(mode="json", exclude={"image_data"})
|
||||||
|
for result in results
|
||||||
|
]
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _node(toc: "dict[str, Any]") -> OutlineNode:
|
def _node(toc: "dict[str, Any]") -> OutlineNode:
|
||||||
return OutlineNode(
|
return OutlineNode(
|
||||||
id=toc["self_ref"],
|
id=toc["self_ref"],
|
||||||
|
|
@ -207,27 +255,30 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP:
|
||||||
include_images: bool = True,
|
include_images: bool = True,
|
||||||
filter: Filter = None,
|
filter: Filter = None,
|
||||||
sources: Sources = None,
|
sources: Sources = None,
|
||||||
) -> list[SearchResult]:
|
) -> ToolResult:
|
||||||
"""Search the knowledge base by meaning and keyword.
|
"""Search the knowledge base by meaning and keyword.
|
||||||
|
|
||||||
Use this first for any question the documents might answer; it needs
|
Use this first for any question the documents might answer; it needs
|
||||||
no model and is the cheapest call. Results come best first, each with
|
no model and is the cheapest call. Results come best first, each with
|
||||||
the document's id, title and collection, the section headings and the
|
its rank, `Document ID`, `Collection` when the server covers several,
|
||||||
matching passage. Scores are not comparable across queries, so read
|
the document title, section headings and the matching passage; pass
|
||||||
the order, not the numbers. If nothing relevant comes back, rephrase
|
the id and collection to the document tools. Pictures in the results
|
||||||
once or narrow with `filter` before concluding the material is absent.
|
follow as images, each labelled with its result. Ranks, not scores,
|
||||||
|
are the signal: scores are not comparable across queries. If nothing
|
||||||
|
relevant comes back, rephrase once or narrow with `filter` before
|
||||||
|
concluding the material is absent.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
query: What to look for, in natural language or keywords.
|
query: What to look for, in natural language or keywords.
|
||||||
limit: How many results to return; the server's configured default
|
limit: How many results to return; the server's configured default
|
||||||
when omitted.
|
when omitted.
|
||||||
include_images: Attach the bytes of pictures in the results as
|
include_images: Return the pictures in the results as images.
|
||||||
base64 PNG under `image_data`. False for a smaller response.
|
False for a smaller response.
|
||||||
"""
|
"""
|
||||||
rag = await _client()
|
rag = await _client()
|
||||||
try:
|
try:
|
||||||
await _check_filter(rag, filter, sources)
|
await _check_filter(rag, filter, sources)
|
||||||
return await rag.search(
|
results = await rag.search(
|
||||||
query,
|
query,
|
||||||
limit=limit,
|
limit=limit,
|
||||||
filter=filter,
|
filter=filter,
|
||||||
|
|
@ -236,6 +287,7 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP:
|
||||||
)
|
)
|
||||||
except UnknownDatabaseError as e:
|
except UnknownDatabaseError as e:
|
||||||
raise ToolError(str(e)) from e
|
raise ToolError(str(e)) from e
|
||||||
|
return _search_result(results, rag.covers_multiple)
|
||||||
|
|
||||||
# Image-as-query tool, only registered when the configured embedder
|
# Image-as-query tool, only registered when the configured embedder
|
||||||
# supports image embeddings. Probed at server-build time when no Store is
|
# supports image embeddings. Probed at server-build time when no Store is
|
||||||
|
|
@ -252,7 +304,7 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP:
|
||||||
include_images: bool = True,
|
include_images: bool = True,
|
||||||
filter: Filter = None,
|
filter: Filter = None,
|
||||||
sources: Sources = None,
|
sources: Sources = None,
|
||||||
) -> list[SearchResult]:
|
) -> ToolResult:
|
||||||
"""Search the knowledge base with an image as the query.
|
"""Search the knowledge base with an image as the query.
|
||||||
|
|
||||||
Use this when the question is about a picture rather than words.
|
Use this when the question is about a picture rather than words.
|
||||||
|
|
@ -264,14 +316,14 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP:
|
||||||
image_base64: The query image, PNG or JPEG bytes as base64.
|
image_base64: The query image, PNG or JPEG bytes as base64.
|
||||||
limit: How many results to return; the server's configured
|
limit: How many results to return; the server's configured
|
||||||
default when omitted.
|
default when omitted.
|
||||||
include_images: Attach the bytes of pictures in the results as
|
include_images: Return the pictures in the results as images.
|
||||||
base64 PNG under `image_data`. False for a smaller response.
|
False for a smaller response.
|
||||||
"""
|
"""
|
||||||
raw = _decode_image(image_base64)
|
raw = _decode_image(image_base64)
|
||||||
rag = await _client()
|
rag = await _client()
|
||||||
try:
|
try:
|
||||||
await _check_filter(rag, filter, sources)
|
await _check_filter(rag, filter, sources)
|
||||||
return await rag.search(
|
results = await rag.search(
|
||||||
raw,
|
raw,
|
||||||
limit=limit,
|
limit=limit,
|
||||||
filter=filter,
|
filter=filter,
|
||||||
|
|
@ -280,6 +332,7 @@ def _covering(scope: "DatabaseScope", config: AppConfig) -> FastMCP:
|
||||||
)
|
)
|
||||||
except UnknownDatabaseError as e:
|
except UnknownDatabaseError as e:
|
||||||
raise ToolError(str(e)) from e
|
raise ToolError(str(e)) from e
|
||||||
|
return _search_result(results, rag.covers_multiple)
|
||||||
|
|
||||||
@mcp.tool(annotations=_read_only("Get document"))
|
@mcp.tool(annotations=_read_only("Get document"))
|
||||||
async def get_document(document_id: str, source: str | None = None) -> Document:
|
async def get_document(document_id: str, source: str | None = None) -> Document:
|
||||||
|
|
|
||||||
|
|
@ -202,6 +202,7 @@ class SearchResult(BaseModel):
|
||||||
total: int | None = None,
|
total: int | None = None,
|
||||||
*,
|
*,
|
||||||
include_collection: bool = False,
|
include_collection: bool = False,
|
||||||
|
include_document_id: bool = False,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Format this search result for inclusion in agent context.
|
"""Format this search result for inclusion in agent context.
|
||||||
|
|
||||||
|
|
@ -215,7 +216,8 @@ class SearchResult(BaseModel):
|
||||||
|
|
||||||
`include_collection` is the caller's decision, not this result's: a
|
`include_collection` is the caller's decision, not this result's: a
|
||||||
search spanning one collection has nothing to distinguish, whether or
|
search spanning one collection has nothing to distinguish, whether or
|
||||||
not that collection is named.
|
not that collection is named. `include_document_id` is for a reader
|
||||||
|
that will fetch the document by id from the text alone.
|
||||||
"""
|
"""
|
||||||
if rank is not None and total is not None:
|
if rank is not None and total is not None:
|
||||||
parts = [f"[{self.chunk_id}] [rank {rank} of {total}]"]
|
parts = [f"[{self.chunk_id}] [rank {rank} of {total}]"]
|
||||||
|
|
@ -224,6 +226,9 @@ class SearchResult(BaseModel):
|
||||||
else:
|
else:
|
||||||
parts = [f"[{self.chunk_id}] (score: {self.score:.2f})"]
|
parts = [f"[{self.chunk_id}] (score: {self.score:.2f})"]
|
||||||
|
|
||||||
|
if include_document_id and self.document_id:
|
||||||
|
parts.append(f"Document ID: {self.document_id}")
|
||||||
|
|
||||||
if include_collection and self.source:
|
if include_collection and self.source:
|
||||||
parts.append(f"Collection: {self.source}")
|
parts.append(f"Collection: {self.source}")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -52,31 +52,16 @@ def decode_picture(data: bytes, self_ref: str) -> BinaryContent | None:
|
||||||
return BinaryContent(data=data, media_type="image/png", identifier=self_ref)
|
return BinaryContent(data=data, media_type="image/png", identifier=self_ref)
|
||||||
|
|
||||||
|
|
||||||
def build_image_content_from_results(
|
def collect_pictures(
|
||||||
results: list[SearchResult],
|
results: list[SearchResult], exclude: AbstractSet[PictureKey] = frozenset()
|
||||||
include_collection: bool = False,
|
) -> tuple[list[tuple[str | None, str | None, str, BinaryContent]], set[PictureKey]]:
|
||||||
exclude: AbstractSet[PictureKey] = frozenset(),
|
"""Every distinct, decodable picture attached to ``results``, in order.
|
||||||
) -> tuple[list[str | BinaryContent], set[PictureKey]]:
|
|
||||||
"""Decode and validate picture bytes attached to search results, labelled.
|
|
||||||
|
|
||||||
Returns the labelled content and the ``PictureKey`` of every picture it
|
Returns ``(source, chunk_id, self_ref, picture)`` per picture and the
|
||||||
emitted. Dedup keyed on ``PictureKey`` so the same picture in
|
``PictureKey`` of each. Dedup keyed on ``PictureKey`` so the same picture in
|
||||||
different chunks is sent once, and a copy in another collection is its
|
different chunks is emitted once, and a copy in another collection is its
|
||||||
own; ``exclude`` seeds that dedup with pictures already sent. Pictures that fail
|
own; ``exclude`` seeds that dedup with pictures already sent. Pictures that
|
||||||
``PIL.Image.verify()`` are skipped — the model adapter renders one
|
fail ``PIL.Image.verify()`` are skipped.
|
||||||
vision placeholder per ``BinaryContent``, so emitting one for an
|
|
||||||
image the server can't decode leaves the processor with an
|
|
||||||
off-by-one count.
|
|
||||||
|
|
||||||
Every picture is preceded by a line naming the result it belongs to.
|
|
||||||
``ToolReturn.content`` reaches the model as a user-role message, so
|
|
||||||
retrieved pictures are otherwise indistinguishable from ones the user
|
|
||||||
attached, and models narrate them as part of the question: unlabelled,
|
|
||||||
gemma4-26b answered about a figure from an unrelated document, and with a
|
|
||||||
single note ahead of the batch it still called them "images in the prompt".
|
|
||||||
The label also names the chunk to cite for a figure, which
|
|
||||||
``BinaryContent.identifier`` cannot do — it does not survive serialization
|
|
||||||
to the vision API.
|
|
||||||
"""
|
"""
|
||||||
collected: list[tuple[str | None, str | None, str, BinaryContent]] = []
|
collected: list[tuple[str | None, str | None, str, BinaryContent]] = []
|
||||||
seen: set[PictureKey] = set(exclude)
|
seen: set[PictureKey] = set(exclude)
|
||||||
|
|
@ -94,7 +79,33 @@ def build_image_content_from_results(
|
||||||
collected.append((result.source, result.chunk_id, self_ref, picture))
|
collected.append((result.source, result.chunk_id, self_ref, picture))
|
||||||
seen.add(key)
|
seen.add(key)
|
||||||
emitted.add(key)
|
emitted.add(key)
|
||||||
|
return collected, emitted
|
||||||
|
|
||||||
|
|
||||||
|
def build_image_content_from_results(
|
||||||
|
results: list[SearchResult],
|
||||||
|
include_collection: bool = False,
|
||||||
|
exclude: AbstractSet[PictureKey] = frozenset(),
|
||||||
|
) -> tuple[list[str | BinaryContent], set[PictureKey]]:
|
||||||
|
"""Decode and validate picture bytes attached to search results, labelled.
|
||||||
|
|
||||||
|
Returns the labelled content and the ``PictureKey`` of every picture it
|
||||||
|
emitted, as ``collect_pictures`` decides them. An undecodable picture is
|
||||||
|
skipped because the model adapter renders one vision placeholder per
|
||||||
|
``BinaryContent``, so emitting one for an image the server can't decode
|
||||||
|
leaves the processor with an off-by-one count.
|
||||||
|
|
||||||
|
Every picture is preceded by a line naming the result it belongs to.
|
||||||
|
``ToolReturn.content`` reaches the model as a user-role message, so
|
||||||
|
retrieved pictures are otherwise indistinguishable from ones the user
|
||||||
|
attached, and models narrate them as part of the question: unlabelled,
|
||||||
|
gemma4-26b answered about a figure from an unrelated document, and with a
|
||||||
|
single note ahead of the batch it still called them "images in the prompt".
|
||||||
|
The label also names the chunk to cite for a figure, which
|
||||||
|
``BinaryContent.identifier`` cannot do — it does not survive serialization
|
||||||
|
to the vision API.
|
||||||
|
"""
|
||||||
|
collected, emitted = collect_pictures(results, exclude)
|
||||||
content: list[str | BinaryContent] = []
|
content: list[str | BinaryContent] = []
|
||||||
total = len(collected)
|
total = len(collected)
|
||||||
for position, (source, chunk_id, self_ref, picture) in enumerate(collected, 1):
|
for position, (source, chunk_id, self_ref, picture) in enumerate(collected, 1):
|
||||||
|
|
|
||||||
|
|
@ -414,6 +414,17 @@ def test_search_result_format_for_agent_source_line(fields, expected_source):
|
||||||
assert expected_source in result.format_for_agent()
|
assert expected_source in result.format_for_agent()
|
||||||
|
|
||||||
|
|
||||||
|
def test_search_result_format_for_agent_document_id_is_opt_in():
|
||||||
|
"""The capabilities' rendering is unchanged; only a caller that asks gets
|
||||||
|
the id it will fetch the document by."""
|
||||||
|
result = SearchResult(content="x", score=0.5, chunk_id="c1", document_id="doc-1")
|
||||||
|
|
||||||
|
assert "Document ID" not in result.format_for_agent(rank=1, total=1)
|
||||||
|
assert "Document ID: doc-1" in result.format_for_agent(
|
||||||
|
rank=1, total=1, include_document_id=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"labels,expected",
|
"labels,expected",
|
||||||
[
|
[
|
||||||
|
|
|
||||||
|
|
@ -97,22 +97,38 @@ async def _call(mcp, name, **kwargs):
|
||||||
return await client.call_tool(name, kwargs, raise_on_error=False)
|
return await client.call_tool(name, kwargs, raise_on_error=False)
|
||||||
|
|
||||||
|
|
||||||
|
def _results(search_result) -> list[dict]:
|
||||||
|
"""The search results a tool returned, as the client sees them."""
|
||||||
|
return search_result.structured_content["result"]
|
||||||
|
|
||||||
|
|
||||||
|
def _png_b64() -> str:
|
||||||
|
import base64
|
||||||
|
from io import BytesIO
|
||||||
|
|
||||||
|
from PIL import Image as PILImage
|
||||||
|
|
||||||
|
buf = BytesIO()
|
||||||
|
PILImage.new("RGB", (4, 4), "red").save(buf, format="PNG")
|
||||||
|
return base64.b64encode(buf.getvalue()).decode()
|
||||||
|
|
||||||
|
|
||||||
class TestMCPReadTools:
|
class TestMCPReadTools:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_search_documents(self, mcp_db):
|
async def test_search_documents(self, mcp_db):
|
||||||
mcp = create_mcp_server(mcp_db)
|
mcp = create_mcp_server(mcp_db)
|
||||||
search = await _get_tool(mcp, "search_documents")
|
search = await _get_tool(mcp, "search_documents")
|
||||||
|
|
||||||
results = await search(query="artificial intelligence")
|
results = _results(await search(query="artificial intelligence"))
|
||||||
assert len(results) > 0
|
assert len(results) > 0
|
||||||
assert all(isinstance(r, SearchResult) for r in results)
|
assert all(r["chunk_id"] and r["content"] for r in results)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_search_documents_with_limit(self, mcp_db):
|
async def test_search_documents_with_limit(self, mcp_db):
|
||||||
mcp = create_mcp_server(mcp_db)
|
mcp = create_mcp_server(mcp_db)
|
||||||
search = await _get_tool(mcp, "search_documents")
|
search = await _get_tool(mcp, "search_documents")
|
||||||
|
|
||||||
results = await search(query="artificial intelligence", limit=1)
|
results = _results(await search(query="artificial intelligence", limit=1))
|
||||||
assert len(results) == 1
|
assert len(results) == 1
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -433,6 +449,155 @@ class TestMCPDocumentNavigation:
|
||||||
await outline(document_id=doc.id, source="alpha")
|
await outline(document_id=doc.id, source="alpha")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.filterwarnings("ignore:Found propagated trace context:RuntimeWarning")
|
||||||
|
class TestMCPSearchResultShape:
|
||||||
|
"""Text as the in-process agents read it, one image per distinct picture,
|
||||||
|
and the results as structured content without picture bytes."""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _serve(monkeypatch, results):
|
||||||
|
async def fake_search(self, *args, **kwargs):
|
||||||
|
return results
|
||||||
|
|
||||||
|
monkeypatch.setattr(HaikuRAG, "search", fake_search)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_text_ranks_then_one_image_per_distinct_picture(
|
||||||
|
self, mcp_db, monkeypatch
|
||||||
|
):
|
||||||
|
from mcp.types import ImageContent, TextContent
|
||||||
|
|
||||||
|
shared = {"#/pictures/0": _png_b64()}
|
||||||
|
self._serve(
|
||||||
|
monkeypatch,
|
||||||
|
[
|
||||||
|
SearchResult(
|
||||||
|
content="a",
|
||||||
|
score=0.9,
|
||||||
|
chunk_id="c1",
|
||||||
|
document_id="d1",
|
||||||
|
image_data=shared,
|
||||||
|
),
|
||||||
|
SearchResult(
|
||||||
|
content="b",
|
||||||
|
score=0.8,
|
||||||
|
chunk_id="c2",
|
||||||
|
document_id="d1",
|
||||||
|
image_data=shared,
|
||||||
|
),
|
||||||
|
SearchResult(
|
||||||
|
content="c",
|
||||||
|
score=0.7,
|
||||||
|
chunk_id="c3",
|
||||||
|
document_id="d2",
|
||||||
|
image_data={"#/pictures/3": _png_b64()},
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await _call(create_mcp_server(mcp_db), "search_documents", query="q")
|
||||||
|
|
||||||
|
text, *rest = result.content
|
||||||
|
assert isinstance(text, TextContent)
|
||||||
|
assert "[rank 1 of 3]" in text.text and "[rank 3 of 3]" in text.text
|
||||||
|
assert "score" not in text.text
|
||||||
|
assert "Document ID: d1" in text.text
|
||||||
|
images = [block for block in rest if isinstance(block, ImageContent)]
|
||||||
|
labels = [block.text for block in rest if isinstance(block, TextContent)]
|
||||||
|
assert len(images) == 2
|
||||||
|
assert all(image.mimeType == "image/png" for image in images)
|
||||||
|
assert [
|
||||||
|
label for label in labels if "[c1]" in label and "#/pictures/0" in label
|
||||||
|
]
|
||||||
|
assert [
|
||||||
|
label for label in labels if "[c3]" in label and "#/pictures/3" in label
|
||||||
|
]
|
||||||
|
structured = _results(result)
|
||||||
|
assert [r["chunk_id"] for r in structured] == ["c1", "c2", "c3"]
|
||||||
|
assert all("image_data" not in r for r in structured)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_an_undecodable_picture_yields_no_image(self, mcp_db, monkeypatch):
|
||||||
|
import base64
|
||||||
|
|
||||||
|
self._serve(
|
||||||
|
monkeypatch,
|
||||||
|
[
|
||||||
|
SearchResult(
|
||||||
|
content="a",
|
||||||
|
score=0.9,
|
||||||
|
chunk_id="c1",
|
||||||
|
document_id="d1",
|
||||||
|
image_data={
|
||||||
|
"#/pictures/0": base64.b64encode(b"not a png").decode()
|
||||||
|
},
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await _call(create_mcp_server(mcp_db), "search_documents", query="q")
|
||||||
|
|
||||||
|
assert len(result.content) == 1
|
||||||
|
assert "[rank 1 of 1]" in result.content[0].text
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_no_results_says_so(self, mcp_db, monkeypatch):
|
||||||
|
self._serve(monkeypatch, [])
|
||||||
|
|
||||||
|
result = await _call(create_mcp_server(mcp_db), "search_documents", query="q")
|
||||||
|
|
||||||
|
assert [block.text for block in result.content] == ["No results found."]
|
||||||
|
assert _results(result) == []
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_search_text_alone_drives_the_document_tools(self, two_dbs):
|
||||||
|
"""Over two databases, every result's `Document ID` and `Collection`
|
||||||
|
parsed from the text are working arguments for the outline and
|
||||||
|
section tools."""
|
||||||
|
import re
|
||||||
|
|
||||||
|
from haiku.rag.store.models.document_item import DocumentItem
|
||||||
|
|
||||||
|
for name in ("alpha", "beta"):
|
||||||
|
async with HaikuRAG(config=two_dbs, sources=[name]) as rag:
|
||||||
|
[doc] = await rag.list_documents()
|
||||||
|
await rag.document_item_repository.create_items(
|
||||||
|
doc.id,
|
||||||
|
[
|
||||||
|
DocumentItem(
|
||||||
|
document_id=doc.id,
|
||||||
|
position=0,
|
||||||
|
self_ref="#/texts/0",
|
||||||
|
label="section_header",
|
||||||
|
text=f"Heading in {name}",
|
||||||
|
heading_level=1,
|
||||||
|
)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
mcp = _covering_all(two_dbs)
|
||||||
|
|
||||||
|
search = await _call(mcp, "search_documents", query="cats")
|
||||||
|
pairs = re.findall(
|
||||||
|
r"Document ID: (\S+)\nCollection: (\S+)", search.content[0].text
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(pairs) == len(_results(search)) == 2
|
||||||
|
assert {source for _, source in pairs} == {"alpha", "beta"}
|
||||||
|
for document_id, source in pairs:
|
||||||
|
outline = await _call(
|
||||||
|
mcp, "get_document_outline", document_id=document_id, source=source
|
||||||
|
)
|
||||||
|
[node] = _results(outline)
|
||||||
|
section = await _call(
|
||||||
|
mcp,
|
||||||
|
"get_document_section",
|
||||||
|
document_id=document_id,
|
||||||
|
section_id=node["id"],
|
||||||
|
source=source,
|
||||||
|
)
|
||||||
|
assert section.structured_content["title"] == f"Heading in {source}"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.filterwarnings("ignore:Found propagated trace context:RuntimeWarning")
|
@pytest.mark.filterwarnings("ignore:Found propagated trace context:RuntimeWarning")
|
||||||
class TestMCPDescribesItself:
|
class TestMCPDescribesItself:
|
||||||
"""What a client learns from initialize and list_tools, over the wire."""
|
"""What a client learns from initialize and list_tools, over the wire."""
|
||||||
|
|
@ -537,19 +702,19 @@ class TestMCPCoversTheConfiguredSet:
|
||||||
mcp = _covering_all(two_dbs)
|
mcp = _covering_all(two_dbs)
|
||||||
search = await _get_tool(mcp, "search_documents")
|
search = await _get_tool(mcp, "search_documents")
|
||||||
|
|
||||||
results = await search(query="cats")
|
results = _results(await search(query="cats"))
|
||||||
|
|
||||||
assert {r.source for r in results} == {"alpha", "beta"}
|
assert {r["source"] for r in results} == {"alpha", "beta"}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_sources_narrows_the_search(self, two_dbs):
|
async def test_sources_narrows_the_search(self, two_dbs):
|
||||||
mcp = _covering_all(two_dbs)
|
mcp = _covering_all(two_dbs)
|
||||||
search = await _get_tool(mcp, "search_documents")
|
search = await _get_tool(mcp, "search_documents")
|
||||||
|
|
||||||
results = await search(query="cats", sources=["beta"])
|
results = _results(await search(query="cats", sources=["beta"]))
|
||||||
|
|
||||||
assert results
|
assert results
|
||||||
assert {r.source for r in results} == {"beta"}
|
assert {r["source"] for r in results} == {"beta"}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
|
|
@ -583,12 +748,13 @@ class TestMCPCoversTheConfiguredSet:
|
||||||
mcp = _covering_all(two_dbs)
|
mcp = _covering_all(two_dbs)
|
||||||
search = await _get_tool(mcp, "search_documents")
|
search = await _get_tool(mcp, "search_documents")
|
||||||
|
|
||||||
results = await search(
|
results = _results(
|
||||||
query="cats", filter="uri LIKE '%beta%'", sources=["beta"]
|
await search(query="cats", filter="uri LIKE '%beta%'", sources=["beta"])
|
||||||
)
|
)
|
||||||
assert results
|
assert results
|
||||||
assert {r.source for r in results} == {"beta"}
|
assert {r["source"] for r in results} == {"beta"}
|
||||||
assert await search(query="cats", filter="uri LIKE '%beta%'", sources=[]) == []
|
none = await search(query="cats", filter="uri LIKE '%beta%'", sources=[])
|
||||||
|
assert _results(none) == []
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_the_listing_covers_every_database(self, two_dbs):
|
async def test_the_listing_covers_every_database(self, two_dbs):
|
||||||
|
|
@ -617,9 +783,9 @@ class TestMCPCoversTheConfiguredSet:
|
||||||
mcp = create_mcp_server(config=two_dbs)
|
mcp = create_mcp_server(config=two_dbs)
|
||||||
search = await _get_tool(mcp, "search_documents")
|
search = await _get_tool(mcp, "search_documents")
|
||||||
|
|
||||||
results = await search(query="cats")
|
results = _results(await search(query="cats"))
|
||||||
|
|
||||||
assert {r.source for r in results} == {"alpha", "beta"}
|
assert {r["source"] for r in results} == {"alpha", "beta"}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_ask_question_names_each_citations_database(
|
async def test_ask_question_names_each_citations_database(
|
||||||
|
|
@ -713,7 +879,7 @@ class TestMCPImageQuery:
|
||||||
sources=[],
|
sources=[],
|
||||||
)
|
)
|
||||||
|
|
||||||
assert results == []
|
assert _results(results) == []
|
||||||
assert seen["query"] == png
|
assert seen["query"] == png
|
||||||
assert seen["filter"] == "uri LIKE 'x%'"
|
assert seen["filter"] == "uri LIKE 'x%'"
|
||||||
assert seen["sources"] == []
|
assert seen["sources"] == []
|
||||||
|
|
@ -1028,12 +1194,12 @@ class TestMCPClientLifetime:
|
||||||
mcp = _mcp_covering(scope, config)
|
mcp = _mcp_covering(scope, config)
|
||||||
async with mcp._lifespan_manager():
|
async with mcp._lifespan_manager():
|
||||||
search = await _get_tool(mcp, "search_documents")
|
search = await _get_tool(mcp, "search_documents")
|
||||||
results = await search(query="artificial intelligence")
|
results = _results(await search(query="artificial intelligence"))
|
||||||
listing = await _get_tool(mcp, "list_documents")
|
listing = await _get_tool(mcp, "list_documents")
|
||||||
documents = await listing()
|
documents = await listing()
|
||||||
|
|
||||||
assert results
|
assert results
|
||||||
assert {r.source for r in results} == {"alpha"}
|
assert {r["source"] for r in results} == {"alpha"}
|
||||||
titles = {d.title for d in documents}
|
titles = {d.title for d in documents}
|
||||||
assert "AI Overview" in titles
|
assert "AI Overview" in titles
|
||||||
assert "Zebras" not in titles
|
assert "Zebras" not in titles
|
||||||
|
|
@ -1118,7 +1284,7 @@ class TestMCPClientLifetime:
|
||||||
assert opens == 1
|
assert opens == 1
|
||||||
|
|
||||||
async with mcp._lifespan_manager():
|
async with mcp._lifespan_manager():
|
||||||
results = await search(query="artificial intelligence")
|
results = _results(await search(query="artificial intelligence"))
|
||||||
assert opens == 2
|
assert opens == 2
|
||||||
assert len(results) > 0
|
assert len(results) > 0
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue