diff --git a/db/migrations/20241211190000_add_entity_type_name_unique.sql b/db/migrations/20241211190000_add_entity_type_name_unique.sql new file mode 100644 index 00000000..1acecece --- /dev/null +++ b/db/migrations/20241211190000_add_entity_type_name_unique.sql @@ -0,0 +1,8 @@ +-- migrate:up +-- Add unique index on entity type and name combination +CREATE UNIQUE INDEX idx_entity_type_name ON entity(entity_type, name); + +-- migrate:down + +-- Restore original schema +DROP INDEX IF EXISTS idx_entity_type_name; \ No newline at end of file diff --git a/db/schema.sql b/db/schema.sql index c9c71fb7..da9a6c1f 100644 --- a/db/schema.sql +++ b/db/schema.sql @@ -30,9 +30,11 @@ CREATE TABLE IF NOT EXISTS "entity" ( created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP ); +CREATE UNIQUE INDEX idx_entity_type_name ON entity(entity_type, name); -- Dbmate schema migrations INSERT INTO "schema_migrations" (version) VALUES ('20240101000000'), ('20241210213454'), ('20241211034719'), - ('20241211052101'); + ('20241211052101'), + ('20241211190000'); diff --git a/src/basic_memory/fileio.py b/src/basic_memory/fileio.py index b3e9d2de..b6c749a0 100644 --- a/src/basic_memory/fileio.py +++ b/src/basic_memory/fileio.py @@ -4,6 +4,8 @@ Handles reading and writing entities and observations to the filesystem. """ from pathlib import Path +from loguru import logger + from basic_memory.schemas import EntityIn, ObservationIn, RelationIn @@ -31,7 +33,9 @@ async def write_entity_file(entities_path: Path, entity: EntityIn) -> bool: Raises: FileOperationError: If file operations fail """ - entity_path = entities_path / f"{entity.id}.md" + logger.debug(f"Writing entity file for {entity.file_path}") + + entity_path = entities_path / entity.file_path # Handle directory creation separately try: @@ -79,7 +83,8 @@ async def write_entity_file(entities_path: Path, entity: EntityIn) -> bool: temp_path.rename(entity_path) except Exception as e: raise FileOperationError(f"Failed to finalize entity file: {str(e)}") from e - + + logger.debug(f"Wrote entity file: {entity.file_path}") return True @@ -154,21 +159,15 @@ async def read_entity_file(entities_path: Path, entity_id: str) -> EntityIn: parts = rest.split(" | ", 1) relation_type = parts[0] context = parts[1] if len(parts) > 1 else None - - # Create temporary entities for the relation - # TODO what is this for? - target_entity = EntityIn(id=target_id, name=target_id, entity_type="unknown") - source_entity = EntityIn(id=entity_id, name=name, entity_type=entity_type) - + relations.append(RelationIn( - from_id=source_entity.id, - to_id=target_entity.id, + from_id=entity_id, + to_id=target_id, relation_type=relation_type, context=context )) return EntityIn( - id=entity_id, name=name, entity_type=entity_type, observations=observations, diff --git a/src/basic_memory/models.py b/src/basic_memory/models.py index 809aea73..d4d97285 100644 --- a/src/basic_memory/models.py +++ b/src/basic_memory/models.py @@ -1,7 +1,7 @@ """Database models for basic-memory.""" from datetime import datetime, UTC from typing import List, Optional -from sqlalchemy import String, DateTime, ForeignKey, Text, TypeDecorator, Integer, text +from sqlalchemy import String, DateTime, ForeignKey, Text, TypeDecorator, Integer, text, UniqueConstraint from sqlalchemy.orm import Mapped, mapped_column, relationship, DeclarativeBase from sqlalchemy.ext.asyncio import AsyncAttrs @@ -39,16 +39,19 @@ class Entity(Base): Core entity in the knowledge graph. Entities are the primary nodes in the knowledge graph. Each entity has: - - A unique identifier (text, for filesystem references) + - A unique identifier (text, based on type/name path) - A name - An entity type (e.g., "person", "organization", "event") - A description (optional) - A list of observations """ __tablename__ = "entity" + __table_args__ = ( + UniqueConstraint('entity_type', 'name', name='uix_entity_type_name'), + ) id: Mapped[str] = mapped_column(String, primary_key=True) - name: Mapped[str] = mapped_column(String, index=True) + name: Mapped[str] = mapped_column(String) entity_type: Mapped[str] = mapped_column(String) description: Mapped[Optional[str]] = mapped_column(Text, nullable=True) created_at: Mapped[datetime] = mapped_column( @@ -80,6 +83,17 @@ class Entity(Base): cascade="all, delete-orphan" ) + @classmethod + def generate_id(cls, entity_type: str, name: str) -> str: + """Generate a filesystem path-based ID for this entity.""" + # Normalize name for filesystem (handle spaces, special chars etc) + safe_name = name.lower().replace(" ", "_") + return f"{entity_type}/{safe_name}" + + def get_file_path(self) -> str: + """Get the filesystem path for this entity.""" + return f"{self.id}.md" # id is already in path format + def __repr__(self) -> str: return f"Entity(id='{self.id}', name='{self.name}', type='{self.entity_type}')" diff --git a/src/basic_memory/repository/__init__.py b/src/basic_memory/repository/__init__.py index b7895c93..88a2aa45 100644 --- a/src/basic_memory/repository/__init__.py +++ b/src/basic_memory/repository/__init__.py @@ -6,7 +6,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Mapped from loguru import logger -from basic_memory.models import Base +from basic_memory.models import Base, Entity T = TypeVar('T', bound=Base) @@ -72,8 +72,16 @@ class Repository[T: Base]: # Only include valid columns that are provided in entity_data model_data = { k: v for k, v in entity_data.items() - if k in self.valid_columns + if k in self.valid_columns and v is not None } + + # Generate ID if this is an Entity model and no ID provided + if model is Entity and 'id' not in model_data: + model_data['id'] = Entity.generate_id( + model_data['entity_type'], + model_data['name'] + ) + logger.debug(f"Filtered data for valid columns: {model_data}") # Create insert statement with only provided data diff --git a/src/basic_memory/schemas.py b/src/basic_memory/schemas.py index 4665cd08..2030b9e1 100644 --- a/src/basic_memory/schemas.py +++ b/src/basic_memory/schemas.py @@ -5,9 +5,8 @@ independent from storage/persistence concerns. """ from datetime import datetime, UTC from typing import List, Optional, Dict, Any, Annotated -from uuid import uuid4 from annotated_types import Gt, Len -from pydantic import BaseModel, Field, model_validator, ConfigDict +from pydantic import BaseModel, Field, ConfigDict # Base output model for SQLAlchemy attribute conversion class SQLAlchemyOut(BaseModel): @@ -57,23 +56,14 @@ class RelationOut(SQLAlchemyOut): model_config = ConfigDict(populate_by_name=True) class EntityBase(BaseModel): - id: str = Field(default=None) # Allow None during creation name: str entity_type: str = Field(alias="entityType") description: Optional[str] = None - @model_validator(mode='after') - def generate_id(self) -> 'EntityBase': - """Generate an ID for this entity if not provided""" - if not self.id: - timestamp = datetime.now(UTC).strftime("%Y%m%d") - normalized_name = self.name.lower().replace(" ", "-") - self.id = f"{timestamp}-{normalized_name}" - return self - - def file_name(self) -> str: - """Get the markdown file name for this entity.""" - return f"{self.id}.md" + @property + def file_path(self) -> str: + """The relative file path for this entity.""" + return f"{self.entity_type}/{self.name}.md" model_config = ConfigDict(from_attributes=True) @@ -88,6 +78,7 @@ class EntityIn(EntityBase): model_config = ConfigDict(populate_by_name=True) class EntityOut(EntityBase, SQLAlchemyOut): + id: str # ID will be set by repository """Schema for entity data returned from the service.""" observations: List[ObservationOut] = [] relations: List[RelationOut] = [] diff --git a/src/basic_memory/services/entity_service.py b/src/basic_memory/services/entity_service.py index dbef6258..7890a5bb 100644 --- a/src/basic_memory/services/entity_service.py +++ b/src/basic_memory/services/entity_service.py @@ -34,7 +34,7 @@ class EntityService: async def create_entity(self, entity: EntityIn) -> Entity: """Create a new entity in the database.""" - logger.debug(f"Creating entity in DB: {entity.id}") + logger.debug(f"Creating entity in DB: {entity}") try: created_entity = await self.entity_repo.create(entity.model_dump()) logger.debug(f"Created base entity: {created_entity.id}") @@ -44,7 +44,7 @@ class EntityService: return created_entity except Exception as e: - logger.exception(f"Failed to create entity: {entity.id}") + logger.exception(f"Failed to create entity: {entity}") raise async def update_entity(self, entity_id: str, update_data: Dict[str, Any]) -> Entity: @@ -80,21 +80,21 @@ class EntityService: logger.exception(f"Failed to get entity: {entity_id}") raise - async def get_by_name(self, name: str) -> Entity: - """Get entity by name.""" - logger.debug(f"Getting entity by name: {name}") + async def get_by_type_and_name(self, entity_type: str, name: str) -> Entity: + """Get entity by type and name combination.""" + logger.debug(f"Getting entity by type/name: {entity_type}/{name}") try: - db_entity = await self.entity_repo.find_by_name(name) + db_entity = await self.entity_repo.find_by_type_and_name(entity_type, name) if not db_entity: - logger.error(f"Entity not found: {name}") - raise EntityNotFoundError(f"Entity not found: {name}") + logger.error(f"Entity not found: {entity_type}/{name}") + raise EntityNotFoundError(f"Entity not found: {entity_type}/{name}") logger.debug(f"Found entity: {db_entity.id}") return db_entity except EntityNotFoundError: raise except Exception as e: - logger.exception(f"Failed to get entity by name: {name}") + logger.exception(f"Failed to get entity by type/name: {entity_type}/{name}") raise async def delete_entity(self, entity_id: str) -> bool: diff --git a/src/basic_memory/services/memory_service.py b/src/basic_memory/services/memory_service.py index 7279a2f8..f0b1c857 100644 --- a/src/basic_memory/services/memory_service.py +++ b/src/basic_memory/services/memory_service.py @@ -35,9 +35,7 @@ class MemoryService: # Write files in parallel (filesystem is source of truth) async def write_file(entity: EntityIn): - logger.debug(f"Writing entity file for {entity.id}") await write_entity_file(self.entities_path, entity) - logger.debug(f"Wrote entity file: {self.entities_path}/{entity.file_name()}") file_writes = [write_file(entity) for entity in entities_in] logger.debug("Starting parallel file writes") @@ -45,27 +43,27 @@ class MemoryService: logger.debug("Completed all file writes") async def create_entity_in_db(entity_in: EntityIn): - logger.debug(f"Creating entity in DB: {entity_in.id}") + logger.debug(f"Creating entity in DB: {entity_in}") try: # Create base entity - await self.entity_service.create_entity(entity_in) - logger.debug(f"Created base entity: {entity_in.id}") + created_entity = await self.entity_service.create_entity(entity_in) + logger.debug(f"Created base entity: {created_entity.id}") # Add observations - await self.observation_service.add_observations(entity_in, entity_in.observations) - logger.debug(f"Added {len(entity_in.observations)} observations to {entity_in.id}") + await self.observation_service.add_observations(created_entity.id, entity_in.observations) + logger.debug(f"Added {len(entity_in.observations)} observations to {created_entity.id}") # Add relations for relation in entity_in.relations: await self.relation_service.create_relation(relation) - logger.debug(f"Added {len(entity_in.relations)} relations for {entity_in.id}") + logger.debug(f"Added {len(entity_in.relations)} relations for {created_entity.id}") # Query final state - final_entity = await self.entity_service.get_entity(entity_in.id) + final_entity = await self.entity_service.get_entity(created_entity.id) logger.debug(f"Retrieved final entity state: {final_entity}") return final_entity except Exception as e: - logger.exception(f"Failed to create entity in DB: {entity_in.id}") + logger.exception(f"Failed to create entity in DB: {entity_in}") raise # Update database index sequentially @@ -89,13 +87,13 @@ class MemoryService: # First read complete entities from filesystem from_entity = await read_entity_file(self.entities_path, relation.from_id) to_entity = await read_entity_file(self.entities_path, relation.to_id) - logger.debug(f"Read entities for relation: {from_entity.id}, {to_entity.id}") + logger.debug(f"Read entities for relation: {from_entity.file_path}, {to_entity.file_path}") # Add the new relation to the source entity if not hasattr(from_entity, 'relations'): from_entity.relations = [] from_entity.relations.append(relation) - logger.debug(f"Added relation to source entity: {from_entity.id}") + logger.debug(f"Added relation to source entity: {from_entity.file_path}") # Write updated entity files (filesystem is source of truth) logger.debug("Writing updated entity files") @@ -126,7 +124,7 @@ class MemoryService: # Read entity from filesystem using the ID entity = await read_entity_file(self.entities_path, db_entity.id) - logger.debug(f"Read entity from filesystem: {entity.id}") + logger.debug(f"Read entity from filesystem: {db_entity.id}") # Create new observations for the entity for obs in observations_in.observations: @@ -139,10 +137,9 @@ class MemoryService: logger.debug("Wrote updated entity file") # Update database index - added_observations = await self.observation_service.add_observations(entity, observations_in.observations) + added_observations = await self.observation_service.add_observations(db_entity.id, observations_in.observations) logger.debug(f"Added {len(added_observations)} observations to DB") - db_entity = await self.entity_service.get_entity(entity.id) return added_observations except Exception as e: logger.exception(f"Failed to add observations to entity: {observations_in.entity_id}") diff --git a/src/basic_memory/services/observation_service.py b/src/basic_memory/services/observation_service.py index 3cc45837..0a99e39c 100644 --- a/src/basic_memory/services/observation_service.py +++ b/src/basic_memory/services/observation_service.py @@ -19,7 +19,7 @@ class ObservationService: self.project_path = project_path self.observation_repo = observation_repo - async def add_observations(self, entity: EntityIn, observations: List[ObservationIn]) -> List[Observation]: + async def add_observations(self, entity_id: str, observations: List[ObservationIn]) -> List[Observation]: """ Add multiple observations to an entity. Returns the created observations with IDs set. @@ -28,7 +28,7 @@ class ObservationService: try: obs = await self.observation_repo.create({ **observation.model_dump(), - 'entity_id': entity.id + 'entity_id': entity_id }) # Ensure observation is flushed and refreshed await self.observation_repo.session.flush() diff --git a/tests/test_entity_repository.py b/tests/test_entity_repository.py index 6c8720e1..cecb3e1f 100644 --- a/tests/test_entity_repository.py +++ b/tests/test_entity_repository.py @@ -2,6 +2,7 @@ import pytest from datetime import datetime, UTC from sqlalchemy import text, select +from sqlalchemy.exc import IntegrityError from basic_memory.models import Entity from basic_memory.repository.entity_repository import EntityRepository @@ -13,7 +14,6 @@ class TestEntityRepository: async def test_create_entity(self, entity_repository: EntityRepository): """Test creating a new entity""" entity_data = { - 'id': '20240102-test', 'name': 'Test', 'entity_type': 'test', 'description': 'Test description', @@ -21,7 +21,7 @@ class TestEntityRepository: entity = await entity_repository.create(entity_data) # Verify returned object - assert entity.id == '20240102-test' + assert entity.id == f'test/test' assert entity.name == 'Test' assert entity.description == 'Test description' assert isinstance(entity.created_at, datetime) @@ -35,6 +35,42 @@ class TestEntityRepository: assert db_entity.name == entity.name assert db_entity.description == entity.description + async def test_entity_type_name_unique_constraint(self, entity_repository: EntityRepository): + """Test the unique constraint on entity_type + name combination.""" + # Create first entity + entity1_data = { + 'id': '20240102-test1', + 'name': 'Test Entity', + 'entity_type': 'type1', + 'description': 'First entity' + } + await entity_repository.create(entity1_data) + + # Try to create another entity with same type and name + entity2_data = { + 'id': '20240102-test2', + 'name': 'Test Entity', # Same name + 'entity_type': 'type1', # Same type + 'description': 'Second entity' + } + + # Should raise IntegrityError + with pytest.raises(IntegrityError) as exc_info: + await entity_repository.create(entity2_data) + assert 'UNIQUE constraint failed: entity.entity_type, entity.name' in str(exc_info.value) + + # But should allow same name with different type + entity3_data = { + 'id': '20240102-test3', + 'name': 'Test Entity', # Same name + 'entity_type': 'type2', # Different type + 'description': 'Third entity' + } + entity3 = await entity_repository.create(entity3_data) + assert entity3 is not None + assert entity3.name == 'Test Entity' + assert entity3.entity_type == 'type2' + async def test_create_entity_null_description(self, entity_repository: EntityRepository): """Test creating an entity with null description""" entity_data = { diff --git a/tests/test_entity_service.py b/tests/test_entity_service.py index 6f381a84..118ac4ff 100644 --- a/tests/test_entity_service.py +++ b/tests/test_entity_service.py @@ -30,6 +30,41 @@ async def test_create_entity_success(entity_service): retrieved = await entity_service.get_entity(entity.id) assert retrieved.description == "A test entity description" +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( + name="Test Entity", + entity_type="type1", + description="First test entity" + ) + entity1 = await entity_service.create_entity(entity1_data) + + entity2_data = EntityIn( + name="Test Entity", # Same name + entity_type="type2", # Different type + description="Second test entity" + ) + entity2 = await entity_service.create_entity(entity2_data) + + # Find by type1 and name + found = await entity_service.get_by_type_and_name("type1", "Test Entity") + assert found is not None + assert found.id == entity1.id + assert found.entity_type == "type1" + assert found.description == "First test entity" + + # Find by type2 and name + found = await entity_service.get_by_type_and_name("type2", "Test Entity") + assert found is not None + assert found.id == entity2.id + assert found.entity_type == "type2" + assert found.description == "Second test entity" + + # Test not found case + with pytest.raises(EntityNotFoundError): + await entity_service.get_by_type_and_name("nonexistent", "Test Entity") + async def test_create_entity_no_description(entity_service): """Test creating entity without description (should be None).""" entity_data = EntityIn( @@ -179,7 +214,7 @@ async def test_entity_id_generation(entity_service): entity = await entity_service.create_entity(entity_data) assert entity.id # ID should be generated - assert "-test-entity" in entity.id # Should contain normalized name + assert "test/test_entity" == entity.id # Should contain normalized name async def test_create_entity_long_description(entity_service): """Test creating entity with a long description.""" diff --git a/tests/test_observation_service.py b/tests/test_observation_service.py index 907d3091..9727c58c 100644 --- a/tests/test_observation_service.py +++ b/tests/test_observation_service.py @@ -13,16 +13,9 @@ async def test_add_observation_success(observation_service, test_entity): content="New observation", context="test-context" ) - - entity_data = EntityIn( - name=test_entity.name, - entity_type=test_entity.entity_type, - id=test_entity.id, - observations=[] - ) - + # Act - observations = await observation_service.add_observations(entity_data, [observation_data]) + observations = await observation_service.add_observations(test_entity.id, [observation_data]) # Assert assert len(observations) == 1 @@ -40,15 +33,8 @@ async def test_add_observation_success(observation_service, test_entity): async def test_search_observations(observation_service, test_entity): """Test searching observations across entities.""" # Arrange - entity_data = EntityIn( - name=test_entity.name, - entity_type=test_entity.entity_type, - id=test_entity.id, - observations=[] - ) - await observation_service.add_observations( - entity_data, + test_entity.id, [ObservationIn(content="Unique test content"), ObservationIn(content="Other content")] ) @@ -63,15 +49,8 @@ async def test_search_observations(observation_service, test_entity): async def test_get_observations_by_context(observation_service, test_entity): """Test retrieving observations by context.""" # Arrange - entity_data = EntityIn( - name=test_entity.name, - entity_type=test_entity.entity_type, - id=test_entity.id, - observations=[] - ) - await observation_service.add_observations( - entity_data, + test_entity.id, [ObservationIn(content="Context observation", context="test-context"), ObservationIn(content="Other observation", context="other-context")] ) @@ -89,26 +68,15 @@ async def test_get_observations_by_context(observation_service, test_entity): async def test_observation_with_special_characters(observation_service, test_entity): """Test handling observations with special characters.""" content = "Test & observation with @#$% special chars!" - entity_data = EntityIn( - name=test_entity.name, - entity_type=test_entity.entity_type, - id=test_entity.id, - ) - observations = await observation_service.add_observations(entity_data, [ObservationIn(content=content)]) + observations = await observation_service.add_observations(test_entity.id, [ObservationIn(content=content)]) assert observations[0].content == content async def test_very_long_observation(observation_service, test_entity): """Test handling very long observation content.""" long_content = "Very long observation " * 100 # ~1800 characters - entity_data = EntityIn( - name=test_entity.name, - entity_type=test_entity.entity_type, - id=test_entity.id, - observations=[] - ) - observations = await observation_service.add_observations(entity_data, [ObservationIn(content=long_content)]) + observations = await observation_service.add_observations(test_entity.id, [ObservationIn(content=long_content)]) assert observations[0].content == long_content \ No newline at end of file diff --git a/tests/test_schemas.py b/tests/test_schemas.py index c9ea255d..0686b5f8 100644 --- a/tests/test_schemas.py +++ b/tests/test_schemas.py @@ -26,7 +26,6 @@ def test_entity_in_minimal(): assert entity.description is None assert entity.observations == [] assert entity.relations == [] - assert entity.id is not None # Should auto-generate def test_entity_in_complete(): """Test creating EntityIn with all fields.""" @@ -125,12 +124,6 @@ def test_create_entities_input(): with pytest.raises(ValidationError): CreateEntitiesInput.model_validate({"entities": []}) -def test_entity_id_generation(): - """Test ID generation for entities.""" - entity = EntityIn.model_validate({"name": "test entity", "entityType": "test"}) - assert entity.id.startswith(datetime.now().strftime("%Y%m%d")) - assert "test-entity" in entity.id - def test_snake_case_to_camel(): """Test conversion from snake_case to camelCase.""" data = {