From 635c2851a759c9cc6881655915deee1bb2a1c251 Mon Sep 17 00:00:00 2001 From: phernandez Date: Sat, 14 Dec 2024 19:14:17 -0600 Subject: [PATCH] renamed all schemas to be Request/Response --- src/basic_memory/api/routers/knowledge.py | 18 +++---- src/basic_memory/fileio.py | 10 ++-- src/basic_memory/mcp/server.py | 12 ++--- src/basic_memory/schemas.py | 50 +++++++++---------- src/basic_memory/services/entity_service.py | 6 +-- src/basic_memory/services/memory_service.py | 14 +++--- src/basic_memory/services/relation_service.py | 6 +-- tests/conftest.py | 4 +- tests/test_entity_service.py | 26 +++++----- tests/test_memory_service.py | 4 +- tests/test_relation_service.py | 10 ++-- tests/test_schemas.py | 30 +++++------ 12 files changed, 95 insertions(+), 95 deletions(-) diff --git a/src/basic_memory/api/routers/knowledge.py b/src/basic_memory/api/routers/knowledge.py index a2206983..bbbc9645 100644 --- a/src/basic_memory/api/routers/knowledge.py +++ b/src/basic_memory/api/routers/knowledge.py @@ -7,7 +7,7 @@ from basic_memory.schemas import ( CreateEntitiesRequest, CreateEntitiesResponse, SearchNodesRequest, SearchNodesResponse, CreateRelationsRequest, CreateRelationsResponse, - EntityOut, RelationOut, AddObservationsRequest, ObservationOut, + EntityResponse, RelationResponse, AddObservationsRequest, ObservationResponse, OpenNodesRequest, OpenNodesResponse, DeleteEntitiesResponse, DeleteObservationsRequest, DeleteObservationsResponse, AddObservationsResponse @@ -23,17 +23,17 @@ async def create_entities( ) -> CreateEntitiesResponse: """Create new entities in the knowledge graph.""" entities = await memory_service.create_entities(data.entities) - return CreateEntitiesResponse(entities=[EntityOut.model_validate(entity) for entity in entities]) + return CreateEntitiesResponse(entities=[EntityResponse.model_validate(entity) for entity in entities]) -@router.get("/entities/{entity_id:path}", response_model=EntityOut) +@router.get("/entities/{entity_id:path}", response_model=EntityResponse) async def get_entity( entity_id: str, memory_service: MemoryServiceDep -) -> EntityOut: +) -> EntityResponse: """Get a specific entity by ID.""" entity = await memory_service.get_entity(entity_id) - return EntityOut.model_validate(entity) + return EntityResponse.model_validate(entity) @router.delete("/entities/{entity_id}", response_model=DeleteEntitiesResponse) @@ -53,7 +53,7 @@ async def open_nodes( ) -> OpenNodesResponse: """Open specific nodes by their names.""" entities = await memory_service.open_nodes(data.names) - return OpenNodesResponse(entities=[EntityOut.model_validate(entity) for entity in entities]) + return OpenNodesResponse(entities=[EntityResponse.model_validate(entity) for entity in entities]) @router.post("/relations", response_model=CreateRelationsResponse) @@ -63,7 +63,7 @@ async def create_relations( ) -> CreateRelationsResponse: """Create relations between entities.""" relations = await memory_service.create_relations(data.relations) - return CreateRelationsResponse(relations=[RelationOut.model_validate(relation) for relation in relations]) + return CreateRelationsResponse(relations=[RelationResponse.model_validate(relation) for relation in relations]) @router.delete("/relations/{relation_id}", response_model=DeleteEntitiesResponse) @@ -83,7 +83,7 @@ async def add_observations( ) -> AddObservationsResponse: """Add observations to an entity.""" observations = await memory_service.add_observations(data) - return AddObservationsResponse(entity_id=data.entity_id, observations=[ObservationOut.model_validate(observation) for observation in observations]) + return AddObservationsResponse(entity_id=data.entity_id, observations=[ObservationResponse.model_validate(observation) for observation in observations]) @router.delete("/observations", response_model=DeleteObservationsResponse) @@ -103,4 +103,4 @@ async def search_nodes( ) -> SearchNodesResponse: """Search for entities in the knowledge graph.""" matches = await memory_service.search_nodes(data.query) - return SearchNodesResponse(matches=[EntityOut.model_validate(entity) for entity in matches], query=data.query) \ No newline at end of file + return SearchNodesResponse(matches=[EntityResponse.model_validate(entity) for entity in matches], query=data.query) \ No newline at end of file diff --git a/src/basic_memory/fileio.py b/src/basic_memory/fileio.py index b00699c1..8c26c0cb 100644 --- a/src/basic_memory/fileio.py +++ b/src/basic_memory/fileio.py @@ -6,7 +6,7 @@ from pathlib import Path from loguru import logger -from basic_memory.schemas import EntityIn, RelationIn +from basic_memory.schemas import EntityRequest, RelationRequest class FileOperationError(Exception): @@ -34,7 +34,7 @@ def get_entity_path(project_entities_path: Path, entity_id: str) -> Path: return Path(f"{project_entities_path}/{entity_id}.md") -async def write_entity_file(project_entities_path: Path, entity_id: str, entity: EntityIn) -> bool: +async def write_entity_file(project_entities_path: Path, entity_id: str, entity: EntityRequest) -> bool: """ Write entity to filesystem in markdown format. @@ -102,7 +102,7 @@ async def write_entity_file(project_entities_path: Path, entity_id: str, entity: return True -async def read_entity_file(project_entities_path: Path, entity_id: str) -> EntityIn: +async def read_entity_file(project_entities_path: Path, entity_id: str) -> EntityRequest: """ Read entity data from filesystem. @@ -174,14 +174,14 @@ async def read_entity_file(project_entities_path: Path, entity_id: str) -> Entit relation_type = parts[0] context = parts[1] if len(parts) > 1 else None - relations.append(RelationIn( # pyright: ignore [reportCallIssue] + relations.append(RelationRequest( # pyright: ignore [reportCallIssue] from_id=entity_id, # pyright: ignore [reportCallIssue] to_id=target_id, # pyright: ignore [reportCallIssue] relation_type=relation_type, # pyright: ignore [reportCallIssue] context=context )) - return EntityIn( # pyright: ignore [reportCallIssue] + return EntityRequest( # pyright: ignore [reportCallIssue] id=entity_id, name=name, entity_type=entity_type, # pyright: ignore [reportCallIssue] diff --git a/src/basic_memory/mcp/server.py b/src/basic_memory/mcp/server.py index 25a41c6f..4214d44a 100644 --- a/src/basic_memory/mcp/server.py +++ b/src/basic_memory/mcp/server.py @@ -27,7 +27,7 @@ from basic_memory.schemas import ( # Tool responses CreateEntitiesResponse, SearchNodesResponse, OpenNodesResponse, AddObservationsResponse, CreateRelationsResponse, DeleteEntitiesResponse, - EntityOut, ObservationOut, RelationOut, AddObservationsRequest + EntityResponse, ObservationResponse, RelationResponse, AddObservationsRequest ) from basic_memory.services import EntityService, ObservationService, RelationService from basic_memory.services.memory_service import MemoryService @@ -111,7 +111,7 @@ async def handle_create_entities( logger.debug(f"Created {len(entities)} entities") # Format response - response = CreateEntitiesResponse(entities=[EntityOut.model_validate(entity) for entity in entities]) + response = CreateEntitiesResponse(entities=[EntityResponse.model_validate(entity) for entity in entities]) logger.debug("Formatted create_entities response") return create_response(response) @@ -126,7 +126,7 @@ async def handle_search_nodes( results = await service.search_nodes(input_args.query) logger.debug(f"Found {len(results)} matches for query '{input_args.query}'") response = SearchNodesResponse( - matches=[EntityOut.model_validate(entity) for entity in results], + matches=[EntityResponse.model_validate(entity) for entity in results], query=input_args.query ) return create_response(response) @@ -141,7 +141,7 @@ async def handle_open_nodes( input_args = OpenNodesRequest.model_validate(args) entities = await service.open_nodes(input_args.names) logger.debug(f"Opened {len(entities)} entities") - response = OpenNodesResponse(entities=[EntityOut.model_validate(entity) for entity in entities]) + response = OpenNodesResponse(entities=[EntityResponse.model_validate(entity) for entity in entities]) return create_response(response) @@ -162,7 +162,7 @@ async def handle_add_observations( # Format response response = AddObservationsResponse( entity_id=input_args.entity_id, - observations=[ObservationOut.model_validate(obs) for obs in observations] + observations=[ObservationResponse.model_validate(obs) for obs in observations] ) return create_response(response) @@ -182,7 +182,7 @@ async def handle_create_relations( logger.debug(f"Created {len(created)} relations") # Format response - response = CreateRelationsResponse(relations=[RelationOut.model_validate(relation) for relation in created]) + response = CreateRelationsResponse(relations=[RelationResponse.model_validate(relation) for relation in created]) return create_response(response) diff --git a/src/basic_memory/schemas.py b/src/basic_memory/schemas.py index 2a793fd3..e2afc7b8 100644 --- a/src/basic_memory/schemas.py +++ b/src/basic_memory/schemas.py @@ -4,7 +4,7 @@ from annotated_types import Len from pydantic import BaseModel, ConfigDict # Base output model for SQLAlchemy attribute conversion -class SQLAlchemyOut(BaseModel): +class SQLAlchemyModel(BaseModel): """Base class for models that read from SQLAlchemy attributes.""" model_config = ConfigDict(from_attributes=True) @@ -16,18 +16,18 @@ class AddObservationsRequest(BaseModel): observations: List[str] model_config = ConfigDict(populate_by_name=True) -class ObservationOut(SQLAlchemyOut): +class ObservationResponse(SQLAlchemyModel): """Schema for observation data returned from the service.""" id: int content: str -class ObservationsOut(SQLAlchemyOut): +class ObservationsResponse(SQLAlchemyModel): """Schema for bulk observation operation results.""" entity_id: str - observations: List[ObservationOut] + observations: List[ObservationResponse] model_config = ConfigDict(populate_by_name=True) -class RelationIn(BaseModel): +class RelationRequest(BaseModel): """ Represents a directed edge between entities in the knowledge graph. Relations are always stored in active voice (e.g. "created", "teaches", etc.) @@ -39,7 +39,7 @@ class RelationIn(BaseModel): model_config = ConfigDict(populate_by_name=True) -class RelationOut(SQLAlchemyOut): +class RelationResponse(SQLAlchemyModel): id: int from_id: str to_id: str @@ -60,26 +60,26 @@ class EntityBase(BaseModel): model_config = ConfigDict(from_attributes=True) -class EntityIn(EntityBase): +class EntityRequest(EntityBase): """ Represents a node in our knowledge graph - could be a person, project, concept, etc. Each entity has a unique name, a type, and a list of associated observations. """ observations: List[str] = [] - relations: List[RelationIn] = [] + relations: List[RelationRequest] = [] model_config = ConfigDict(populate_by_name=True) -class EntityOut(EntityBase, SQLAlchemyOut): +class EntityResponse(EntityBase, SQLAlchemyModel): """Schema for entity data returned from the service.""" - observations: List[ObservationOut] = [] - relations: List[RelationOut] = [] + observations: List[ObservationResponse] = [] + relations: List[RelationResponse] = [] model_config = ConfigDict(populate_by_name=True) # Tool Input Schemas class CreateEntitiesRequest(BaseModel): """Input schema for create_entities tool.""" - entities: Annotated[List[EntityIn], Len(min_length=1)] + entities: Annotated[List[EntityRequest], Len(min_length=1)] class SearchNodesRequest(BaseModel): """Input schema for search_nodes tool.""" @@ -91,7 +91,7 @@ class OpenNodesRequest(BaseModel): class CreateRelationsRequest(BaseModel): """Input schema for create_relations tool.""" - relations: List[RelationIn] + relations: List[RelationRequest] class DeleteEntitiesRequest(BaseModel): """Input schema for delete_entities tool.""" @@ -102,33 +102,33 @@ class DeleteObservationsRequest(BaseModel): entity_id: str deletions: List[str] # TODO: Make this more specific -class CreateEntitiesResponse(SQLAlchemyOut): +class CreateEntitiesResponse(SQLAlchemyModel): """Response for create_entities tool.""" - entities: List[EntityOut] + entities: List[EntityResponse] -class SearchNodesResponse(SQLAlchemyOut): +class SearchNodesResponse(SQLAlchemyModel): """Response for search_nodes tool.""" - matches: List[EntityOut] + matches: List[EntityResponse] query: str -class OpenNodesResponse(SQLAlchemyOut): +class OpenNodesResponse(SQLAlchemyModel): """Response for open_nodes tool.""" - entities: List[EntityOut] + entities: List[EntityResponse] -class AddObservationsResponse(SQLAlchemyOut): +class AddObservationsResponse(SQLAlchemyModel): """Response for add_observations tool.""" entity_id: str - observations: List[ObservationOut] + observations: List[ObservationResponse] -class CreateRelationsResponse(SQLAlchemyOut): +class CreateRelationsResponse(SQLAlchemyModel): """Response for create_relations tool.""" - relations: List[RelationOut] + relations: List[RelationResponse] -class DeleteEntitiesResponse(SQLAlchemyOut): +class DeleteEntitiesResponse(SQLAlchemyModel): """Response for delete_entities tool.""" deleted: List[str] -class DeleteObservationsResponse(SQLAlchemyOut): +class DeleteObservationsResponse(SQLAlchemyModel): """Response for delete_observations tool.""" entity_id: str deleted: List[str] \ No newline at end of file diff --git a/src/basic_memory/services/entity_service.py b/src/basic_memory/services/entity_service.py index 9fa3f1b5..e57e5a49 100644 --- a/src/basic_memory/services/entity_service.py +++ b/src/basic_memory/services/entity_service.py @@ -1,9 +1,9 @@ """Service for managing entities in the database.""" from pathlib import Path -from typing import List, Dict, Any, Sequence +from typing import Dict, Any, Sequence from basic_memory.repository.entity_repository import EntityRepository -from basic_memory.schemas import EntityIn +from basic_memory.schemas import EntityRequest from basic_memory.models import Entity from basic_memory.fileio import EntityNotFoundError from loguru import logger @@ -31,7 +31,7 @@ class EntityService: logger.exception(f"Failed to search entities with query: {query}") raise - async def create_entity(self, entity: EntityIn) -> Entity: + async def create_entity(self, entity: EntityRequest) -> Entity: """Create a new entity in the database.""" logger.debug(f"Creating entity in DB: {entity}") try: diff --git a/src/basic_memory/services/memory_service.py b/src/basic_memory/services/memory_service.py index a3d2e25b..762c3ab5 100644 --- a/src/basic_memory/services/memory_service.py +++ b/src/basic_memory/services/memory_service.py @@ -5,7 +5,7 @@ from pathlib import Path from basic_memory.models import Entity, Observation from basic_memory.schemas import ( - AddObservationsRequest, EntityIn, RelationIn + AddObservationsRequest, EntityRequest, RelationRequest ) from basic_memory.fileio import write_entity_file, read_entity_file, EntityNotFoundError from basic_memory.services import EntityService, RelationService, ObservationService @@ -32,12 +32,12 @@ class MemoryService: self.observation_service = observation_service logger.debug(f"Initialized MemoryService with path: {project_path}") - async def create_entities(self, entities_in: List[EntityIn]) -> List[Entity]: + async def create_entities(self, entities_in: List[EntityRequest]) -> List[Entity]: """Create multiple entities with their observations.""" logger.debug(f"Creating {len(entities_in)} entities") # Write files in parallel (filesystem is source of truth) - async def write_file(entity: EntityIn): + async def write_file(entity: EntityRequest): try: existing = await self.entity_service.get_by_type_and_name( entity.entity_type, @@ -60,7 +60,7 @@ class MemoryService: await asyncio.gather(*file_writes) logger.debug("Completed all file writes") - async def create_entity_in_db(entity_in: EntityIn): + async def create_entity_in_db(entity_in: EntityRequest): logger.debug(f"Creating entity in DB: {entity_in}") try: # Create base entity @@ -112,7 +112,7 @@ class MemoryService: logger.debug(f"Found entity {entity}") return entity - async def create_relations(self, relations_data: List[RelationIn]) -> List[RelationIn]: + async def create_relations(self, relations_data: List[RelationRequest]) -> List[RelationRequest]: """Create multiple relations between entities.""" logger.debug(f"Creating {len(relations_data)} relations") @@ -214,11 +214,11 @@ class MemoryService: logger.exception(f"Failed to search nodes with query: {query}") raise - async def open_nodes(self, names: List[str]) -> List[EntityIn]: + async def open_nodes(self, names: List[str]) -> List[EntityRequest]: """Get specific nodes and their relationships.""" logger.debug(f"Opening nodes: {names}") - async def read_node(name: str) -> Optional[EntityIn]: + async def read_node(name: str) -> Optional[EntityRequest]: try: # Get ID from name first logger.debug(f"Looking up entity: {name}") diff --git a/src/basic_memory/services/relation_service.py b/src/basic_memory/services/relation_service.py index 92861c05..9d50b700 100644 --- a/src/basic_memory/services/relation_service.py +++ b/src/basic_memory/services/relation_service.py @@ -2,7 +2,7 @@ from pathlib import Path from basic_memory.repository.relation_repository import RelationRepository -from basic_memory.schemas import EntityIn, RelationIn +from basic_memory.schemas import EntityRequest, RelationRequest from . import DatabaseSyncError @@ -16,14 +16,14 @@ class RelationService: self.project_path = project_path self.relation_repo = relation_repo - async def create_relation(self, relation: RelationIn) -> RelationIn: + async def create_relation(self, relation: RelationRequest) -> RelationRequest: """Create a new relation in the database.""" try: return await self.relation_repo.create(relation.model_dump()) except Exception as e: raise DatabaseSyncError(f"Failed to sync relation to database: {str(e)}") from e - async def delete_relation(self, from_entity: EntityIn, to_entity: EntityIn, relation_type: str) -> bool: + async def delete_relation(self, from_entity: EntityRequest, to_entity: EntityRequest, relation_type: str) -> bool: """Delete a specific relation between entities.""" try: # Use repository to find and delete the relation diff --git a/tests/conftest.py b/tests/conftest.py index 2e74a351..4616fa85 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -14,7 +14,7 @@ from basic_memory.deps import ( get_relation_service, get_relation_repo, get_observation_repo, get_entity_repo ) -from basic_memory.schemas import EntityIn +from basic_memory.schemas import EntityRequest from basic_memory.config import ProjectConfig from basic_memory.services import MemoryService @@ -115,7 +115,7 @@ async def sample_entity(entity_repository: EntityRepository): @pytest_asyncio.fixture async def test_entity(entity_service): """Create a test entity for reuse in tests.""" - entity_data = EntityIn( # pyright: ignore [reportCallIssue] + entity_data = EntityRequest( # pyright: ignore [reportCallIssue] name="Test Entity", entity_type="test", # pyright: ignore [reportCallIssue] ) diff --git a/tests/test_entity_service.py b/tests/test_entity_service.py index 118ac4ff..bececc89 100644 --- a/tests/test_entity_service.py +++ b/tests/test_entity_service.py @@ -3,14 +3,14 @@ import pytest from basic_memory.fileio import EntityNotFoundError from basic_memory.models import Entity -from basic_memory.schemas import EntityIn +from basic_memory.schemas import EntityRequest pytestmark = pytest.mark.asyncio async def test_create_entity_success(entity_service): """Test successful entity creation.""" - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", description="A test entity description" @@ -33,14 +33,14 @@ async def test_create_entity_success(entity_service): async def test_get_by_type_and_name(entity_service): """Test finding entity by type and name combination.""" # Create two entities with same name but different types - entity1_data = EntityIn( + entity1_data = EntityRequest( name="Test Entity", entity_type="type1", description="First test entity" ) entity1 = await entity_service.create_entity(entity1_data) - entity2_data = EntityIn( + entity2_data = EntityRequest( name="Test Entity", # Same name entity_type="type2", # Different type description="Second test entity" @@ -67,7 +67,7 @@ async def test_get_by_type_and_name(entity_service): async def test_create_entity_no_description(entity_service): """Test creating entity without description (should be None).""" - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", ) @@ -82,7 +82,7 @@ async def test_create_entity_no_description(entity_service): async def test_get_entity_success(entity_service): """Test successful entity retrieval.""" # Arrange - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", description="Test description" @@ -103,7 +103,7 @@ async def test_get_entity_success(entity_service): async def test_update_entity_description(entity_service): """Test updating an entity's description.""" # Create entity with description - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", description="Initial description" @@ -121,7 +121,7 @@ async def test_update_entity_description(entity_service): async def test_update_entity_description_to_none(entity_service): """Test updating an entity's description to None.""" # Create entity with description - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", description="Initial description" @@ -139,7 +139,7 @@ async def test_update_entity_description_to_none(entity_service): async def test_delete_entity_success(entity_service): """Test successful entity deletion.""" # Arrange - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", ) @@ -167,7 +167,7 @@ async def test_create_entity_db_error(entity_service, monkeypatch): raise Exception("Mock DB error") monkeypatch.setattr(entity_service.entity_repo, "create", mock_create) - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", description="Test description" @@ -188,7 +188,7 @@ async def test_create_entity_with_special_chars(entity_service): """Test entity creation with special characters in name and description.""" name = "Test & Entity! With @ Special #Chars" description = "Description with $pecial chars & symbols!" - entity_data = EntityIn( + entity_data = EntityRequest( name=name, entity_type="test", description=description @@ -204,7 +204,7 @@ async def test_create_entity_with_special_chars(entity_service): async def test_entity_id_generation(entity_service): """Test that entities get unique IDs generated correctly.""" - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", description="Test description", @@ -219,7 +219,7 @@ async def test_entity_id_generation(entity_service): async def test_create_entity_long_description(entity_service): """Test creating entity with a long description.""" long_description = "A" * 1000 # 1000 character description - entity_data = EntityIn( + entity_data = EntityRequest( name="Test Entity", entity_type="test", description=long_description diff --git a/tests/test_memory_service.py b/tests/test_memory_service.py index 042a3e62..83302f87 100644 --- a/tests/test_memory_service.py +++ b/tests/test_memory_service.py @@ -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 CreateEntitiesRequest, CreateRelationsRequest, AddObservationsRequest, RelationIn +from basic_memory.schemas import CreateEntitiesRequest, CreateRelationsRequest, AddObservationsRequest, RelationRequest test_entities_data = [ { @@ -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([RelationIn.model_validate(bad_relation)]) \ No newline at end of file + await memory_service.create_relations([RelationRequest.model_validate(bad_relation)]) \ No newline at end of file diff --git a/tests/test_relation_service.py b/tests/test_relation_service.py index 96397edd..f76f05d8 100644 --- a/tests/test_relation_service.py +++ b/tests/test_relation_service.py @@ -2,7 +2,7 @@ import pytest import pytest_asyncio -from basic_memory.schemas import EntityIn, RelationIn +from basic_memory.schemas import EntityRequest, RelationRequest pytestmark = pytest.mark.asyncio @@ -10,13 +10,13 @@ pytestmark = pytest.mark.asyncio @pytest_asyncio.fixture async def sample_entities(entity_service): """Create two sample entities for testing relations""" - entity1_data = EntityIn( + entity1_data = EntityRequest( name="test_entity_1", entity_type="test_type", observations=[], relations=[] ) - entity2_data = EntityIn( + entity2_data = EntityRequest( name="test_entity_2", entity_type="test_type", observations=[], @@ -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 = RelationIn( + relation_data = RelationRequest( 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 = RelationIn( + relation_data = RelationRequest( from_id=entity1.id, to_id=entity2.id, relation_type="test_relation", diff --git a/tests/test_schemas.py b/tests/test_schemas.py index fe6ba017..87f4e682 100644 --- a/tests/test_schemas.py +++ b/tests/test_schemas.py @@ -2,9 +2,9 @@ import pytest from pydantic import ValidationError from basic_memory.schemas import ( - EntityIn, - EntityOut, - RelationIn, + EntityRequest, + EntityResponse, + RelationRequest, CreateEntitiesRequest, SearchNodesRequest, OpenNodesRequest, @@ -16,7 +16,7 @@ def test_entity_in_minimal(): "name": "test_entity", "entity_type": "test" } - entity = EntityIn.model_validate(data) + entity = EntityRequest.model_validate(data) assert entity.name == "test_entity" assert entity.entity_type == "test" assert entity.description is None @@ -40,7 +40,7 @@ def test_entity_in_complete(): } ] } - entity = EntityIn.model_validate(data) + entity = EntityRequest.model_validate(data) assert entity.name == "test_entity" assert entity.entity_type == "test" assert entity.description == "A test entity" @@ -52,13 +52,13 @@ def test_entity_in_complete(): def test_entity_in_validation(): """Test validation errors for EntityIn.""" with pytest.raises(ValidationError): - EntityIn.model_validate({}) # Missing required fields + EntityRequest.model_validate({}) # Missing required fields with pytest.raises(ValidationError): - EntityIn.model_validate({"name": "test"}) # Missing entityType + EntityRequest.model_validate({"name": "test"}) # Missing entityType with pytest.raises(ValidationError): - EntityIn.model_validate({"entityType": "test"}) # Missing name + EntityRequest.model_validate({"entityType": "test"}) # Missing name def test_relation_in_validation(): """Test RelationIn validation.""" @@ -67,7 +67,7 @@ def test_relation_in_validation(): "to_id": "456", "relation_type": "test" } - relation = RelationIn.model_validate(data) + relation = RelationRequest.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 = RelationIn.model_validate(data) + relation = RelationRequest.model_validate(data) assert relation.context == "test context" # Missing required fields with pytest.raises(ValidationError): - RelationIn.model_validate({"from_id": "123", "to_id": "456"}) # Missing relationType + RelationRequest.model_validate({"from_id": "123", "to_id": "456"}) # Missing relationType def test_create_entities_input(): """Test CreateEntitiesInput validation.""" @@ -126,7 +126,7 @@ def test_entity_out_from_attributes(): } ] } - entity = EntityOut.model_validate(db_data) + entity = EntityResponse.model_validate(db_data) assert entity.id == "123" assert entity.description == "test description" assert len(entity.observations) == 1 @@ -137,13 +137,13 @@ def test_entity_out_from_attributes(): def test_optional_fields(): """Test handling of optional fields.""" # Create with no optional fields - entity = EntityIn.model_validate({"name": "test", "entity_type": "test"}) + entity = EntityRequest.model_validate({"name": "test", "entity_type": "test"}) assert entity.description is None assert entity.observations == [] assert entity.relations == [] # Create with empty optional fields - entity = EntityIn.model_validate({ + entity = EntityRequest.model_validate({ "name": "test", "entity_type": "test", "description": None, @@ -155,7 +155,7 @@ def test_optional_fields(): assert entity.relations == [] # Create with some optional fields - entity = EntityIn.model_validate({ + entity = EntityRequest.model_validate({ "name": "test", "entity_type": "test", "description": "test",