diff --git a/tests/test_mcp.py b/tests/test_mcp.py new file mode 100644 index 00000000..02b1b275 --- /dev/null +++ b/tests/test_mcp.py @@ -0,0 +1,322 @@ +import tempfile +from pathlib import Path +from unittest.mock import AsyncMock, patch + +import pytest + +from haiku.rag.mcp import create_mcp_server +from haiku.rag.research.models import ResearchReport +from haiku.rag.store.models.document import Document + + +@pytest.mark.asyncio +async def test_mcp_add_document_from_file(): + """Test add_document_from_file tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + mock_doc = Document(content="test", uri="file:///test.txt") + mock_doc.id = "doc123" + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.create_document_from_source = AsyncMock(return_value=mock_doc) + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + add_file_tool = next( + t for t in tools.values() if t.name == "add_document_from_file" + ) + + result = await add_file_tool.fn(file_path="/test.txt") # type: ignore[attr-defined] + + assert result == "doc123" + mock_rag.create_document_from_source.assert_called_once() + + +@pytest.mark.asyncio +async def test_mcp_ask_question(): + """Test ask_question tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.ask = AsyncMock(return_value="This is the answer") + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + ask_tool = next(t for t in tools.values() if t.name == "ask_question") + + result = await ask_tool.fn( # type: ignore[attr-defined] + question="What is this?", cite=False, deep=False + ) + + assert result == "This is the answer" + mock_rag.ask.assert_called_once_with("What is this?", cite=False) + + +@pytest.mark.asyncio +async def test_mcp_add_document_from_url(): + """Test add_document_from_url tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + mock_doc = Document(content="test", uri="https://example.com") + mock_doc.id = "doc456" + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.create_document_from_source = AsyncMock(return_value=mock_doc) + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + add_url_tool = next( + t for t in tools.values() if t.name == "add_document_from_url" + ) + + result = await add_url_tool.fn(url="https://example.com") # type: ignore[attr-defined] + + assert result == "doc456" + mock_rag.create_document_from_source.assert_called_once() + + +@pytest.mark.asyncio +async def test_mcp_add_document_from_text(): + """Test add_document_from_text tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + mock_doc = Document(content="test content", uri="text://test") + mock_doc.id = "doc789" + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.create_document = AsyncMock(return_value=mock_doc) + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + add_text_tool = next( + t for t in tools.values() if t.name == "add_document_from_text" + ) + + result = await add_text_tool.fn( # type: ignore[attr-defined] + content="test content", uri="text://test" + ) + + assert result == "doc789" + mock_rag.create_document.assert_called_once() + + +@pytest.mark.asyncio +async def test_mcp_search_documents(): + """Test search_documents tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + from haiku.rag.store.models.chunk import Chunk + + mock_chunk1 = Chunk(content="Result 1", document_id="doc1") + mock_chunk2 = Chunk(content="Result 2", document_id="doc2") + mock_results = [(mock_chunk1, 0.9), (mock_chunk2, 0.8)] + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.search = AsyncMock(return_value=mock_results) + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + search_tool = next( + t for t in tools.values() if t.name == "search_documents" + ) + + result = await search_tool.fn(query="test query", limit=5) # type: ignore[attr-defined] + + assert len(result) == 2 + assert result[0].document_id == "doc1" + assert result[0].content == "Result 1" + assert result[0].score == 0.9 + assert result[1].document_id == "doc2" + mock_rag.search.assert_called_once_with("test query", 5) + + +@pytest.mark.asyncio +async def test_mcp_get_document(): + """Test get_document tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + from datetime import UTC, datetime + + mock_doc = Document(content="test", uri="file:///test.txt", title="Test Doc") + mock_doc.id = "doc123" + mock_doc.created_at = datetime(2024, 1, 1, tzinfo=UTC) + mock_doc.updated_at = datetime(2024, 1, 2, tzinfo=UTC) + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.get_document_by_id = AsyncMock(return_value=mock_doc) + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + get_tool = next(t for t in tools.values() if t.name == "get_document") + + result = await get_tool.fn(document_id="doc123") # type: ignore[attr-defined] + + assert result is not None + assert result.id == "doc123" + assert result.content == "test" + assert result.title == "Test Doc" + mock_rag.get_document_by_id.assert_called_once_with("doc123") + + +@pytest.mark.asyncio +async def test_mcp_list_documents(): + """Test list_documents tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + from datetime import UTC, datetime + + mock_doc1 = Document(content="test1", uri="file:///test1.txt") + mock_doc1.id = "doc1" + mock_doc1.created_at = datetime(2024, 1, 1, tzinfo=UTC) + mock_doc1.updated_at = datetime(2024, 1, 1, tzinfo=UTC) + + mock_doc2 = Document(content="test2", uri="file:///test2.txt") + mock_doc2.id = "doc2" + mock_doc2.created_at = datetime(2024, 1, 2, tzinfo=UTC) + mock_doc2.updated_at = datetime(2024, 1, 2, tzinfo=UTC) + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.list_documents = AsyncMock(return_value=[mock_doc1, mock_doc2]) + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + list_tool = next(t for t in tools.values() if t.name == "list_documents") + + result = await list_tool.fn(limit=10, offset=0) # type: ignore[attr-defined] + + assert len(result) == 2 + assert result[0].id == "doc1" + assert result[1].id == "doc2" + mock_rag.list_documents.assert_called_once_with(10, 0) + + +@pytest.mark.asyncio +async def test_mcp_delete_document(): + """Test delete_document tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + with patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class: + mock_rag = AsyncMock() + mock_rag.delete_document = AsyncMock(return_value=True) + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + tools = await mcp.get_tools() + delete_tool = next(t for t in tools.values() if t.name == "delete_document") + + result = await delete_tool.fn(document_id="doc123") # type: ignore[attr-defined] + + assert result is True + mock_rag.delete_document.assert_called_once_with("doc123") + + +@pytest.mark.asyncio +async def test_mcp_ask_question_deep(): + """Test ask_question tool with deep=True is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + with ( + patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class, + patch("haiku.rag.qa.deep.graph.build_deep_qa_graph") as mock_graph_builder, + ): + mock_rag = AsyncMock() + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + mock_graph = AsyncMock() + mock_result = AsyncMock() + mock_result.output.answer = "Deep answer" + mock_graph.run = AsyncMock(return_value=mock_result) + mock_graph_builder.return_value = mock_graph + + tools = await mcp.get_tools() + ask_tool = next(t for t in tools.values() if t.name == "ask_question") + + result = await ask_tool.fn( # type: ignore[attr-defined] + question="Deep question?", cite=True, deep=True + ) + + assert result == "Deep answer" + mock_graph.run.assert_called_once() + + +@pytest.mark.asyncio +async def test_mcp_research_question(): + """Test research_question tool is properly wired.""" + with tempfile.TemporaryDirectory() as temp_dir: + db_path = Path(temp_dir) / "test.lancedb" + mcp = create_mcp_server(db_path) + + mock_report = ResearchReport( + title="Research Title", + executive_summary="Summary", + main_findings=["Finding 1"], + conclusions=["Conclusion 1"], + recommendations=["Recommendation 1"], + sources_summary="Sources used", + ) + + with ( + patch("haiku.rag.mcp.HaikuRAG") as mock_rag_class, + patch( + "haiku.rag.research.graph.build_research_graph" + ) as mock_graph_builder, + ): + mock_rag = AsyncMock() + mock_rag_class.return_value.__aenter__ = AsyncMock(return_value=mock_rag) + mock_rag_class.return_value.__aexit__ = AsyncMock(return_value=None) + + mock_graph = AsyncMock() + mock_result = AsyncMock() + mock_result.output = mock_report + mock_graph.run = AsyncMock(return_value=mock_result) + mock_graph_builder.return_value = mock_graph + + tools = await mcp.get_tools() + research_tool = next( + t for t in tools.values() if t.name == "research_question" + ) + + result = await research_tool.fn( # type: ignore[attr-defined] + question="Research question?", + max_iterations=1, + confidence_threshold=0.5, + max_concurrency=1, + ) + + assert result is not None + assert result.title == "Research Title" + assert result.executive_summary == "Summary" + mock_graph.run.assert_called_once()