finish implementing tests for tools

This commit is contained in:
phernandez
2024-12-28 14:59:20 -06:00
parent 0d86bfeb7a
commit cc5d0d103a
8 changed files with 959 additions and 208 deletions
+16 -3
View File
@@ -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()
+93 -71
View File
@@ -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"
}]
)
await add_observations(request)
+67 -48
View File
@@ -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"
+123 -86
View File
@@ -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
second_result = await create_relations(relation_request)
# Current behavior: No entities returned when duplicate relation fails
assert len(second_result.entities) == 0
+191
View File
@@ -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"}
+132
View File
@@ -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")
+164
View File
@@ -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
+173
View File
@@ -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])