mirror of
https://github.com/basicmachines-co/basic-memory
synced 2026-06-21 13:47:35 +00:00
cleaned up schemas for api/mcp server
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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)])
|
||||
@@ -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
@@ -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": []})
|
||||
Reference in New Issue
Block a user