diff --git a/src/basic_memory/repository/repository.py b/src/basic_memory/repository/repository.py index 0b0bb192..381f15c0 100644 --- a/src/basic_memory/repository/repository.py +++ b/src/basic_memory/repository/repository.py @@ -1,6 +1,6 @@ """Base repository implementation.""" -from datetime import datetime, timezone +from datetime import datetime from typing import Type, Optional, Any, Sequence, TypeVar, List from loguru import logger @@ -65,7 +65,7 @@ class Repository[T: Base]: :param model: the model to add :return: the added model instance """ - async with db.scoped_session(self.session_maker) as session: + async with db.scoped_session(self.session_maker) as session: session.add(model) await session.flush() @@ -184,7 +184,7 @@ class Repository[T: Base]: async with db.scoped_session(self.session_maker) as session: # Only include valid columns that are provided in entity_data model_data = self.get_model_data(data) - model = self.Model(**model_data) + model = self.Model(**model_data) session.add(model) await session.flush() @@ -209,7 +209,7 @@ class Repository[T: Base]: return await self.select_by_ids(session, [model.id for model in model_list]) # pyright: ignore [reportAttributeAccessIssue] - async def update(self, entity_id: int, entity_data: dict) -> Optional[T]: + async def update(self, entity_id: int, entity_data: dict | T) -> Optional[T]: """Update an entity with the given data.""" logger.debug(f"Updating {self.Model.__name__} {entity_id} with data: {entity_data}") async with db.scoped_session(self.session_maker) as session: @@ -219,10 +219,15 @@ class Repository[T: Base]: ) entity = result.scalars().one() - for key, value in entity_data.items(): - if key in self.valid_columns: - setattr(entity, key, value) - + if isinstance(entity_data, dict): + for key, value in entity_data.items(): + if key in self.valid_columns: + setattr(entity, key, value) + + elif isinstance(entity_data, self.Model): + for column in self.Model.__table__.columns.keys(): + setattr(entity, column, getattr(entity_data, column)) + await session.flush() # Make sure changes are flushed await session.refresh(entity) # Refresh diff --git a/src/basic_memory/services/entity_service.py b/src/basic_memory/services/entity_service.py index 302d348a..30e2e40b 100644 --- a/src/basic_memory/services/entity_service.py +++ b/src/basic_memory/services/entity_service.py @@ -220,9 +220,6 @@ class EntityService(BaseService[EntityModel]): # Clear observations for entity await self.observation_repository.delete_by_fields(entity_id=db_entity.id) - # update values from markdown - db_entity = entity_model_from_markdown(file_path, markdown, db_entity) - # add new observations observations = [ Observation( @@ -236,6 +233,9 @@ class EntityService(BaseService[EntityModel]): ] await self.observation_repository.add_all(observations) + # update values from markdown + db_entity = entity_model_from_markdown(file_path, markdown, db_entity) + # checksum value is None == not finished with sync db_entity.checksum = None @@ -243,16 +243,7 @@ class EntityService(BaseService[EntityModel]): # checksum value is None == not finished with sync return await self.repository.update( db_entity.id, - { - "title": db_entity.title, - "entity_type": db_entity.entity_type, - "entity_metadata": db_entity.entity_metadata, - # TODO redo update, get created, modified from file - "created_at": markdown.frontmatter.created, - "updated_at": markdown.frontmatter.modified, - # Mark as incomplete - "checksum": None, - }, + db_entity, ) async def update_entity_relations( diff --git a/tests/repository/test_repository.py b/tests/repository/test_repository.py index a8685ede..7b4cff72 100644 --- a/tests/repository/test_repository.py +++ b/tests/repository/test_repository.py @@ -123,6 +123,36 @@ async def test_delete_by_ids(repository): assert await repository.find_by_id(ids_to_delete[2]) is None +@pytest.mark.asyncio +async def test_update(repository): + """Test finding entities modified since a timestamp.""" + # Create initial test data + instance = TestModel(id="test_add", name="Test Add") + await repository.add(instance) + + instance = TestModel(id="test_add", name="Updated") + + # Find recently modified + modified = await repository.update(instance.id, {"name": "Updated"}) + assert modified is not None + assert modified.name == "Updated" + + +@pytest.mark.asyncio +async def test_update_model(repository): + """Test finding entities modified since a timestamp.""" + # Create initial test data + instance = TestModel(id="test_add", name="Test Add") + await repository.add(instance) + + instance.name = "Updated" + + # Find recently modified + modified = await repository.update(instance.id, instance) + assert modified is not None + assert modified.name == "Updated" + + @pytest.mark.asyncio async def test_find_modified_since(repository): """Test finding entities modified since a timestamp."""