from __future__ import annotations import shutil import sys import tempfile import unittest from pathlib import Path from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import stdio_client from mcp.shared.memory import create_connected_server_and_client_session from docforge.index import ProjectIndex from docforge.mcp_server import CONTENT_WARNING, READ_TOOLS, create_server from docforge.project import Project ROOT = Path(__file__).resolve().parents[1] FIXTURES = ROOT / "tests" / "fixtures" class DocForgeMcpTests(unittest.IsolatedAsyncioTestCase): def copy_fixture(self, name: str, destination: Path) -> Path: root = destination / name shutil.copytree(FIXTURES / name, root) return root async def test_protocol_lists_only_the_fixed_read_surface(self) -> None: with tempfile.TemporaryDirectory() as directory: root = self.copy_fixture("alpha", Path(directory)) ProjectIndex(Project.open(root)).build() async with create_connected_server_and_client_session( create_server(root), raise_exceptions=True ) as session: response = await session.list_tools() names = tuple(tool.name for tool in response.tools) self.assertEqual(READ_TOOLS, names) self.assertFalse( any( token in name for name in names for token in ("write", "apply", "commit", "push", "deploy", "propose") ) ) async def test_every_read_tool_returns_scoped_structured_results(self) -> None: with tempfile.TemporaryDirectory() as directory: root = self.copy_fixture("alpha", Path(directory)) ProjectIndex(Project.open(root)).build() calls = ( ("docforge_project_info", {}), ("docforge_get_contract", {}), ("docforge_get_node", {"node_id": "guide.workflow"}), ("docforge_search", {"query": "canonical nodes", "limit": 5}), ("docforge_filter_nodes", {"family": "proof", "tag": "validation"}), ("docforge_backlinks", {"node_id": "guide.workflow"}), ("docforge_dependencies", {"node_id": "guide.workflow", "depth": 2}), ("docforge_impact", {"node_id": "guide.foundation", "depth": 2}), ("docforge_get_context", {"profile": "active", "budget": 180}), ("docforge_validate_project", {}), ("docforge_render_status", {}), ) async with create_connected_server_and_client_session( create_server(root), raise_exceptions=True ) as session: results = [await session.call_tool(name, arguments) for name, arguments in calls] for result in results: self.assertFalse(result.isError) self.assertIsNotNone(result.structuredContent) payload = result.structuredContent self.assertEqual("ok", payload["status"]) self.assertEqual("alpha-docs", payload["project_id"]) self.assertEqual(CONTENT_WARNING, payload["content_warning"]) self.assertTrue(payload["project_root_fingerprint"]) self.assertEqual("current", payload["staleness"]) contract = results[1].structuredContent self.assertFalse(contract["canonical_writes_allowed"]) self.assertFalse(contract["project_switching_allowed"]) self.assertIn("canonical_writes", contract["excluded_operations"]) context = results[8].structuredContent self.assertLessEqual(context["estimated_tokens"], 180) self.assertTrue(context["omissions"]) async def test_missing_node_and_stale_index_are_structured_failures(self) -> None: with tempfile.TemporaryDirectory() as directory: root = self.copy_fixture("alpha", Path(directory)) ProjectIndex(Project.open(root)).build() server = create_server(root) async with create_connected_server_and_client_session( server, raise_exceptions=True ) as session: missing = await session.call_tool( "docforge_get_node", {"node_id": "research.question"} ) workflow = root / "docs" / "content" / "workflow.md" workflow.write_text( workflow.read_text(encoding="utf-8") + "\nChanged after startup.\n", encoding="utf-8", ) stale = await session.call_tool("docforge_get_node", {"node_id": "guide.workflow"}) self.assertEqual("missing_node", missing.structuredContent["error"]["code"]) self.assertEqual("stale_index", stale.structuredContent["error"]["code"]) self.assertEqual("current", missing.structuredContent["staleness"]) self.assertEqual("stale", stale.structuredContent["staleness"]) self.assertTrue(missing.structuredContent["source_hash"]) self.assertTrue(stale.structuredContent["source_hash"]) self.assertFalse(missing.isError) self.assertFalse(stale.isError) async def test_output_limit_fails_without_returning_partial_content(self) -> None: with tempfile.TemporaryDirectory() as directory: root = self.copy_fixture("alpha", Path(directory)) descriptor = root / ".docforge" / "project.toml" descriptor.write_text( descriptor.read_text(encoding="utf-8").replace( "max_context_tokens = 2000", "max_context_tokens = 2000\nmax_tool_output_chars = 700", ), encoding="utf-8", ) ProjectIndex(Project.open(root)).build() async with create_connected_server_and_client_session( create_server(root), raise_exceptions=True ) as session: result = await session.call_tool("docforge_get_contract", {}) payload = result.structuredContent self.assertEqual("error", payload["status"]) self.assertEqual("result_too_large", payload["error"]["code"]) self.assertNotIn("canonical_paths", payload) async def test_stdio_transport_serves_the_same_project_bound_contract(self) -> None: with tempfile.TemporaryDirectory() as directory: root = self.copy_fixture("beta", Path(directory)) ProjectIndex(Project.open(root)).build() parameters = StdioServerParameters( command=sys.executable, args=["-m", "docforge.mcp_server", "--project-root", str(root)], ) async with ( stdio_client(parameters) as (read, write), ClientSession(read, write) as session, ): await session.initialize() result = await session.call_tool("docforge_project_info", {}) self.assertFalse(result.isError) self.assertEqual("beta-notes", result.structuredContent["project_id"]) self.assertEqual(1, result.structuredContent["node_count"]) if __name__ == "__main__": unittest.main()