cleaned up schemas for api/mcp server

This commit is contained in:
phernandez
2024-12-14 19:09:56 -06:00
parent 9b359b884c
commit 2b4bd19da1
10 changed files with 81 additions and 89 deletions
+2 -2
View File
@@ -186,8 +186,8 @@ async def test_add_observations(test_entity_data, memory_service, test_config):
response = AddObservationsResponse.model_validate_json(result[0].resource.text) # pyright: ignore [reportAttributeAccessIssue]
assert response.entity_id == entity_id
assert len(response.added_observations) == 1
assert response.added_observations[0].content == "A new observation"
assert len(response.observations) == 1
assert response.observations[0].content == "A new observation"
@pytest.mark.anyio
async def test_invalid_tool_name(test_config):
+9 -9
View File
@@ -2,7 +2,7 @@
import pytest
from basic_memory.services import MemoryService
from basic_memory.fileio import read_entity_file
from basic_memory.schemas import CreateEntitiesInput, CreateRelationsInput, ObservationsIn, Relation
from basic_memory.schemas import CreateEntitiesRequest, CreateRelationsRequest, AddObservationsRequest, RelationIn
test_entities_data = [
{
@@ -21,7 +21,7 @@ test_entities_data = [
async def test_create_entities(memory_service: MemoryService):
"""Should create multiple entities in parallel with their observations."""
entity_input = CreateEntitiesInput.model_validate({"entities": test_entities_data})
entity_input = CreateEntitiesRequest.model_validate({"entities": test_entities_data})
entities = await memory_service.create_entities(entity_input.entities)
# Verify the SQLAlchemy models were created
@@ -51,7 +51,7 @@ async def test_create_entities(memory_service: MemoryService):
@pytest.mark.asyncio
async def test_add_observations(memory_service: MemoryService):
"""Should add observations to an existing entity."""
entity_input = CreateEntitiesInput.model_validate({"entities": test_entities_data})
entity_input = CreateEntitiesRequest.model_validate({"entities": test_entities_data})
entities = await memory_service.create_entities([entity_input.entities[0]])
entity = entities[0]
@@ -65,7 +65,7 @@ async def test_add_observations(memory_service: MemoryService):
}
# Add observations - returns List[models.Observation]
observation_input = ObservationsIn.model_validate(observations_data)
observation_input = AddObservationsRequest.model_validate(observations_data)
added_observations = await memory_service.add_observations(observation_input)
# Check the SQLAlchemy model results
@@ -94,13 +94,13 @@ async def test_add_observations_nonexistent_entity(memory_service: MemoryService
}
with pytest.raises(Exception) as exc: # We might want to define a specific error type
observation_input = ObservationsIn.model_validate(observations_data)
observation_input = AddObservationsRequest.model_validate(observations_data)
await memory_service.add_observations(observation_input)
@pytest.mark.asyncio
async def test_create_relations(memory_service: MemoryService):
"""Should create relations between entities and update both filesystem and database."""
entity_input = CreateEntitiesInput.model_validate({"entities": test_entities_data})
entity_input = CreateEntitiesRequest.model_validate({"entities": test_entities_data})
entities = await memory_service.create_entities(entity_input.entities)
entity1, entity2 = entities
@@ -120,7 +120,7 @@ async def test_create_relations(memory_service: MemoryService):
]
# Create relations - returns List[models.Relation]
input_args = CreateRelationsInput.model_validate({"relations": test_relations_data})
input_args = CreateRelationsRequest.model_validate({"relations": test_relations_data})
relations = await memory_service.create_relations(input_args.relations)
# Verify SQLAlchemy Relation models were created
@@ -183,7 +183,7 @@ async def test_create_relations(memory_service: MemoryService):
async def test_create_relations_with_invalid_entity_id(memory_service: MemoryService):
"""Should raise an appropriate error when trying to create relations with non-existent entity IDs."""
# Create one entity - returns SQLAlchemy Entity
entity_input = CreateEntitiesInput.model_validate({"entities": test_entities_data})
entity_input = CreateEntitiesRequest.model_validate({"entities": test_entities_data})
entities = await memory_service.create_entities([entity_input.entities[0]])
entity1 = entities[0]
@@ -195,4 +195,4 @@ async def test_create_relations_with_invalid_entity_id(memory_service: MemorySer
}
with pytest.raises(Exception) as exc: # We might want to define a specific error type
await memory_service.create_relations([Relation.model_validate(bad_relation)])
await memory_service.create_relations([RelationIn.model_validate(bad_relation)])
+3 -3
View File
@@ -2,7 +2,7 @@
import pytest
import pytest_asyncio
from basic_memory.schemas import EntityIn, Relation
from basic_memory.schemas import EntityIn, RelationIn
pytestmark = pytest.mark.asyncio
@@ -36,7 +36,7 @@ async def test_create_relation(relation_service, sample_entities):
"""Test creating a basic relation between two entities"""
entity1, entity2 = sample_entities
relation_data = Relation(
relation_data = RelationIn(
from_id=entity1.id,
to_id=entity2.id,
relation_type="test_relation"
@@ -62,7 +62,7 @@ async def test_create_relation_with_context(relation_service, sample_entities):
"""Test creating a relation with context information"""
entity1, entity2 = sample_entities
relation_data = Relation(
relation_data = RelationIn(
from_id=entity1.id,
to_id=entity2.id,
relation_type="test_relation",
+13 -13
View File
@@ -4,10 +4,10 @@ from pydantic import ValidationError
from basic_memory.schemas import (
EntityIn,
EntityOut,
Relation,
CreateEntitiesInput,
SearchNodesInput,
OpenNodesInput,
RelationIn,
CreateEntitiesRequest,
SearchNodesRequest,
OpenNodesRequest,
)
def test_entity_in_minimal():
@@ -67,7 +67,7 @@ def test_relation_in_validation():
"to_id": "456",
"relation_type": "test"
}
relation = Relation.model_validate(data)
relation = RelationIn.model_validate(data)
assert relation.from_id == "123"
assert relation.to_id == "456"
assert relation.relation_type == "test"
@@ -75,12 +75,12 @@ def test_relation_in_validation():
# With context
data["context"] = "test context"
relation = Relation.model_validate(data)
relation = RelationIn.model_validate(data)
assert relation.context == "test context"
# Missing required fields
with pytest.raises(ValidationError):
Relation.model_validate({"from_id": "123", "to_id": "456"}) # Missing relationType
RelationIn.model_validate({"from_id": "123", "to_id": "456"}) # Missing relationType
def test_create_entities_input():
"""Test CreateEntitiesInput validation."""
@@ -97,13 +97,13 @@ def test_create_entities_input():
}
]
}
create_input = CreateEntitiesInput.model_validate(data)
create_input = CreateEntitiesRequest.model_validate(data)
assert len(create_input.entities) == 2
assert create_input.entities[1].description == "test description"
# Empty entities list should fail
with pytest.raises(ValidationError):
CreateEntitiesInput.model_validate({"entities": []})
CreateEntitiesRequest.model_validate({"entities": []})
def test_entity_out_from_attributes():
"""Test EntityOut creation from database model attributes."""
@@ -167,17 +167,17 @@ def test_optional_fields():
def test_search_nodes_input():
"""Test SearchNodesInput validation."""
search = SearchNodesInput.model_validate({"query": "test query"})
search = SearchNodesRequest.model_validate({"query": "test query"})
assert search.query == "test query"
with pytest.raises(ValidationError):
SearchNodesInput.model_validate({}) # Missing required query
SearchNodesRequest.model_validate({}) # Missing required query
def test_open_nodes_input():
"""Test OpenNodesInput validation."""
open_input = OpenNodesInput.model_validate({"names": ["entity1", "entity2"]})
open_input = OpenNodesRequest.model_validate({"names": ["entity1", "entity2"]})
assert len(open_input.names) == 2
# Empty names list should fail
with pytest.raises(ValidationError):
OpenNodesInput.model_validate({"names": []})
OpenNodesRequest.model_validate({"names": []})