From 705ca3e4c87e7914022ef6108f6af9a05a6d9fac Mon Sep 17 00:00:00 2001 From: phernandez Date: Mon, 30 Dec 2024 23:42:57 -0600 Subject: [PATCH] fix tests --- src/basic_memory/mcp/server.py | 161 ++++++++++++++++++++++-- src/basic_memory/mcp/tools/activity.py | 4 +- src/basic_memory/mcp/tools/discovery.py | 8 +- src/basic_memory/mcp/tools/documents.py | 42 +++---- src/basic_memory/mcp/tools/enhanced.py | 137 -------------------- src/basic_memory/mcp/tools/help.py | 39 +++--- src/basic_memory/mcp/tools/knowledge.py | 90 ++++++------- src/basic_memory/mcp/tools/search.py | 6 +- tests/mcp/test_schema.py | 27 ++-- 9 files changed, 255 insertions(+), 259 deletions(-) delete mode 100644 src/basic_memory/mcp/tools/enhanced.py diff --git a/src/basic_memory/mcp/server.py b/src/basic_memory/mcp/server.py index 9ead88bd..604c363c 100644 --- a/src/basic_memory/mcp/server.py +++ b/src/basic_memory/mcp/server.py @@ -1,12 +1,157 @@ -"""Shared MCP instance for Basic Memory.""" +"""Enhanced FastMCP server instance for Basic Memory.""" + +from typing import Any, Callable, Dict, List, Optional from fastmcp import FastMCP -from basic_memory.mcp.tools.enhanced import EnhancedToolManager +import inspect +from pydantic import BaseModel, Field -# Create and configure the shared MCP instance -mcp = FastMCP("Basic Memory") +from fastmcp.tools import Tool as FastMCPTool +from fastmcp.tools.tool_manager import ToolManager as FastMCPToolManager -# Replace the default tool manager with our enhanced version -mcp._tool_manager = EnhancedToolManager( - warn_on_duplicate_tools=mcp.settings.warn_on_duplicate_tools -) \ No newline at end of file + +class BasicMemoryServer(FastMCP): + """Enhanced FastMCP server with schema support.""" + + def __init__(self, name: str | None = None, **settings: Any): + super().__init__(name=name, **settings) + # Replace default tool manager with our enhanced version + self._tool_manager = EnhancedToolManager( + warn_on_duplicate_tools=self.settings.warn_on_duplicate_tools + ) + + def tool( + self, + name: Optional[str] = None, + description: Optional[str] = None, + examples: Optional[List[Dict]] = None, + category: Optional[str] = None, + ): + """Decorator to register an enhanced tool. + + Example: + @server.tool( + name="search", + description="Search for entities", + category="core", + examples=[{ + "name": "Basic Search", + "description": "Search by text", + "code": 'results = await search({"query": "test"})' + }] + ) + async def search(request: SearchRequest) -> SearchResults: + return await search_service.search(request) + """ + + def decorator(fn: Callable) -> Callable: + tool = self._tool_manager.add_tool( + fn, name=name, description=description, examples=examples, category=category + ) + return fn + + return decorator + + +class ToolExample(BaseModel): + """Example usage of a tool.""" + + name: str = Field(description="Name of the example") + description: str = Field(description="Description of what the example demonstrates") + code: str = Field(description="Example code") + + +class EnhancedTool(FastMCPTool): + """Extended tool registration with rich metadata.""" + + examples: List[ToolExample] = Field(default_factory=list) + category: Optional[str] = Field(None) + input_schema: Optional[Dict] = Field(None) + output_schema: Optional[Dict] = Field(None) + + @classmethod + def from_function( + cls, + fn: Callable, + name: Optional[str] = None, + description: Optional[str] = None, + examples: Optional[List[Dict]] = None, + category: Optional[str] = None, + ) -> "EnhancedTool": + """Create an enhanced tool from a function.""" + # First create the base tool + base_tool = super().from_function(fn, name=name, description=description) + + # Extract return type schema if available + return_schema = None + sig = inspect.signature(fn) + return_type = sig.return_annotation + + if hasattr(return_type, "model_json_schema"): + return_schema = return_type.model_json_schema() + + # Convert examples to ToolExample models + tool_examples = [ToolExample(**ex) for ex in (examples or [])] + + return cls( + fn=fn, + name=base_tool.name, + description=base_tool.description, + parameters=base_tool.parameters, + fn_metadata=base_tool.fn_metadata, + is_async=base_tool.is_async, + context_kwarg=base_tool.context_kwarg, + examples=tool_examples, + category=category, + input_schema=base_tool.parameters, + output_schema=return_schema, + ) + + def get_schema(self) -> Dict: + """Get complete tool schema including examples and metadata.""" + return { + "name": self.name, + "description": self.description, + "category": self.category, + "inputSchema": self.input_schema, + "outputSchema": self.output_schema, + "examples": [ex.model_dump() for ex in self.examples], + } + + +class EnhancedToolManager(FastMCPToolManager): + """Tool manager with enhanced metadata support.""" + + def add_tool( + self, + fn: Callable, + name: Optional[str] = None, + description: Optional[str] = None, + examples: Optional[List[Dict]] = None, + category: Optional[str] = None, + ) -> EnhancedTool: + """Add a tool with enhanced metadata.""" + tool = EnhancedTool.from_function( + fn, name=name, description=description, examples=examples, category=category + ) + self._tools[tool.name] = tool + return tool + + def get_schema_catalog(self) -> Dict: + """Get complete schema catalog for all tools.""" + catalog = {"tools": {}, "categories": {}} + + for tool in self._tools.values(): + if isinstance(tool, EnhancedTool): + catalog["tools"][tool.name] = tool.get_schema() + + if tool.category: + if tool.category not in catalog["categories"]: + catalog["categories"][tool.category] = {"name": tool.category, "tools": []} + catalog["categories"][tool.category]["tools"].append(tool.name) + + return catalog + + +# Create the shared server instance +mcp = BasicMemoryServer("Basic Memory") diff --git a/src/basic_memory/mcp/tools/activity.py b/src/basic_memory/mcp/tools/activity.py index 3b7a54a6..8a3c3d54 100644 --- a/src/basic_memory/mcp/tools/activity.py +++ b/src/basic_memory/mcp/tools/activity.py @@ -5,11 +5,11 @@ from typing import List, Optional from loguru import logger from basic_memory.mcp.async_client import client -from basic_memory.mcp.tools.enhanced import enhanced_tool +from basic_memory.mcp.server import mcp from basic_memory.schemas.activity import ActivityType, RecentActivity -@enhanced_tool() +@mcp.tool() async def get_recent_activity( timeframe: str = "1d", activity_types: Optional[List[ActivityType]] = None, diff --git a/src/basic_memory/mcp/tools/discovery.py b/src/basic_memory/mcp/tools/discovery.py index abbb732c..0192d425 100644 --- a/src/basic_memory/mcp/tools/discovery.py +++ b/src/basic_memory/mcp/tools/discovery.py @@ -4,12 +4,12 @@ from typing import List, Optional from loguru import logger -from basic_memory.mcp.tools.enhanced import enhanced_tool +from basic_memory.mcp.server import mcp from basic_memory.schemas import EntityTypeList, ObservationCategoryList, TypedEntityList from basic_memory.mcp.async_client import client -@enhanced_tool() +@mcp.tool() async def get_entity_types() -> List[str]: """List all unique entity types in use across the knowledge graph. @@ -33,7 +33,7 @@ async def get_entity_types() -> List[str]: return EntityTypeList.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def get_observation_categories() -> List[str]: """List all unique observation categories in use across the knowledge graph. @@ -57,7 +57,7 @@ async def get_observation_categories() -> List[str]: return ObservationCategoryList.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def list_by_type( entity_type: str, include_related: bool = False, sort_by: Optional[str] = "updated_at" ) -> TypedEntityList: diff --git a/src/basic_memory/mcp/tools/documents.py b/src/basic_memory/mcp/tools/documents.py index a3d1e38d..da331f21 100644 --- a/src/basic_memory/mcp/tools/documents.py +++ b/src/basic_memory/mcp/tools/documents.py @@ -2,25 +2,25 @@ from typing import Dict, List -from basic_memory.mcp.tools.enhanced import enhanced_tool +from basic_memory.mcp.server import mcp from basic_memory.schemas.request import DocumentRequest, DocumentPathId from basic_memory.schemas.response import DocumentResponse, DocumentCreateResponse from basic_memory.mcp.async_client import client -@enhanced_tool() +@mcp.tool() async def create_document(request: DocumentRequest) -> DocumentCreateResponse: """Create a new markdown document. - + Examples: # Create a technical specification request = DocumentRequest( path="specs/memory_format.md", content='''# Memory Format Specification - + ## Overview This document defines the standard format for memory files. - + ## Format - Markdown with frontmatter - UTF-8 encoding @@ -33,7 +33,7 @@ async def create_document(request: DocumentRequest) -> DocumentCreateResponse: } ) response = await create_document(request) - + # Response contains document info: # DocumentCreateResponse( # path="specs/memory_format.md", @@ -48,16 +48,16 @@ async def create_document(request: DocumentRequest) -> DocumentCreateResponse: return DocumentCreateResponse.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def update_document(request: DocumentRequest) -> DocumentResponse: """Update an existing document. - + Examples: # Update implementation docs with new details request = DocumentRequest( path="docs/implementation.md", content='''# Implementation Details - + ## Recent Changes - Added FTS5 support - Improved error handling @@ -69,14 +69,14 @@ async def update_document(request: DocumentRequest) -> DocumentResponse: } ) response = await update_document(request) - + # Response contains updated document: # DocumentResponse( # path="docs/implementation.md", # content="# Implementation Details\n...", # checksum="def456...", # doc_metadata={...}, - # created_at="2024-12-20T10:00:00Z", + # created_at="2024-12-20T10:00:00Z", # updated_at="2024-12-25T14:30:00Z" # ) """ @@ -85,14 +85,14 @@ async def update_document(request: DocumentRequest) -> DocumentResponse: return DocumentResponse.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def get_document(path: DocumentPathId) -> DocumentResponse: """Get a document by its path. - + Examples: # Load an API specification response = await get_document("specs/api_format.md") - + # Response contains complete document: # DocumentResponse( # path="specs/api_format.md", @@ -114,14 +114,14 @@ async def get_document(path: DocumentPathId) -> DocumentResponse: return DocumentResponse.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def list_documents() -> List[DocumentCreateResponse]: """List all documents in the system. - + Examples: # Get all documents with metadata documents = await list_documents() - + # Response is list of document info: # [ # DocumentCreateResponse( @@ -145,14 +145,14 @@ async def list_documents() -> List[DocumentCreateResponse]: return [DocumentCreateResponse.model_validate(doc) for doc in response.json()] -@enhanced_tool() +@mcp.tool() async def delete_document(path: DocumentPathId) -> Dict[str, bool]: """Delete a document. - + Examples: # Remove an obsolete document result = await delete_document("docs/outdated_spec.md") - + # Response indicates success: # { # "deleted": true @@ -162,4 +162,4 @@ async def delete_document(path: DocumentPathId) -> Dict[str, bool]: response = await client.delete(url) if response.status_code == 204: return {"deleted": True} - return response.json() \ No newline at end of file + return response.json() diff --git a/src/basic_memory/mcp/tools/enhanced.py b/src/basic_memory/mcp/tools/enhanced.py deleted file mode 100644 index 58e6a9c9..00000000 --- a/src/basic_memory/mcp/tools/enhanced.py +++ /dev/null @@ -1,137 +0,0 @@ -"""Enhanced MCP tool support with rich schema information.""" - -from typing import Any, Callable, Dict, List, Optional, Type -import inspect -from pydantic import BaseModel, Field - -from fastmcp.tools import Tool as FastMCPTool -from fastmcp.tools.tool_manager import ToolManager as FastMCPToolManager - - -class ToolExample(BaseModel): - """Example usage of a tool.""" - name: str = Field(description="Name of the example") - description: str = Field(description="Description of what the example demonstrates") - code: str = Field(description="Example code") - - -class EnhancedTool(FastMCPTool): - """Extended tool registration with rich metadata.""" - examples: List[ToolExample] = Field(default_factory=list) - category: Optional[str] = Field(None) - input_schema: Optional[Dict] = Field(None) - output_schema: Optional[Dict] = Field(None) - - @classmethod - def from_function( - cls, - fn: Callable, - name: Optional[str] = None, - description: Optional[str] = None, - examples: Optional[List[Dict]] = None, - category: Optional[str] = None, - ) -> "EnhancedTool": - """Create an enhanced tool from a function.""" - # First create the base tool - base_tool = super().from_function(fn, name=name, description=description) - - # Extract return type schema if available - return_schema = None - sig = inspect.signature(fn) - return_type = sig.return_annotation - - if hasattr(return_type, "model_json_schema"): - return_schema = return_type.model_json_schema() - - # Convert examples to ToolExample models - tool_examples = [ToolExample(**ex) for ex in (examples or [])] - - return cls( - fn=fn, - name=base_tool.name, - description=base_tool.description, - parameters=base_tool.parameters, - fn_metadata=base_tool.fn_metadata, - is_async=base_tool.is_async, - context_kwarg=base_tool.context_kwarg, - examples=tool_examples, - category=category, - input_schema=base_tool.parameters, - output_schema=return_schema - ) - - def get_schema(self) -> Dict: - """Get complete tool schema including examples and metadata.""" - return { - "name": self.name, - "description": self.description, - "category": self.category, - "inputSchema": self.input_schema, - "outputSchema": self.output_schema, - "examples": [ex.model_dump() for ex in self.examples] - } - - -class EnhancedToolManager(FastMCPToolManager): - """Tool manager with enhanced metadata support.""" - - def add_tool( - self, - fn: Callable, - name: Optional[str] = None, - description: Optional[str] = None, - examples: Optional[List[Dict]] = None, - category: Optional[str] = None, - ) -> EnhancedTool: - """Add a tool with enhanced metadata.""" - tool = EnhancedTool.from_function( - fn, - name=name, - description=description, - examples=examples, - category=category - ) - self._tools[tool.name] = tool - return tool - - def get_schema_catalog(self) -> Dict: - """Get complete schema catalog for all tools.""" - catalog = { - "tools": {}, - "categories": {} - } - - for tool in self._tools.values(): - if isinstance(tool, EnhancedTool): - catalog["tools"][tool.name] = tool.get_schema() - - if tool.category: - if tool.category not in catalog["categories"]: - catalog["categories"][tool.category] = { - "name": tool.category, - "tools": [] - } - catalog["categories"][tool.category]["tools"].append(tool.name) - - return catalog - - -def enhanced_tool( - name: Optional[str] = None, - description: Optional[str] = None, - examples: Optional[List[Dict]] = None, - category: Optional[str] = None, -): - """Decorator for registering enhanced tools.""" - def decorator(fn: Callable) -> Callable: - # Store metadata on the function for later - if not hasattr(fn, "_tool_metadata"): - fn._tool_metadata = {} - fn._tool_metadata.update({ - "name": name, - "description": description, - "examples": examples, - "category": category - }) - return fn - return decorator \ No newline at end of file diff --git a/src/basic_memory/mcp/tools/help.py b/src/basic_memory/mcp/tools/help.py index 8b9e3117..5faab999 100644 --- a/src/basic_memory/mcp/tools/help.py +++ b/src/basic_memory/mcp/tools/help.py @@ -3,40 +3,37 @@ from typing import Dict, Optional from basic_memory.mcp.server import mcp -from basic_memory.mcp.tools.enhanced import enhanced_tool -@enhanced_tool( +@mcp.tool( category="system", examples=[ { "name": "Get All Tools", "description": "Get complete schema catalog for all tools", - "code": "catalog = await get_schema()" + "code": "catalog = await get_schema()", }, { "name": "Get Specific Tool", "description": "Get schema for a specific tool", - "code": 'tool_schema = await get_schema("create_entity")' - } - ] + "code": 'tool_schema = await get_schema("create_entity")', + }, + ], ) async def get_schema( - tool_name: Optional[str] = None, - include_examples: bool = True, - include_referenced: bool = True + tool_name: Optional[str] = None, include_examples: bool = True, include_referenced: bool = True ) -> Dict: """Get schema information about available tools. - + Args: tool_name: Optional name of specific tool to get schema for include_examples: Whether to include usage examples include_referenced: Whether to include referenced model schemas - + Returns: Complete tool catalog if tool_name is None, or specific tool schema if tool_name is provided. - + Tool catalog includes: - Tool descriptions and purposes - Input/output schemas @@ -46,30 +43,30 @@ async def get_schema( """ # Our tool manager has the enhanced schema support catalog = mcp._tool_manager.get_schema_catalog() - + # Filter if specific tool requested if tool_name: if tool_name not in catalog["tools"]: raise ValueError(f"Unknown tool: {tool_name}") - + tool_schema = catalog["tools"][tool_name] - + if not include_examples: tool_schema.pop("examples", None) - + if include_referenced: return { "tools": {tool_name: tool_schema}, - "referencedModels": tool_schema.get("referencedModels", {}) + "referencedModels": tool_schema.get("referencedModels", {}), } else: return {"tools": {tool_name: tool_schema}} - + # Return full catalog with requested inclusions result = catalog - + if not include_examples: for tool in result["tools"].values(): tool.pop("examples", None) - - return result \ No newline at end of file + + return result diff --git a/src/basic_memory/mcp/tools/knowledge.py b/src/basic_memory/mcp/tools/knowledge.py index e0b567c0..894dcb7c 100644 --- a/src/basic_memory/mcp/tools/knowledge.py +++ b/src/basic_memory/mcp/tools/knowledge.py @@ -4,7 +4,8 @@ from typing import Dict import httpx -from basic_memory.schemas.base import Entity, Relation, ObservationCategory, PathId +from basic_memory.mcp.server import mcp +from basic_memory.schemas.base import PathId from basic_memory.schemas.request import ( CreateEntityRequest, CreateRelationsRequest, @@ -13,20 +14,20 @@ from basic_memory.schemas.request import ( from basic_memory.schemas.delete import ( DeleteEntitiesRequest, DeleteObservationsRequest, - DeleteRelationsRequest + DeleteRelationsRequest, ) from basic_memory.schemas.response import EntityListResponse, EntityResponse from basic_memory.mcp.async_client import client from basic_memory.services.exceptions import EntityNotFoundError -from basic_memory.mcp.tools.enhanced import enhanced_tool -@enhanced_tool( +@mcp.tool( category="knowledge", - examples=[{ - "name": "Create Component", - "description": "Create a new technical component", - "code": """ + examples=[ + { + "name": "Create Component", + "description": "Create a new technical component", + "code": """ await create_entities({ "entities": [{ "name": "SearchService", @@ -38,12 +39,13 @@ await create_entities({ ] }] }) -""" - }] +""", + } + ], ) async def create_entities(request: CreateEntityRequest) -> EntityListResponse: """Create new entities in the knowledge graph. - + Entities can include initial observations and properties. Entity IDs are automatically generated from the type and name. """ @@ -52,12 +54,13 @@ async def create_entities(request: CreateEntityRequest) -> EntityListResponse: return EntityListResponse.model_validate(response.json()) -@enhanced_tool( +@mcp.tool( category="knowledge", - examples=[{ - "name": "Add Dependency", - "description": "Create dependency relationship between components", - "code": """ + examples=[ + { + "name": "Add Dependency", + "description": "Create dependency relationship between components", + "code": """ await create_relations({ "relations": [{ "from_id": "component/search_service", @@ -66,8 +69,9 @@ await create_relations({ "context": "Needs storage for search indexes" }] }) -""" - }] +""", + } + ], ) async def create_relations(request: CreateRelationsRequest) -> EntityListResponse: """Create relations between existing entities.""" @@ -76,28 +80,30 @@ async def create_relations(request: CreateRelationsRequest) -> EntityListRespons return EntityListResponse.model_validate(response.json()) -@enhanced_tool( +@mcp.tool( category="knowledge", - examples=[{ - "name": "Get Entity Details", - "description": "Load complete entity information", - "code": """ + examples=[ + { + "name": "Get Entity Details", + "description": "Load complete entity information", + "code": """ # Get component details entity = await get_entity("component/search_service") print(f"Name: {entity.name}") print(f"Type: {entity.entity_type}") for obs in entity.observations: print(f"- {obs.content}") -""" - }] +""", + } + ], ) async def get_entity(path_id: PathId) -> EntityResponse: """Get a specific entity by its path_id. - + Examples: # Load implementation details response = await get_entity("component/memory_service") - + # Response contains complete entity: # EntityResponse( # path_id="component/memory_service", @@ -124,7 +130,7 @@ async def get_entity(path_id: PathId) -> EntityResponse: # Load and analyze a design spec spec = await get_entity("specification/file_format") - decisions = [obs for obs in spec.observations + decisions = [obs for obs in spec.observations if obs.category == ObservationCategory.DESIGN] """ try: @@ -142,12 +148,10 @@ async def get_entity(path_id: PathId) -> EntityResponse: raise - - -@enhanced_tool() +@mcp.tool() async def add_observations(request: AddObservationsRequest) -> EntityResponse: """Add observations to an existing entity. - + Examples: # Document implementation decisions with context request = AddObservationsRequest( @@ -169,7 +173,7 @@ async def add_observations(request: AddObservationsRequest) -> EntityResponse: ] ) response = await add_observations(request) - + # Response shows entity with new observations: # EntityResponse( # path_id="component/search_service", @@ -188,10 +192,10 @@ async def add_observations(request: AddObservationsRequest) -> EntityResponse: return EntityResponse.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def delete_observations(request: DeleteObservationsRequest) -> EntityResponse: """Delete specific observations from an entity. - + Examples: # Remove obsolete implementation notes request = DeleteObservationsRequest( @@ -202,7 +206,7 @@ async def delete_observations(request: DeleteObservationsRequest) -> EntityRespo ] ) response = await delete_observations(request) - + # Response shows entity with observations removed: # EntityResponse( # path_id="component/indexer", @@ -214,10 +218,10 @@ async def delete_observations(request: DeleteObservationsRequest) -> EntityRespo return EntityResponse.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def delete_relations(request: DeleteRelationsRequest) -> EntityListResponse: """Delete relations between entities. - + Examples: # Remove obsolete dependency request = DeleteRelationsRequest( @@ -230,7 +234,7 @@ async def delete_relations(request: DeleteRelationsRequest) -> EntityListRespons ] ) response = await delete_relations(request) - + # Response shows updated entities: # EntityListResponse( # entities=[ @@ -248,10 +252,10 @@ async def delete_relations(request: DeleteRelationsRequest) -> EntityListRespons return EntityListResponse.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def delete_entities(request: DeleteEntitiesRequest) -> Dict[str, bool]: """Delete entities from the knowledge graph. - + Examples: # Remove obsolete components request = DeleteEntitiesRequest( @@ -261,7 +265,7 @@ async def delete_entities(request: DeleteEntitiesRequest) -> Dict[str, bool]: ] ) response = await delete_entities(request) - + # Response indicates success: # { # "deleted": true @@ -269,4 +273,4 @@ async def delete_entities(request: DeleteEntitiesRequest) -> Dict[str, bool]: """ url = "/knowledge/entities/delete" response = await client.post(url, json=request.model_dump()) - return response.json() \ No newline at end of file + return response.json() diff --git a/src/basic_memory/mcp/tools/search.py b/src/basic_memory/mcp/tools/search.py index 194dc37f..344282a6 100644 --- a/src/basic_memory/mcp/tools/search.py +++ b/src/basic_memory/mcp/tools/search.py @@ -2,13 +2,13 @@ from typing import Dict -from basic_memory.mcp.tools.enhanced import enhanced_tool +from basic_memory.mcp.server import mcp from basic_memory.schemas.request import SearchNodesRequest, OpenNodesRequest from basic_memory.schemas.response import SearchNodesResponse, EntityResponse from basic_memory.mcp.async_client import client -@enhanced_tool() +@mcp.tool() async def search_nodes(request: SearchNodesRequest) -> SearchNodesResponse: """Search for entities in the knowledge graph. @@ -51,7 +51,7 @@ async def search_nodes(request: SearchNodesRequest) -> SearchNodesResponse: return SearchNodesResponse.model_validate(response.json()) -@enhanced_tool() +@mcp.tool() async def open_nodes(request: OpenNodesRequest) -> Dict[str, EntityResponse]: """Load multiple entities by their path_ids. diff --git a/tests/mcp/test_schema.py b/tests/mcp/test_schema.py index e50445db..d1781a7c 100644 --- a/tests/mcp/test_schema.py +++ b/tests/mcp/test_schema.py @@ -2,7 +2,10 @@ import pytest from pydantic import BaseModel -from typing import List +from typing import List, Optional + +from basic_memory.mcp.server import mcp +from basic_memory.mcp.tools.help import get_schema class TestInput(BaseModel): @@ -20,12 +23,8 @@ class TestOutput(BaseModel): @pytest.mark.asyncio async def test_get_schema_all(): """Test getting complete tool catalog.""" - from basic_memory.mcp.server import mcp - from basic_memory.mcp.tools.enhanced import enhanced_tool - from basic_memory.mcp.tools.help import get_schema - # Create test tools - @enhanced_tool( + @mcp.tool( name="test_tool", description="A test tool with enhanced metadata", category="test", @@ -63,11 +62,7 @@ async def test_get_schema_all(): @pytest.mark.asyncio async def test_get_schema_enhanced_tool(): """Test getting schema for enhanced tool.""" - from basic_memory.mcp.server import mcp - from basic_memory.mcp.tools.enhanced import enhanced_tool - from basic_memory.mcp.tools.help import get_schema - - @enhanced_tool( + @mcp.tool( name="test_tool", description="A test tool with enhanced metadata", category="test", @@ -104,9 +99,6 @@ async def test_get_schema_enhanced_tool(): @pytest.mark.asyncio async def test_get_schema_basic_tool(): """Test getting schema for basic tool.""" - from basic_memory.mcp.server import mcp - from basic_memory.mcp.tools.help import get_schema - @mcp.tool(name="basic_tool") async def basic_tool(value: str) -> str: """A basic tool without enhanced metadata""" @@ -126,10 +118,7 @@ async def test_get_schema_basic_tool(): @pytest.mark.asyncio async def test_get_schema_filter_examples(): """Test filtering out examples from schema.""" - from basic_memory.mcp.tools.enhanced import enhanced_tool - from basic_memory.mcp.tools.help import get_schema - - @enhanced_tool( + @mcp.tool( name="test_tool", description="A test tool with enhanced metadata", category="test", @@ -153,7 +142,5 @@ async def test_get_schema_filter_examples(): @pytest.mark.asyncio async def test_get_schema_unknown_tool(): """Test getting schema for unknown tool.""" - from basic_memory.mcp.tools.help import get_schema - with pytest.raises(ValueError, match="Unknown tool: unknown_tool"): await get_schema("unknown_tool") \ No newline at end of file