Refactor MCP to use proper models, test

This commit is contained in:
Yiorgis Gozadinos 2026-02-20 17:38:53 +02:00
parent 69083c8cb4
commit d1fdbdeb21
No known key found for this signature in database
3 changed files with 207 additions and 41 deletions

View file

@ -2,26 +2,16 @@ from pathlib import Path
from typing import Any
from fastmcp import FastMCP
from pydantic import BaseModel
from haiku.rag.agents.research.models import ResearchReport
from haiku.rag.client import HaikuRAG
from haiku.rag.config import AppConfig, Config
from haiku.rag.store.models import SearchResult
from haiku.rag.store.models import Document, SearchResult
from haiku.rag.tools.document import DocumentInfo
from haiku.rag.utils import format_citations
class DocumentResult(BaseModel):
id: str | None
content: str
uri: str | None = None
title: str | None = None
metadata: dict[str, Any] = {}
created_at: str
updated_at: str
def create_mcp_server( # pragma: no cover
def create_mcp_server(
db_path: Path, config: AppConfig = Config, read_only: bool = False
) -> FastMCP:
"""Create an MCP server with the specified database path.
@ -111,24 +101,11 @@ def create_mcp_server( # pragma: no cover
return []
@mcp.tool()
async def get_document(document_id: str) -> DocumentResult | None:
async def get_document(document_id: str) -> Document | None:
"""Get a document by its ID."""
try:
async with HaikuRAG(db_path, config=config, read_only=read_only) as rag:
document = await rag.get_document_by_id(document_id)
if document is None:
return None
return DocumentResult(
id=document.id,
content=document.content,
uri=document.uri,
title=document.title,
metadata=document.metadata,
created_at=str(document.created_at),
updated_at=str(document.updated_at),
)
return await rag.get_document_by_id(document_id)
except Exception:
return None
@ -137,30 +114,24 @@ def create_mcp_server( # pragma: no cover
limit: int | None = None,
offset: int | None = None,
filter: str | None = None,
) -> list[DocumentResult]:
) -> list[DocumentInfo]:
"""List all documents with optional pagination and filtering.
Args:
limit: Maximum number of documents to return.
offset: Number of documents to skip.
filter: Optional SQL WHERE clause to filter documents.
Returns:
List of DocumentResult instances matching the criteria.
"""
try:
async with HaikuRAG(db_path, config=config, read_only=read_only) as rag:
documents = await rag.list_documents(limit, offset, filter)
return [
DocumentResult(
DocumentInfo(
id=doc.id,
content=doc.content,
uri=doc.uri,
title=doc.title,
metadata=doc.metadata,
created_at=str(doc.created_at),
updated_at=str(doc.updated_at),
title=doc.title or "Untitled",
uri=doc.uri or "",
created=doc.created_at.strftime("%Y-%m-%d"),
)
for doc in documents
]

View file

@ -43,8 +43,8 @@ class Document(BaseModel):
uri: str | None = None
title: str | None = None
metadata: dict = {}
docling_document: bytes | None = None
docling_version: str | None = None
docling_document: bytes | None = Field(default=None, exclude=True)
docling_version: str | None = Field(default=None, exclude=True)
created_at: datetime = Field(default_factory=datetime.now)
updated_at: datetime = Field(default_factory=datetime.now)

195
tests/test_mcp.py Normal file
View file

@ -0,0 +1,195 @@
import pytest
from haiku.rag.client import HaikuRAG
from haiku.rag.mcp import create_mcp_server
from haiku.rag.store.models import Document, SearchResult
from haiku.rag.tools.document import DocumentInfo
@pytest.fixture(autouse=True)
def mock_embedder(monkeypatch):
"""Monkeypatch the embedder to return deterministic vectors."""
import random
from haiku.rag.embeddings import EmbedderWrapper
async def fake_embed_query(self, text):
random.seed(hash(text) % (2**32))
return [random.random() for _ in range(2560)]
async def fake_embed_documents(self, texts):
result = []
for t in texts:
random.seed(hash(t) % (2**32))
result.append([random.random() for _ in range(2560)])
return result
monkeypatch.setattr(EmbedderWrapper, "embed_query", fake_embed_query)
monkeypatch.setattr(EmbedderWrapper, "embed_documents", fake_embed_documents)
@pytest.fixture
async def mcp_db(temp_db_path):
"""Create a test database with sample documents."""
async with HaikuRAG(temp_db_path, create=True) as rag:
await rag.create_document(
"Artificial intelligence is transforming industries worldwide.",
title="AI Overview",
uri="test://ai-overview",
)
await rag.create_document(
"Machine learning is a subset of artificial intelligence.",
title="ML Basics",
uri="test://ml-basics",
)
return temp_db_path
def _get_tool(mcp, name):
"""Get a tool function from an MCP server by name."""
for tool in mcp._tool_manager._tools.values():
if tool.fn.__name__ == name:
return tool.fn
raise ValueError(f"Tool {name!r} not found in MCP server")
class TestMCPReadTools:
@pytest.mark.asyncio
async def test_search_documents(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=True)
search = _get_tool(mcp, "search_documents")
results = await search(query="artificial intelligence")
assert len(results) > 0
assert all(isinstance(r, SearchResult) for r in results)
@pytest.mark.asyncio
async def test_search_documents_with_limit(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=True)
search = _get_tool(mcp, "search_documents")
results = await search(query="artificial intelligence", limit=1)
assert len(results) == 1
@pytest.mark.asyncio
async def test_get_document(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=True)
get_doc = _get_tool(mcp, "get_document")
# First get the ID via list
list_docs = _get_tool(mcp, "list_documents")
docs = await list_docs()
doc_id = docs[0].id
result = await get_doc(document_id=doc_id)
assert isinstance(result, Document)
assert result.content != ""
assert result.title is not None
@pytest.mark.asyncio
async def test_get_document_excludes_docling_fields(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=True)
get_doc = _get_tool(mcp, "get_document")
list_docs = _get_tool(mcp, "list_documents")
docs = await list_docs()
doc_id = docs[0].id
result = await get_doc(document_id=doc_id)
serialized = result.model_dump(mode="json")
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, read_only=True)
get_doc = _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, read_only=True)
list_docs = _get_tool(mcp, "list_documents")
results = await list_docs()
assert len(results) == 2
assert all(isinstance(r, DocumentInfo) for r in results)
@pytest.mark.asyncio
async def test_list_documents_with_limit(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=True)
list_docs = _get_tool(mcp, "list_documents")
results = await list_docs(limit=1)
assert len(results) == 1
@pytest.mark.asyncio
async def test_list_documents_with_filter(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=True)
list_docs = _get_tool(mcp, "list_documents")
results = await list_docs(filter="title = 'AI Overview'")
assert len(results) == 1
assert results[0].title == "AI Overview"
class TestMCPWriteTools:
@pytest.mark.asyncio
async def test_write_tools_registered_when_not_read_only(self, temp_db_path):
async with HaikuRAG(temp_db_path, create=True):
pass
mcp = create_mcp_server(temp_db_path, read_only=False)
tool_names = list(mcp._tool_manager._tools.keys())
assert "add_document_from_text" in tool_names
assert "add_document_from_file" in tool_names
assert "add_document_from_url" in tool_names
assert "delete_document" in tool_names
@pytest.mark.asyncio
async def test_write_tools_not_registered_when_read_only(self, temp_db_path):
async with HaikuRAG(temp_db_path, create=True):
pass
mcp = create_mcp_server(temp_db_path, read_only=True)
tool_names = list(mcp._tool_manager._tools.keys())
assert "add_document_from_text" not in tool_names
assert "delete_document" not in tool_names
@pytest.mark.asyncio
async def test_add_document_from_text(self, temp_db_path):
async with HaikuRAG(temp_db_path, create=True):
pass
mcp = create_mcp_server(temp_db_path, read_only=False)
add_text = _get_tool(mcp, "add_document_from_text")
doc_id = await add_text(content="Test content for MCP", title="MCP Test Doc")
assert doc_id is not None
get_doc = _get_tool(mcp, "get_document")
doc = await get_doc(document_id=doc_id)
assert doc.title == "MCP Test Doc"
assert doc.content == "Test content for MCP"
@pytest.mark.asyncio
async def test_delete_document(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=False)
list_docs = _get_tool(mcp, "list_documents")
delete_doc = _get_tool(mcp, "delete_document")
docs = await list_docs()
assert len(docs) == 2
result = await delete_doc(document_id=docs[0].id)
assert result is True
docs_after = await list_docs()
assert len(docs_after) == 1
@pytest.mark.asyncio
async def test_delete_document_not_found(self, mcp_db):
mcp = create_mcp_server(mcp_db, read_only=False)
delete_doc = _get_tool(mcp, "delete_document")
result = await delete_doc(document_id="nonexistent-id")
assert result is False