Refactor MCP to use proper models, test
This commit is contained in:
parent
69083c8cb4
commit
d1fdbdeb21
3 changed files with 207 additions and 41 deletions
|
|
@ -2,26 +2,16 @@ from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from fastmcp import FastMCP
|
from fastmcp import FastMCP
|
||||||
from pydantic import BaseModel
|
|
||||||
|
|
||||||
from haiku.rag.agents.research.models import ResearchReport
|
from haiku.rag.agents.research.models import ResearchReport
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import AppConfig, Config
|
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
|
from haiku.rag.utils import format_citations
|
||||||
|
|
||||||
|
|
||||||
class DocumentResult(BaseModel):
|
def create_mcp_server(
|
||||||
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
|
|
||||||
db_path: Path, config: AppConfig = Config, read_only: bool = False
|
db_path: Path, config: AppConfig = Config, read_only: bool = False
|
||||||
) -> FastMCP:
|
) -> FastMCP:
|
||||||
"""Create an MCP server with the specified database path.
|
"""Create an MCP server with the specified database path.
|
||||||
|
|
@ -111,24 +101,11 @@ def create_mcp_server( # pragma: no cover
|
||||||
return []
|
return []
|
||||||
|
|
||||||
@mcp.tool()
|
@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."""
|
"""Get a document by its ID."""
|
||||||
try:
|
try:
|
||||||
async with HaikuRAG(db_path, config=config, read_only=read_only) as rag:
|
async with HaikuRAG(db_path, config=config, read_only=read_only) as rag:
|
||||||
document = await rag.get_document_by_id(document_id)
|
return 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),
|
|
||||||
)
|
|
||||||
except Exception:
|
except Exception:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
@ -137,30 +114,24 @@ def create_mcp_server( # pragma: no cover
|
||||||
limit: int | None = None,
|
limit: int | None = None,
|
||||||
offset: int | None = None,
|
offset: int | None = None,
|
||||||
filter: str | None = None,
|
filter: str | None = None,
|
||||||
) -> list[DocumentResult]:
|
) -> list[DocumentInfo]:
|
||||||
"""List all documents with optional pagination and filtering.
|
"""List all documents with optional pagination and filtering.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
limit: Maximum number of documents to return.
|
limit: Maximum number of documents to return.
|
||||||
offset: Number of documents to skip.
|
offset: Number of documents to skip.
|
||||||
filter: Optional SQL WHERE clause to filter documents.
|
filter: Optional SQL WHERE clause to filter documents.
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of DocumentResult instances matching the criteria.
|
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
async with HaikuRAG(db_path, config=config, read_only=read_only) as rag:
|
async with HaikuRAG(db_path, config=config, read_only=read_only) as rag:
|
||||||
documents = await rag.list_documents(limit, offset, filter)
|
documents = await rag.list_documents(limit, offset, filter)
|
||||||
|
|
||||||
return [
|
return [
|
||||||
DocumentResult(
|
DocumentInfo(
|
||||||
id=doc.id,
|
id=doc.id,
|
||||||
content=doc.content,
|
title=doc.title or "Untitled",
|
||||||
uri=doc.uri,
|
uri=doc.uri or "",
|
||||||
title=doc.title,
|
created=doc.created_at.strftime("%Y-%m-%d"),
|
||||||
metadata=doc.metadata,
|
|
||||||
created_at=str(doc.created_at),
|
|
||||||
updated_at=str(doc.updated_at),
|
|
||||||
)
|
)
|
||||||
for doc in documents
|
for doc in documents
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -43,8 +43,8 @@ class Document(BaseModel):
|
||||||
uri: str | None = None
|
uri: str | None = None
|
||||||
title: str | None = None
|
title: str | None = None
|
||||||
metadata: dict = {}
|
metadata: dict = {}
|
||||||
docling_document: bytes | None = None
|
docling_document: bytes | None = Field(default=None, exclude=True)
|
||||||
docling_version: str | None = None
|
docling_version: str | None = Field(default=None, exclude=True)
|
||||||
created_at: datetime = Field(default_factory=datetime.now)
|
created_at: datetime = Field(default_factory=datetime.now)
|
||||||
updated_at: datetime = Field(default_factory=datetime.now)
|
updated_at: datetime = Field(default_factory=datetime.now)
|
||||||
|
|
||||||
|
|
|
||||||
195
tests/test_mcp.py
Normal file
195
tests/test_mcp.py
Normal 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
|
||||||
Loading…
Reference in a new issue