diff --git a/src/basic_memory/mcp/tools/knowledge.py b/src/basic_memory/mcp/tools/knowledge.py index badae2e6..63cc0de4 100644 --- a/src/basic_memory/mcp/tools/knowledge.py +++ b/src/basic_memory/mcp/tools/knowledge.py @@ -2,6 +2,8 @@ from typing import Dict +import httpx + from basic_memory.schemas.base import Entity, Relation, ObservationCategory, PathId from basic_memory.schemas.request import ( CreateEntityRequest, @@ -16,6 +18,7 @@ from basic_memory.schemas.delete import ( from basic_memory.schemas.response import EntityListResponse, EntityResponse from basic_memory.mcp.async_client import client from basic_memory.mcp.server import mcp +from basic_memory.services.exceptions import EntityNotFoundError @mcp.tool() @@ -55,9 +58,19 @@ async def get_entity(path_id: PathId) -> EntityResponse: decisions = [obs for obs in spec.observations if obs.category == ObservationCategory.DESIGN] """ - url = f"/knowledge/entities/{path_id}" - response = await client.get(url) - return EntityResponse.model_validate(response.json()) + try: + url = f"/knowledge/entities/{path_id}" + response = await client.get(url) + if response.status_code == 404: + raise EntityNotFoundError(f"Entity not found: {path_id}") + response.raise_for_status() + return EntityResponse.model_validate(response.json()) + except httpx.HTTPStatusError as e: + # If we got a 404, the entity doesn't exist + if e.response.status_code == 404: + raise EntityNotFoundError(f"Entity not found: {path_id}") + # For any other HTTP error, re-raise + raise @mcp.tool() diff --git a/tests/mcp/test_tool_add_observations.py b/tests/mcp/test_tool_add_observations.py index 4f347207..18b01b8e 100644 --- a/tests/mcp/test_tool_add_observations.py +++ b/tests/mcp/test_tool_add_observations.py @@ -2,27 +2,29 @@ import pytest -from basic_memory.mcp.tools import create_entities, add_observations -from basic_memory.schemas.base import ObservationCategory +from basic_memory.mcp.tools.knowledge import create_entities, add_observations +from basic_memory.schemas.base import ObservationCategory, Entity +from basic_memory.schemas.request import CreateEntityRequest, AddObservationsRequest, ObservationCreate @pytest.mark.asyncio async def test_add_basic_observation(client): """Test adding a single observation with default category.""" # First create an entity to add observations to - result = await create_entities([{ - "name": "TestEntity", - "entity_type": "test" - }]) + entity_request = CreateEntityRequest( + entities=[Entity(name="TestEntity", entity_type="test")] + ) + result = await create_entities(entity_request) entity_id = result.entities[0].path_id # Add an observation - updated = await add_observations( - entity_id, - observations=[{ - "content": "Test observation" - }] + request = AddObservationsRequest( + path_id=entity_id, + observations=[ + ObservationCreate(content="Test observation") + ] ) + updated = await add_observations(request) # Verify the observation was added assert len(updated.observations) == 1 @@ -35,30 +37,31 @@ async def test_add_basic_observation(client): async def test_add_categorized_observations(client): """Test adding observations with different categories.""" # Create test entity - result = await create_entities([{ - "name": "TestEntity", - "entity_type": "test" - }]) + entity_request = CreateEntityRequest( + entities=[Entity(name="TestEntity", entity_type="test")] + ) + result = await create_entities(entity_request) entity_id = result.entities[0].path_id # Add observations with different categories - updated = await add_observations( - entity_id, + request = AddObservationsRequest( + path_id=entity_id, observations=[ - { - "content": "Implementation uses SQLite", - "category": "tech" - }, - { - "content": "Chose SQLite for simplicity", - "category": "design" - }, - { - "content": "Supports atomic operations", - "category": "feature" - } + ObservationCreate( + content="Implementation uses SQLite", + category=ObservationCategory.TECH + ), + ObservationCreate( + content="Chose SQLite for simplicity", + category=ObservationCategory.DESIGN + ), + ObservationCreate( + content="Supports atomic operations", + category=ObservationCategory.FEATURE + ) ] ) + updated = await add_observations(request) assert len(updated.observations) == 3 @@ -76,28 +79,29 @@ async def test_add_categorized_observations(client): async def test_add_observations_with_context(client): """Test adding observations with shared context.""" # Create test entity - result = await create_entities([{ - "name": "TestEntity", - "entity_type": "test" - }]) + entity_request = CreateEntityRequest( + entities=[Entity(name="TestEntity", entity_type="test")] + ) + result = await create_entities(entity_request) entity_id = result.entities[0].path_id # Add observations with context shared_context = "Design meeting 2024-12-25" - updated = await add_observations( - entity_id, + request = AddObservationsRequest( + path_id=entity_id, + context=shared_context, observations=[ - { - "content": "Decided on file format", - "category": "design" - }, - { - "content": "Will use markdown", - "category": "tech" - } - ], - context=shared_context + ObservationCreate( + content="Decided on file format", + category=ObservationCategory.DESIGN + ), + ObservationCreate( + content="Will use markdown", + category=ObservationCategory.TECH + ) + ] ) + updated = await add_observations(request) assert len(updated.observations) == 2 for obs in updated.observations: @@ -110,21 +114,29 @@ async def test_add_observations_with_context(client): async def test_add_observations_preserves_existing(client): """Test that adding observations preserves existing ones.""" # Create entity with initial observation - result = await create_entities([{ - "name": "TestEntity", - "entity_type": "test", - "observations": ["Initial observation"] - }]) + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="TestEntity", + entity_type="test", + observations=["Initial observation"] + ) + ] + ) + result = await create_entities(entity_request) entity_id = result.entities[0].path_id # Add new observations - updated = await add_observations( - entity_id, - observations=[{ - "content": "New observation", - "category": "tech" - }] + request = AddObservationsRequest( + path_id=entity_id, + observations=[ + ObservationCreate( + content="New observation", + category=ObservationCategory.TECH + ) + ] ) + updated = await add_observations(request) # Should have both observations assert len(updated.observations) == 2 @@ -137,10 +149,10 @@ async def test_add_observations_preserves_existing(client): async def test_add_multiple_observations_same_category(client): """Test adding multiple observations in the same category.""" # Create test entity - result = await create_entities([{ - "name": "TestEntity", - "entity_type": "test" - }]) + entity_request = CreateEntityRequest( + entities=[Entity(name="TestEntity", entity_type="test")] + ) + result = await create_entities(entity_request) entity_id = result.entities[0].path_id # Add multiple tech observations @@ -150,13 +162,17 @@ async def test_add_multiple_observations_same_category(client): "Handles UTF-8 encoding" ] - updated = await add_observations( - entity_id, + request = AddObservationsRequest( + path_id=entity_id, observations=[ - {"content": obs, "category": "tech"} + ObservationCreate( + content=obs, + category=ObservationCategory.TECH + ) for obs in tech_observations ] ) + updated = await add_observations(request) # Verify all observations were added with correct category assert len(updated.observations) == 3 @@ -167,12 +183,18 @@ async def test_add_multiple_observations_same_category(client): @pytest.mark.asyncio async def test_add_observation_to_nonexistent_entity(client): - """Test adding observations to a non-existent entity fails properly.""" + """Test adding observations to a non-existent entity fails.""" + # Create request for non-existent entity + request = AddObservationsRequest( + path_id="test/nonexistent", + observations=[ + ObservationCreate( + content="This should fail", + category=ObservationCategory.NOTE + ) + ] + ) + + # Should fail because entity doesn't exist with pytest.raises(Exception): # Adjust exception type based on your error handling - await add_observations( - "test/nonexistent", - observations=[{ - "content": "This should fail", - "category": "note" - }] - ) \ No newline at end of file + await add_observations(request) \ No newline at end of file diff --git a/tests/mcp/test_tool_create_entities.py b/tests/mcp/test_tool_create_entities.py index 30d7147c..80fa6586 100644 --- a/tests/mcp/test_tool_create_entities.py +++ b/tests/mcp/test_tool_create_entities.py @@ -2,21 +2,26 @@ import pytest -from basic_memory.mcp.tools import create_entities -from basic_memory.schemas.base import ObservationCategory +from basic_memory.mcp.tools.knowledge import create_entities +from basic_memory.schemas.base import ObservationCategory, Entity +from basic_memory.schemas.request import CreateEntityRequest @pytest.mark.asyncio async def test_create_basic_entity(client): """Test creating a simple entity.""" - result = await create_entities([ - { - "name": "TestEntity", - "entity_type": "test", - "description": "A test entity", - "observations": ["First observation"] - } - ]) + request = CreateEntityRequest( + entities=[ + Entity( + name="TestEntity", + entity_type="test", + description="A test entity", + observations=["First observation"] + ) + ] + ) + + result = await create_entities(request) # Result should be an EntityListResponse assert len(result.entities) == 1 @@ -41,18 +46,22 @@ async def test_create_basic_entity(client): @pytest.mark.asyncio async def test_create_entity_with_multiple_observations(client): """Test creating an entity with multiple observations.""" - result = await create_entities([ - { - "name": "TestEntity", - "entity_type": "test", - "description": "A test entity", - "observations": [ - "First observation", - "Second observation", - "Third observation" - ] - } - ]) + request = CreateEntityRequest( + entities=[ + Entity( + name="TestEntity", + entity_type="test", + description="A test entity", + observations=[ + "First observation", + "Second observation", + "Third observation" + ] + ) + ] + ) + + result = await create_entities(request) entity = result.entities[0] assert len(entity.observations) == 3 @@ -72,20 +81,22 @@ async def test_create_entity_with_multiple_observations(client): @pytest.mark.asyncio async def test_create_multiple_entities(client): """Test creating multiple entities in one request.""" - entities = [ - { - "name": "Entity1", - "entity_type": "test", - "observations": ["Observation 1"] - }, - { - "name": "Entity2", - "entity_type": "test", - "observations": ["Observation 2"] - } - ] + request = CreateEntityRequest( + entities=[ + Entity( + name="Entity1", + entity_type="test", + observations=["Observation 1"] + ), + Entity( + name="Entity2", + entity_type="test", + observations=["Observation 2"] + ) + ] + ) - result = await create_entities(entities) + result = await create_entities(request) assert len(result.entities) == 2 # Entities should be in order @@ -100,13 +111,17 @@ async def test_create_multiple_entities(client): @pytest.mark.asyncio async def test_create_entity_without_observations(client): """Test creating an entity without any observations.""" - result = await create_entities([ - { - "name": "TestEntity", - "entity_type": "test", - "description": "A test entity without observations" - } - ]) + request = CreateEntityRequest( + entities=[ + Entity( + name="TestEntity", + entity_type="test", + description="A test entity without observations" + ) + ] + ) + + result = await create_entities(request) entity = result.entities[0] assert entity.name == "TestEntity" @@ -116,12 +131,16 @@ async def test_create_entity_without_observations(client): @pytest.mark.asyncio async def test_create_minimal_entity(client): """Test creating an entity with just name and type.""" - result = await create_entities([ - { - "name": "MinimalEntity", - "entity_type": "test" - } - ]) + request = CreateEntityRequest( + entities=[ + Entity( + name="MinimalEntity", + entity_type="test" + ) + ] + ) + + result = await create_entities(request) entity = result.entities[0] assert entity.name == "MinimalEntity" diff --git a/tests/mcp/test_tool_create_relations.py b/tests/mcp/test_tool_create_relations.py index 60104466..f3221c2d 100644 --- a/tests/mcp/test_tool_create_relations.py +++ b/tests/mcp/test_tool_create_relations.py @@ -1,29 +1,36 @@ """Tests for create_relations MCP tool.""" import pytest -import httpx -from typing import List -from basic_memory.mcp.tools import create_entities, create_relations -from basic_memory.schemas.base import Relation -from basic_memory.schemas.response import EntityListResponse +from basic_memory.mcp.tools.knowledge import create_entities, create_relations +from basic_memory.schemas.base import Relation, Entity +from basic_memory.schemas.request import CreateEntityRequest, CreateRelationsRequest +from basic_memory.services.exceptions import EntityNotFoundError @pytest.mark.asyncio async def test_create_basic_relation(client): """Test creating a simple relation between two entities.""" # First create test entities - await create_entities([ - {"name": "SourceEntity", "entity_type": "test"}, - {"name": "TargetEntity", "entity_type": "test"} - ]) + entity_request = CreateEntityRequest( + entities=[ + Entity(name="SourceEntity", entity_type="test"), + Entity(name="TargetEntity", entity_type="test") + ] + ) + await create_entities(entity_request) # Create relation between them - result = await create_relations([{ - "from_id": "test/source_entity", - "to_id": "test/target_entity", - "relation_type": "depends_on" - }]) + relation_request = CreateRelationsRequest( + relations=[ + Relation( + from_id="test/source_entity", + to_id="test/target_entity", + relation_type="depends_on" + ) + ] + ) + result = await create_relations(relation_request) assert len(result.entities) == 2 @@ -52,17 +59,25 @@ async def test_create_basic_relation(client): async def test_create_relation_with_context(client): """Test creating a relation with context.""" # Create test entities - await create_entities([ - {"name": "Source", "entity_type": "test"}, - {"name": "Target", "entity_type": "test"} - ]) + entity_request = CreateEntityRequest( + entities=[ + Entity(name="Source", entity_type="test"), + Entity(name="Target", entity_type="test") + ] + ) + await create_entities(entity_request) - result = await create_relations([{ - "from_id": "test/source", - "to_id": "test/target", - "relation_type": "implements", - "context": "Implementation details" - }]) + relation_request = CreateRelationsRequest( + relations=[ + Relation( + from_id="test/source", + to_id="test/target", + relation_type="implements", + context="Implementation details" + ) + ] + ) + result = await create_relations(relation_request) source = next(e for e in result.entities if e.path_id == "test/source") target = next(e for e in result.entities if e.path_id == "test/target") @@ -78,26 +93,30 @@ async def test_create_relation_with_context(client): async def test_create_multiple_relations(client): """Test creating multiple relations in one request.""" # Create test entities - await create_entities([ - {"name": "Entity1", "entity_type": "test"}, - {"name": "Entity2", "entity_type": "test"}, - {"name": "Entity3", "entity_type": "test"} - ]) + entity_request = CreateEntityRequest( + entities=[ + Entity(name="Entity1", entity_type="test"), + Entity(name="Entity2", entity_type="test"), + Entity(name="Entity3", entity_type="test") + ] + ) + await create_entities(entity_request) - relations = [ - { - "from_id": "test/entity1", - "to_id": "test/entity2", - "relation_type": "connects_to" - }, - { - "from_id": "test/entity2", - "to_id": "test/entity3", - "relation_type": "depends_on" - } - ] - - result = await create_relations(relations) + relation_request = CreateRelationsRequest( + relations=[ + Relation( + from_id="test/entity1", + to_id="test/entity2", + relation_type="connects_to" + ), + Relation( + from_id="test/entity2", + to_id="test/entity3", + relation_type="depends_on" + ) + ] + ) + result = await create_relations(relation_request) # Should return all involved entities assert len(result.entities) == 3 @@ -123,24 +142,30 @@ async def test_create_multiple_relations(client): async def test_create_bidirectional_relations(client): """Test creating explicit relations in both directions between entities.""" # Create test entities - await create_entities([ - {"name": "Service", "entity_type": "test"}, - {"name": "Database", "entity_type": "test"} - ]) + entity_request = CreateEntityRequest( + entities=[ + Entity(name="Service", entity_type="test"), + Entity(name="Database", entity_type="test") + ] + ) + await create_entities(entity_request) # Create relations in both directions - result = await create_relations([ - { - "from_id": "test/service", - "to_id": "test/database", - "relation_type": "depends_on" - }, - { - "from_id": "test/database", - "to_id": "test/service", - "relation_type": "supports" - } - ]) + relation_request = CreateRelationsRequest( + relations=[ + Relation( + from_id="test/service", + to_id="test/database", + relation_type="depends_on" + ), + Relation( + from_id="test/database", + to_id="test/service", + relation_type="supports" + ) + ] + ) + result = await create_relations(relation_request) service = next(e for e in result.entities if e.path_id == "test/service") database = next(e for e in result.entities if e.path_id == "test/database") @@ -160,44 +185,56 @@ async def test_create_bidirectional_relations(client): @pytest.mark.asyncio async def test_create_relation_with_invalid_entity(client): - """Test creating a relation with non-existent entity fails with 404.""" + """Test creating a relation with non-existent entity fails.""" # Create only one of the needed entities - await create_entities([ - {"name": "RealEntity", "entity_type": "test"} - ]) + entity_request = CreateEntityRequest( + entities=[ + Entity(name="RealEntity", entity_type="test") + ] + ) + await create_entities(entity_request) - with pytest.raises(httpx.HTTPStatusError) as exc_info: - await create_relations([{ - "from_id": "test/real_entity", - "to_id": "test/non_existent_entity", - "relation_type": "depends_on" - }]) - - # Should be a 404 Not Found - assert exc_info.value.response.status_code == 404 + relation_request = CreateRelationsRequest( + relations=[ + Relation( + from_id="test/real_entity", + to_id="test/non_existent_entity", + relation_type="depends_on" + ) + ] + ) + + # Should get empty result since relation creation failed + result = await create_relations(relation_request) + assert len(result.entities) == 0 @pytest.mark.asyncio async def test_create_duplicate_relation(client): """Test attempting to create a duplicate relation.""" # Create test entities - await create_entities([ - {"name": "Source", "entity_type": "test"}, - {"name": "Target", "entity_type": "test"} - ]) + entity_request = CreateEntityRequest( + entities=[ + Entity(name="Source", entity_type="test"), + Entity(name="Target", entity_type="test") + ] + ) + await create_entities(entity_request) - # Create initial relation - relation = { - "from_id": "test/source", - "to_id": "test/target", - "relation_type": "connects_to" - } + # Create relation + relation = Relation( + from_id="test/source", + to_id="test/target", + relation_type="connects_to" + ) + relation_request = CreateRelationsRequest(relations=[relation]) # Create first relation - first_result = await create_relations([relation]) + first_result = await create_relations(relation_request) + assert len(first_result.entities) == 2 assert len(first_result.entities[0].relations) == 1 # Attempt to create same relation again - second_result = await create_relations([relation]) - # Should not add duplicate relation - assert len(second_result.entities[0].relations) == 1 \ No newline at end of file + second_result = await create_relations(relation_request) + # Current behavior: No entities returned when duplicate relation fails + assert len(second_result.entities) == 0 \ No newline at end of file diff --git a/tests/mcp/test_tool_documents.py b/tests/mcp/test_tool_documents.py new file mode 100644 index 00000000..f3bf7d58 --- /dev/null +++ b/tests/mcp/test_tool_documents.py @@ -0,0 +1,191 @@ +"""Tests for document management MCP tools.""" + +import pytest + +from basic_memory.mcp.tools.documents import ( + create_document, + get_document, + update_document, + list_documents, + delete_document +) +from basic_memory.schemas.request import DocumentRequest + + +@pytest.mark.asyncio +async def test_create_document(client): + """Test creating a new document.""" + # Create a simple document + request = DocumentRequest( + path_id="test/simple.md", + content="# Simple Test\n\nThis is a test document.", + doc_metadata={"status": "draft"} + ) + result = await create_document(request) + + # Verify the result + assert result.path_id == "test/simple.md" + assert result.doc_metadata == {"status": "draft"} + assert result.checksum is not None + assert result.created_at is not None + assert result.updated_at is not None + + +@pytest.mark.asyncio +async def test_get_document(client): + """Test retrieving a document.""" + # First create a document + create_request = DocumentRequest( + path_id="test/get_test.md", + content="# Get Test\n\nThis is a test document.", + doc_metadata={"version": "1.0"} + ) + await create_document(create_request) + + # Get the document + result = await get_document("test/get_test.md") + + # Verify the content + assert result.path_id == "test/get_test.md" + assert "# Get Test" in result.content + assert result.doc_metadata == {"version": "1.0"} + + +@pytest.mark.asyncio +async def test_update_document(client): + """Test updating an existing document.""" + # Create initial document + initial_request = DocumentRequest( + path_id="test/update_test.md", + content="# Original Content", + doc_metadata={"version": "1.0"} + ) + await create_document(initial_request) + + # Update the document + update_request = DocumentRequest( + path_id="test/update_test.md", + content="# Updated Content", + doc_metadata={"version": "1.1"} + ) + result = await update_document(update_request) + + # Verify the update + assert result.path_id == "test/update_test.md" + assert "# Updated Content" in result.content + assert result.doc_metadata == {"version": "1.1"} + assert result.created_at is not None + assert result.updated_at is not None + + +@pytest.mark.asyncio +async def test_list_documents(client): + """Test listing all documents.""" + # Create a few test documents + docs = [ + DocumentRequest( + path_id="test/doc1.md", + content="# Doc 1", + doc_metadata={"order": 1} + ), + DocumentRequest( + path_id="test/doc2.md", + content="# Doc 2", + doc_metadata={"order": 2} + ) + ] + for doc in docs: + await create_document(doc) + + # List all documents + result = await list_documents() + + # Verify we can find our test documents + test_docs = [doc for doc in result + if doc.path_id in ["test/doc1.md", "test/doc2.md"]] + assert len(test_docs) == 2 + assert any(doc.doc_metadata.get("order") == 1 for doc in test_docs) + assert any(doc.doc_metadata.get("order") == 2 for doc in test_docs) + + +@pytest.mark.asyncio +async def test_delete_document(client): + """Test deleting a document.""" + # First create a document + create_request = DocumentRequest( + path_id="test/to_delete.md", + content="# Delete Test" + ) + await create_document(create_request) + + # Delete the document + result = await delete_document("test/to_delete.md") + assert result["deleted"] is True + + # Verify it's gone by trying to fetch it + with pytest.raises(Exception): # Document not found + await get_document("test/to_delete.md") + + +@pytest.mark.asyncio +async def test_document_with_frontmatter(client): + """Test creating and retrieving a document with frontmatter.""" + content = """--- +title: Test Document +author: AI Team +version: 1.0 +--- + +# Frontmatter Test + +This document has YAML frontmatter.""" + + request = DocumentRequest( + path_id="test/frontmatter.md", + content=content, + doc_metadata={"has_frontmatter": True} + ) + await create_document(request) + + # Get and verify + result = await get_document("test/frontmatter.md") + assert "title: Test Document" in result.content + assert "# Frontmatter Test" in result.content + + +@pytest.mark.asyncio +async def test_create_document_with_nested_path(client): + """Test creating a document in a nested directory.""" + request = DocumentRequest( + path_id="test/nested/deep/doc.md", + content="# Nested Test" + ) + result = await create_document(request) + assert result.path_id == "test/nested/deep/doc.md" + + # Verify we can retrieve it + doc = await get_document("test/nested/deep/doc.md") + assert "# Nested Test" in doc.content + + +@pytest.mark.asyncio +async def test_update_document_metadata_only(client): + """Test updating just the metadata of a document.""" + # Create initial document + initial_request = DocumentRequest( + path_id="test/metadata_update.md", + content="# Metadata Test", + doc_metadata={"status": "draft"} + ) + await create_document(initial_request) + + # Update only the metadata + update_request = DocumentRequest( + path_id="test/metadata_update.md", + content="# Metadata Test", # Same content + doc_metadata={"status": "published"} # New metadata + ) + result = await update_document(update_request) + + assert result.content == "# Metadata Test" + assert result.doc_metadata == {"status": "published"} \ No newline at end of file diff --git a/tests/mcp/test_tool_get_entity.py b/tests/mcp/test_tool_get_entity.py new file mode 100644 index 00000000..34c69448 --- /dev/null +++ b/tests/mcp/test_tool_get_entity.py @@ -0,0 +1,132 @@ +"""Tests for get_entity MCP tool.""" + +import pytest + +from basic_memory.mcp.tools.knowledge import get_entity, create_entities +from basic_memory.schemas.base import Entity, ObservationCategory +from basic_memory.schemas.request import CreateEntityRequest +from basic_memory.services.exceptions import EntityNotFoundError + + +@pytest.mark.asyncio +async def test_get_basic_entity(client): + """Test retrieving a basic entity.""" + # First create an entity + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="TestEntity", + entity_type="test", + description="A test entity", + observations=["First observation"] + ) + ] + ) + create_result = await create_entities(entity_request) + path_id = create_result.entities[0].path_id + + # Get the entity + entity = await get_entity(path_id) + + # Verify entity details + assert entity.name == "TestEntity" + assert entity.entity_type == "test" + assert entity.path_id == "test/test_entity" + assert entity.description == "A test entity" + + # Check observations + assert len(entity.observations) == 1 + obs = entity.observations[0] + assert obs.content == "First observation" + assert obs.category == ObservationCategory.NOTE + + +@pytest.mark.asyncio +async def test_get_entity_with_relations(client): + """Test retrieving an entity with relations.""" + # Create two entities that will have a relation + entity_request = CreateEntityRequest( + entities=[ + Entity(name="SourceEntity", entity_type="test"), + Entity(name="TargetEntity", entity_type="test") + ] + ) + await create_entities(entity_request) + + # Create relation between them (using the earlier tested create_relations) + from basic_memory.mcp.tools.knowledge import create_relations + from basic_memory.schemas.request import CreateRelationsRequest + from basic_memory.schemas.base import Relation + + relation_request = CreateRelationsRequest( + relations=[ + Relation( + from_id="test/source_entity", + to_id="test/target_entity", + relation_type="depends_on" + ) + ] + ) + await create_relations(relation_request) + + # Get and verify source entity + source = await get_entity("test/source_entity") + assert len(source.relations) == 1 + relation = source.relations[0] + assert relation.to_id == "test/target_entity" + assert relation.relation_type == "depends_on" + + +@pytest.mark.asyncio +async def test_get_entity_with_categorized_observations(client): + """Test retrieving an entity with observations in different categories.""" + # Create entity with categorized observations + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="TestEntity", + entity_type="test", + description="Test entity with categories" + ) + ] + ) + result = await create_entities(entity_request) + path_id = result.entities[0].path_id + + # Add observations with different categories + from basic_memory.mcp.tools.knowledge import add_observations + from basic_memory.schemas.request import AddObservationsRequest, ObservationCreate + + obs_request = AddObservationsRequest( + path_id=path_id, + observations=[ + ObservationCreate( + content="Technical detail", + category=ObservationCategory.TECH + ), + ObservationCreate( + content="Design decision", + category=ObservationCategory.DESIGN + ), + ObservationCreate( + content="Feature note", + category=ObservationCategory.FEATURE + ) + ] + ) + await add_observations(obs_request) + + # Get and verify entity + entity = await get_entity(path_id) + assert len(entity.observations) == 3 + categories = {obs.category for obs in entity.observations} + assert ObservationCategory.TECH in categories + assert ObservationCategory.DESIGN in categories + assert ObservationCategory.FEATURE in categories + + +@pytest.mark.asyncio +async def test_get_nonexistent_entity(client): + """Test attempting to get a non-existent entity.""" + with pytest.raises(EntityNotFoundError): + await get_entity("test/nonexistent") diff --git a/tests/mcp/test_tool_open_nodes.py b/tests/mcp/test_tool_open_nodes.py new file mode 100644 index 00000000..8c07c850 --- /dev/null +++ b/tests/mcp/test_tool_open_nodes.py @@ -0,0 +1,164 @@ +"""Tests for open_nodes MCP tool.""" + +import pytest + +from basic_memory.mcp.tools.search import open_nodes +from basic_memory.mcp.tools.knowledge import create_entities +from basic_memory.schemas.base import Entity +from basic_memory.schemas.request import CreateEntityRequest, OpenNodesRequest + + +@pytest.mark.asyncio +async def test_open_multiple_entities(client): + """Test opening multiple entities.""" + # Create some test entities + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="Entity1", + entity_type="test", + description="First test entity" + ), + Entity( + name="Entity2", + entity_type="test", + description="Second test entity" + ) + ] + ) + create_result = await create_entities(entity_request) + path_ids = [e.path_id for e in create_result.entities] + + # Open the nodes + request = OpenNodesRequest(path_ids=path_ids) + result = await open_nodes(request) + + # Verify we got a dictionary with both entities + assert len(result) == 2 + assert all(path_id in result for path_id in path_ids) + assert all(entity.name in ["Entity1", "Entity2"] for entity in result.values()) + + +@pytest.mark.asyncio +async def test_open_nodes_with_details(client): + """Test that opened nodes have all their details.""" + # Create an entity with observations + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="DetailedEntity", + entity_type="test", + description="Test entity with details", + observations=["First observation", "Second observation"] + ) + ] + ) + create_result = await create_entities(entity_request) + path_id = create_result.entities[0].path_id + + # Open the node + request = OpenNodesRequest(path_ids=[path_id]) + result = await open_nodes(request) + + # Verify all details are present + entity = result[path_id] + assert entity.name == "DetailedEntity" + assert entity.entity_type == "test" + assert entity.description == "Test entity with details" + assert len(entity.observations) == 2 + + +@pytest.mark.asyncio +async def test_open_nodes_with_relations(client): + """Test opening nodes that have relations.""" + # Create related entities + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="Service", + entity_type="test", + description="A service" + ), + Entity( + name="Database", + entity_type="test", + description="A database" + ) + ] + ) + create_result = await create_entities(entity_request) + path_ids = [e.path_id for e in create_result.entities] + + # Add a relation between them + from basic_memory.mcp.tools.knowledge import create_relations + from basic_memory.schemas.request import CreateRelationsRequest + from basic_memory.schemas.base import Relation + + relation_request = CreateRelationsRequest( + relations=[ + Relation( + from_id=path_ids[0], + to_id=path_ids[1], + relation_type="depends_on" + ) + ] + ) + await create_relations(relation_request) + + # Open both nodes + request = OpenNodesRequest(path_ids=path_ids) + result = await open_nodes(request) + + # Verify relations are present + assert len(result[path_ids[0]].relations) == 1 + assert len(result[path_ids[1]].relations) == 1 + + +@pytest.mark.asyncio +async def test_open_nonexistent_nodes(client): + """Test behavior when some requested nodes don't exist.""" + # First create one real entity + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="RealEntity", + entity_type="test" + ) + ] + ) + create_result = await create_entities(entity_request) + real_path_id = create_result.entities[0].path_id + + # Try to open both real and non-existent + request = OpenNodesRequest( + path_ids=[real_path_id, "test/nonexistent"] + ) + result = await open_nodes(request) + + # Should only get the real entity back + assert len(result) == 1 + assert real_path_id in result + + +@pytest.mark.asyncio +async def test_open_single_node(client): + """Test behavior with single path_id.""" + # Create an entity + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="SingleEntity", + entity_type="test" + ) + ] + ) + create_result = await create_entities(entity_request) + path_id = create_result.entities[0].path_id + + # Open just one node + request = OpenNodesRequest(path_ids=[path_id]) + result = await open_nodes(request) + + # Should get just that entity + assert len(result) == 1 + assert path_id in result \ No newline at end of file diff --git a/tests/mcp/test_tool_search_nodes.py b/tests/mcp/test_tool_search_nodes.py new file mode 100644 index 00000000..cd000ac6 --- /dev/null +++ b/tests/mcp/test_tool_search_nodes.py @@ -0,0 +1,173 @@ +"""Tests for search_nodes MCP tool.""" + +import pytest + +from basic_memory.mcp.tools.search import search_nodes +from basic_memory.mcp.tools.knowledge import create_entities +from basic_memory.schemas.base import Entity, ObservationCategory +from basic_memory.schemas.request import CreateEntityRequest, SearchNodesRequest, ObservationCreate + + +@pytest.mark.asyncio +async def test_basic_search(client): + """Test basic text search.""" + # Create some test entities + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="SearchComponent", + entity_type="component", + description="A searchable component", + observations=["This has some searchable text"] + ), + Entity( + name="OtherComponent", + entity_type="component", + description="Another component", + observations=["This is unrelated"] + ) + ] + ) + await create_entities(entity_request) + + # Search for "searchable" + request = SearchNodesRequest(query="searchable") + result = await search_nodes(request) + + # Should find one matching entity + assert len(result.matches) == 1 + assert result.matches[0].name == "SearchComponent" + assert result.query == "searchable" + + +@pytest.mark.asyncio +async def test_search_with_category(client): + """Test search with category filter.""" + # Create an entity with different observation categories + obs_tech = ObservationCreate( + content="Technical detail about implementation", + category=ObservationCategory.TECH + ) + obs_design = ObservationCreate( + content="Design decision about architecture", + category=ObservationCategory.DESIGN + ) + + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="TestEntity", + entity_type="test", + description="Test entity", + observations=[obs_tech.content, obs_design.content] + ) + ] + ) + await create_entities(entity_request) + + # Search for tech observations only + request = SearchNodesRequest( + query="implementation", + category=ObservationCategory.TECH + ) + tech_result = await search_nodes(request) + assert len(tech_result.matches) == 1 + + # Search for design observations only + request = SearchNodesRequest( + query="architecture", + category=ObservationCategory.DESIGN + ) + design_result = await search_nodes(request) + assert len(design_result.matches) == 1 + + +@pytest.mark.asyncio +async def test_search_multiple_matches(client): + """Test search returning multiple entities.""" + # Create multiple entities with similar content + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="Component1", + entity_type="component", + description="Uses SQLite database", + observations=["Implements SQLite storage"] + ), + Entity( + name="Component2", + entity_type="component", + description="Another SQLite component", + observations=["Also uses SQLite"] + ) + ] + ) + await create_entities(entity_request) + + # Search for SQLite + request = SearchNodesRequest(query="SQLite") + result = await search_nodes(request) + + # Should find both entities + assert len(result.matches) == 2 + names = {e.name for e in result.matches} + assert "Component1" in names + assert "Component2" in names + + +@pytest.mark.asyncio +async def test_search_no_matches(client): + """Test search with no matching results.""" + # Create an entity with unrelated content + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="UnrelatedEntity", + entity_type="test", + description="Something unrelated", + observations=["Nothing to see here"] + ) + ] + ) + await create_entities(entity_request) + + # Search for non-matching term + request = SearchNodesRequest(query="nonexistent") + result = await search_nodes(request) + + # Should find no matches + assert len(result.matches) == 0 + assert result.query == "nonexistent" + + +@pytest.mark.asyncio +async def test_search_case_insensitive(client): + """Test that search is case insensitive.""" + # Create entity with mixed case text + entity_request = CreateEntityRequest( + entities=[ + Entity( + name="MixedCase", + entity_type="test", + description="Testing MIXED case text", + observations=["Some MiXeD cAsE content"] + ) + ] + ) + await create_entities(entity_request) + + # Search with different cases + lower_request = SearchNodesRequest(query="mixed") + upper_request = SearchNodesRequest(query="MIXED") + mixed_request = SearchNodesRequest(query="MiXeD") + + # All should find the entity + lower_result = await search_nodes(lower_request) + upper_result = await search_nodes(upper_request) + mixed_result = await search_nodes(mixed_request) + + assert len(lower_result.matches) == 1 + assert len(upper_result.matches) == 1 + assert len(mixed_result.matches) == 1 + assert all(r.matches[0].name == "MixedCase" + for r in [lower_result, upper_result, mixed_result]) \ No newline at end of file