diff --git a/src/server.py b/src/server.py index d843d70..db852bd 100644 --- a/src/server.py +++ b/src/server.py @@ -20,6 +20,7 @@ import asyncmy import anyio from fastmcp import FastMCP, Context +from mcp.types import ToolAnnotations # Import custom connection pool that disables MULTI_STATEMENTS from custom_connection import create_safe_pool @@ -937,58 +938,60 @@ def register_tools(self): logger.error("Cannot register tools: Database pool is not initialized.") raise RuntimeError("Database pool must be initialized before registering tools.") - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=True)) async def list_databases() -> List[str]: """Lists all accessible databases on the connected MariaDB server.""" return await self.list_databases() - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=True)) async def list_tables(database_name: str) -> List[str]: """Lists all tables within the specified database.""" return await self.list_tables(database_name) - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=True)) async def get_table_schema(database_name: str, table_name: str) -> Dict[str, Any]: """Retrieves the schema for a specific table in a database.""" return await self.get_table_schema(database_name, table_name) - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=True)) async def get_table_schema_with_relations(database_name: str, table_name: str) -> Dict[str, Any]: """Retrieves table schema with foreign key relationship information.""" return await self.get_table_schema_with_relations(database_name, table_name) - @self.mcp.tool + # In read-only mode, execute_sql enforces a statement allowlist, so the hint + # reports enforced behaviour rather than promising it. + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=self.is_read_only)) async def execute_sql(sql_query: str, database_name: str, parameters: Optional[List[Any]] = None) -> List[Dict[str, Any]]: """Executes a read-only SQL query against a specified database.""" return await self.execute_sql(sql_query, database_name, parameters) - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=False, destructiveHint=False)) async def create_database(database_name: str) -> Dict[str, Any]: """Creates a new database if it doesn't exist.""" return await self.create_database(database_name) if EMBEDDING_PROVIDER is not None: - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=False, destructiveHint=False)) async def create_vector_store(database_name: str, vector_store_name: str, model_name: Optional[str] = None, distance_function: Optional[str] = None) -> dict: """Creates a table which stores embeddings.""" return await self.create_vector_store(database_name, vector_store_name, model_name, distance_function) - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=True)) async def list_vector_stores(database_name: str) -> List[str]: """Lists all vector stores in a database.""" return await self.list_vector_stores(database_name) - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=False, destructiveHint=True)) async def delete_vector_store(database_name: str, vector_store_name: str) -> Dict[str, Any]: """Deletes a vector store from the specified database.""" return await self.delete_vector_store(database_name, vector_store_name) - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=False, destructiveHint=False)) async def insert_docs_vector_store(database_name: str, vector_store_name: str, documents: List[str], metadata: Optional[List[dict]] = None) -> dict: """Insert a batch of documents into a vector store.""" return await self.insert_docs_vector_store(database_name, vector_store_name, documents, metadata) - @self.mcp.tool + @self.mcp.tool(annotations=ToolAnnotations(readOnlyHint=True)) async def search_vector_store(user_query: str, database_name: str, vector_store_name: str, k: int = 7) -> list: """Search a vector store for similar documents.""" return await self.search_vector_store(user_query, database_name, vector_store_name, k) diff --git a/src/tests/test_mcp_server.py b/src/tests/test_mcp_server.py index 35b4f54..2c9feeb 100644 --- a/src/tests/test_mcp_server.py +++ b/src/tests/test_mcp_server.py @@ -259,5 +259,34 @@ async def test_readonly_mode(self): pass tg.cancel_scope.cancel() + async def annotations_by_tool_name(self): + async with anyio.create_task_group() as tg: + await self.task_group_helper(tg) + async with self.client: + tools = await self.client.list_tools() + annotations = {tool.name: tool.annotations for tool in tools} + tg.cancel_scope.cancel() + return annotations + + async def test_read_only_tools_annotated_read_only(self): + annotations = await self.annotations_by_tool_name() + for tool_name in ['list_databases', 'list_tables', 'get_table_schema', 'get_table_schema_with_relations']: + self.assertTrue(annotations[tool_name].readOnlyHint, tool_name) + + async def test_writing_tools_annotated_not_read_only(self): + annotations = await self.annotations_by_tool_name() + self.assertFalse(annotations['create_database'].readOnlyHint) + self.assertFalse(annotations['create_database'].destructiveHint) + + async def test_execute_sql_annotated_read_only_in_read_only_mode(self): + self.server.is_read_only = True + annotations = await self.annotations_by_tool_name() + self.assertTrue(annotations['execute_sql'].readOnlyHint) + + async def test_execute_sql_annotated_writing_outside_read_only_mode(self): + self.server.is_read_only = False + annotations = await self.annotations_by_tool_name() + self.assertFalse(annotations['execute_sql'].readOnlyHint) + if __name__ == "__main__": unittest.main() \ No newline at end of file