mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
fix tests
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
"""Tests for the MCP server implementation using FastAPI TestClient."""
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from fastapi import FastAPI
|
||||
from httpx import AsyncClient, ASGITransport
|
||||
|
||||
@@ -11,12 +12,7 @@ from basic_memory.deps import get_project_config, get_engine
|
||||
from basic_memory.schemas import CreateEntityResponse, SearchNodesResponse, AddObservationsResponse
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def anyio_backend():
|
||||
return "asyncio"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@pytest_asyncio.fixture
|
||||
def app(test_config, engine) -> FastAPI:
|
||||
"""Create test FastAPI application."""
|
||||
app = fastapi_app
|
||||
@@ -24,8 +20,13 @@ def app(test_config, engine) -> FastAPI:
|
||||
app.dependency_overrides[get_engine] = lambda: engine
|
||||
return app
|
||||
|
||||
@pytest_asyncio.fixture()
|
||||
async def server(app) -> MemoryServer:
|
||||
server = MemoryServer()
|
||||
await server.setup()
|
||||
return server
|
||||
|
||||
@pytest.fixture
|
||||
@pytest_asyncio.fixture
|
||||
async def client(app: FastAPI):
|
||||
"""Create test client that both MCP and tests will use."""
|
||||
async with AsyncClient(
|
||||
@@ -35,7 +36,7 @@ async def client(app: FastAPI):
|
||||
yield client
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@pytest_asyncio.fixture
|
||||
def test_entity_data():
|
||||
"""Sample data for creating a test entity."""
|
||||
return {
|
||||
@@ -48,7 +49,7 @@ def test_entity_data():
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@pytest_asyncio.fixture
|
||||
def test_directory_entity_data():
|
||||
"""Real data that caused failure in the tool."""
|
||||
return {
|
||||
@@ -65,10 +66,10 @@ def test_directory_entity_data():
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_list_tools(app):
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools(server):
|
||||
"""Test that server exposes expected tools."""
|
||||
server = MemoryServer()
|
||||
|
||||
tools = await server.handle_list_tools()
|
||||
|
||||
# Check each expected tool is present
|
||||
@@ -87,10 +88,9 @@ async def test_list_tools(app):
|
||||
assert search_schema["required"] == ["query"]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_directory_entity(test_directory_entity_data, client):
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_directory_entity(test_directory_entity_data, client, server):
|
||||
"""Test creating entity with exactly the data that failed in the tool."""
|
||||
server = MemoryServer()
|
||||
result = await server.handle_call_tool(
|
||||
"create_entities",
|
||||
test_directory_entity_data
|
||||
@@ -102,7 +102,7 @@ async def test_create_directory_entity(test_directory_entity_data, client):
|
||||
assert result[0].type == "resource"
|
||||
|
||||
# Verify entity creation
|
||||
response = CreateEntityResponse.model_validate_json(result[0].resource.text)
|
||||
response = CreateEntityResponse.model_validate_json(result[0].resource.text) # pyright: ignore [reportAttributeAccessIssue]
|
||||
assert len(response.entities) == 1
|
||||
created = response.entities[0]
|
||||
assert created.name == "Directory Organization"
|
||||
@@ -116,10 +116,9 @@ async def test_create_directory_entity(test_directory_entity_data, client):
|
||||
assert entity["name"] == "Directory Organization"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_search_nodes(test_entity_data, client):
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_nodes(test_entity_data, client, server):
|
||||
"""Test searching for an entity after creating it."""
|
||||
server = MemoryServer()
|
||||
|
||||
# First create an entity
|
||||
await server.handle_call_tool("create_entities", test_entity_data)
|
||||
@@ -139,7 +138,7 @@ async def test_search_nodes(test_entity_data, client):
|
||||
assert result[0].resource.mimeType == MIME_TYPE
|
||||
|
||||
# Verify search results
|
||||
response = SearchNodesResponse.model_validate_json(result[0].resource.text)
|
||||
response = SearchNodesResponse.model_validate_json(result[0].resource.text) # pyright: ignore [reportAttributeAccessIssue]
|
||||
assert len(response.matches) == 1
|
||||
assert response.matches[0].name == "Test Entity"
|
||||
assert response.query == "Test Entity"
|
||||
@@ -152,14 +151,13 @@ async def test_search_nodes(test_entity_data, client):
|
||||
assert data["matches"][0]["name"] == "Test Entity"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_add_observations(test_entity_data, client):
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_observations(test_entity_data, client, server):
|
||||
"""Test adding observations to an existing entity."""
|
||||
server = MemoryServer()
|
||||
|
||||
# First create an entity
|
||||
create_result = await server.handle_call_tool("create_entities", test_entity_data)
|
||||
create_response = CreateEntityResponse.model_validate_json(create_result[0].resource.text)
|
||||
create_response = CreateEntityResponse.model_validate_json(create_result[0].resource.text) # pyright: ignore [reportAttributeAccessIssue]
|
||||
entity_id = create_response.entities[0].id
|
||||
|
||||
# Add new observation
|
||||
@@ -190,22 +188,20 @@ async def test_add_observations(test_entity_data, client):
|
||||
assert "A new observation" in [o["content"] for o in entity["observations"]]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_invalid_tool_name():
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_tool_name(server):
|
||||
"""Test calling a non-existent tool."""
|
||||
server = MemoryServer()
|
||||
with pytest.raises(McpError) as exc:
|
||||
await server.handle_call_tool("not_a_tool", {})
|
||||
assert "Unknown tool" in str(exc.value)
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.asyncio
|
||||
class TestInputValidation:
|
||||
"""Test input validation for various tools."""
|
||||
|
||||
async def test_missing_required_field(self):
|
||||
async def test_missing_required_field(self, server):
|
||||
"""Test validation when required fields are missing."""
|
||||
server = MemoryServer()
|
||||
|
||||
with pytest.raises(McpError) as exc:
|
||||
await server.handle_call_tool("search_nodes", {})
|
||||
@@ -215,9 +211,8 @@ class TestInputValidation:
|
||||
await server.handle_call_tool("create_entities", {})
|
||||
assert "entities" in str(exc.value).lower()
|
||||
|
||||
async def test_empty_arrays(self):
|
||||
async def test_empty_arrays(self, server):
|
||||
"""Test validation of array fields that can't be empty."""
|
||||
server = MemoryServer()
|
||||
|
||||
with pytest.raises(McpError) as exc:
|
||||
await server.handle_call_tool("create_entities", {"entities": []})
|
||||
@@ -227,9 +222,8 @@ class TestInputValidation:
|
||||
await server.handle_call_tool("open_nodes", {"names": []})
|
||||
assert INVALID_PARAMS == exc.value.args[0]
|
||||
|
||||
async def test_invalid_field_types(self):
|
||||
async def test_invalid_field_types(self, server):
|
||||
"""Test validation when fields have wrong types."""
|
||||
server = MemoryServer()
|
||||
|
||||
with pytest.raises(McpError) as exc:
|
||||
await server.handle_call_tool("search_nodes", {"query": 123})
|
||||
@@ -239,9 +233,8 @@ class TestInputValidation:
|
||||
await server.handle_call_tool("create_entities", {"entities": "not an array"})
|
||||
assert "array" in str(exc.value).lower() or "list" in str(exc.value).lower()
|
||||
|
||||
async def test_invalid_nested_fields(self):
|
||||
async def test_invalid_nested_fields(self, server):
|
||||
"""Test validation of nested object fields."""
|
||||
server = MemoryServer()
|
||||
|
||||
with pytest.raises(McpError) as exc:
|
||||
await server.handle_call_tool("create_entities", {
|
||||
Reference in New Issue
Block a user