fix tests

This commit is contained in:
phernandez
2024-12-15 16:17:25 -06:00
parent cc6d9d9fc3
commit ae5d7d52c5
2 changed files with 29 additions and 36 deletions
@@ -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", {