From cbbfa6e637c29ac3bd2fd21ef2bdde390afccb8f Mon Sep 17 00:00:00 2001 From: phernandez Date: Mon, 23 Dec 2024 10:11:56 -0600 Subject: [PATCH] refactor knowledge_service to return Entity objects on updates --- src/basic_memory/api/routers/knowledge.py | 15 ++--- .../services/knowledge/observations.py | 32 +++++++++-- .../services/knowledge/relations.py | 56 ++++++++++++++++--- tests/services/test_knowledge_service.py | 12 ++-- 4 files changed, 89 insertions(+), 26 deletions(-) diff --git a/src/basic_memory/api/routers/knowledge.py b/src/basic_memory/api/routers/knowledge.py index 5cf6997d..e94a5492 100644 --- a/src/basic_memory/api/routers/knowledge.py +++ b/src/basic_memory/api/routers/knowledge.py @@ -45,8 +45,6 @@ async def create_relations( data: CreateRelationsRequest, knowledge_service: KnowledgeServiceDep ) -> CreateEntityResponse: """Create relations between entities.""" - - # TODO knowledge_service.create_relations should return updated Entities updated_entities = await knowledge_service.create_relations(data.relations) return CreateEntityResponse( entities=[EntityResponse.model_validate(entity) for entity in updated_entities] @@ -119,16 +117,14 @@ async def delete_observations( ) -> EntityResponse: """Delete observations from an entity.""" entity_id = data.entity_id - - # TODO add knowledge_service.delete_observations updated_entity = await knowledge_service.delete_observations(entity_id, data.deletions) return EntityResponse.model_validate(updated_entity) -@router.post("/relations/delete", response_model=EntityResponse) +@router.post("/relations/delete", response_model=CreateEntityResponse) async def delete_relations( data: DeleteRelationsRequest, knowledge_service: KnowledgeServiceDep -) -> EntityResponse: +) -> CreateEntityResponse: """Delete relations between entities.""" to_delete = [ { @@ -138,6 +134,7 @@ async def delete_relations( } for relation in data.relations ] - # TODO add knowledge_service.delete_relations - updated_entity = await knowledge_service.delete_relations(to_delete) - return EntityResponse.model_validate(updated_entity) + updated_entities = await knowledge_service.delete_relations(to_delete) + return CreateEntityResponse( + entities=[EntityResponse.model_validate(entity) for entity in updated_entities] + ) \ No newline at end of file diff --git a/src/basic_memory/services/knowledge/observations.py b/src/basic_memory/services/knowledge/observations.py index dc3c0017..1f49a9f7 100644 --- a/src/basic_memory/services/knowledge/observations.py +++ b/src/basic_memory/services/knowledge/observations.py @@ -24,7 +24,7 @@ class ObservationOperations(RelationOperations): logger.debug(f"Adding observations to entity {entity_id}") try: - # Get updated entity + # Get entity to update entity = await self.entity_service.get_entity(entity_id) if not entity: raise EntityNotFoundError(f"Entity not found: {entity_id}") @@ -32,15 +32,39 @@ class ObservationOperations(RelationOperations): # Add observations to DB await self.observation_service.add_observations(entity_id, observations, context) + # Get updated entity + updated_entity = await self.entity_service.get_entity(entity_id) + + # Write updated file and checksum + checksum = await self.write_entity_file(entity) + await self.entity_service.update_entity(entity_id, {"checksum": checksum}) + + return updated_entity + + except Exception as e: + logger.error(f"Failed to add observations: {e}") + raise + + async def delete_observations(self, entity_id: int, observations: List[str]) -> EntityModel: + """Delete observations from entity and update its file.""" + logger.debug(f"Deleting observations from entity {entity_id}") + + try: + # Get updated entity + entity = await self.entity_service.get_entity(entity_id) + if not entity: + raise EntityNotFoundError(f"Entity not found: {entity_id}") + + # Delete observations from DB + await self.observation_service.delete_observations(entity_id, observations) + # Write updated file checksum = await self.write_entity_file(entity) - - # Update checksum in DB await self.entity_service.update_entity(entity_id, {"checksum": checksum}) # Get final entity with all updates return await self.entity_service.get_entity(entity_id) except Exception as e: - logger.error(f"Failed to add observations: {e}") + logger.error(f"Failed to delete observations: {e}") raise diff --git a/src/basic_memory/services/knowledge/relations.py b/src/basic_memory/services/knowledge/relations.py index f2ce613d..a4eb90a8 100644 --- a/src/basic_memory/services/knowledge/relations.py +++ b/src/basic_memory/services/knowledge/relations.py @@ -1,6 +1,6 @@ """Relation operations for knowledge service.""" -from typing import Sequence, List +from typing import Sequence, List, Dict, Any from loguru import logger @@ -21,8 +21,8 @@ class RelationOperations(EntityOperations): async def create_relations(self, relations: List[RelationSchema]) -> Sequence[EntityModel]: """Create relations and return updated entities.""" logger.debug(f"Creating {len(relations)} relations") - updated_entities = [] - update_entity_ids = set() + created_entities = [] + updated_entity_ids = set() for relation in relations: try: @@ -30,15 +30,15 @@ class RelationOperations(EntityOperations): await self.relation_service.create_relation(relation) # Keep track of entities we need to update - update_entity_ids.add(relation.from_id) - update_entity_ids.add(relation.to_id) + updated_entity_ids.add(relation.from_id) + updated_entity_ids.add(relation.to_id) except Exception as e: logger.error(f"Failed to create relation: {e}") continue # Get fresh copies of all updated entities - for entity_id in update_entity_ids: + for entity_id in updated_entity_ids: try: # Get fresh entity entity = await self.entity_service.get_entity(entity_id) @@ -49,10 +49,50 @@ class RelationOperations(EntityOperations): checksum = await self.write_entity_file(entity) updated = await self.entity_service.update_entity(entity_id, {"checksum": checksum}) - updated_entities.append(updated) + created_entities.append(updated) except Exception as e: logger.error(f"Failed to update entity {entity_id}: {e}") continue - return updated_entities + return created_entities + + async def delete_relations(self, to_delete: List[Dict[str, Any]]) -> Sequence[EntityModel]: + """Delete relations and return all updated entities.""" + logger.debug(f"Deleting {len(to_delete)} relations") + updated_entity_ids = set() + + try: + # Delete relations from DB + for relation in to_delete: + updated_entity_ids.add(relation["from_id"]) + updated_entity_ids.add(relation["to_id"]) + + deleted = await self.relation_service.delete_relations(to_delete) + if not deleted: + logger.warning("No relations were deleted") + + # Get fresh copies of all updated entities + updated_entities = [] + for entity_id in updated_entity_ids: + try: + # Get fresh entity + entity = await self.entity_service.get_entity(entity_id) + if not entity: + raise EntityNotFoundError(f"Entity not found: {entity_id}") + + # Write updated file + checksum = await self.write_entity_file(entity) + updated = await self.entity_service.update_entity(entity_id, {"checksum": checksum}) + + updated_entities.append(updated) + + except Exception as e: + logger.error(f"Failed to update entity {entity_id}: {e}") + continue + + return updated_entities + + except Exception as e: + logger.error(f"Failed to delete relations: {e}") + raise \ No newline at end of file diff --git a/tests/services/test_knowledge_service.py b/tests/services/test_knowledge_service.py index f2c4b80c..d5eafa90 100644 --- a/tests/services/test_knowledge_service.py +++ b/tests/services/test_knowledge_service.py @@ -99,15 +99,17 @@ async def test_add_observations(knowledge_service: KnowledgeService): # Add observations observations = ["Test observation 1", "Test observation 2"] - updated = await knowledge_service.add_observations(entity.id, observations, "Test context") + updated_entity = await knowledge_service.add_observations( + entity.id, observations, "Test context" + ) # Verify observations in DB - assert len(updated.observations) == 2 - assert updated.observations[0].content == "Test observation 1" - assert updated.observations[1].content == "Test observation 2" + assert len(updated_entity.observations) == 2 + assert updated_entity.observations[0].content == "Test observation 1" + assert updated_entity.observations[1].content == "Test observation 2" # Verify file was updated - file_path = knowledge_service.get_entity_path(updated) + file_path = knowledge_service.get_entity_path(updated_entity) content, _ = await knowledge_service.file_service.read_file(file_path) for obs in observations: assert obs in content