From a5bb6f8b0ce22e9b8e1afd3dbbd7103eedbf1507 Mon Sep 17 00:00:00 2001 From: phernandez Date: Mon, 23 Dec 2024 10:01:04 -0600 Subject: [PATCH] refactor knowledge_service.create_relations --- src/basic_memory/schemas/__init__.py | 8 ---- .../services/knowledge/relations.py | 45 ++++++++++++------- tests/services/test_knowledge_service.py | 10 +---- 3 files changed, 30 insertions(+), 33 deletions(-) diff --git a/src/basic_memory/schemas/__init__.py b/src/basic_memory/schemas/__init__.py index ba67f396..0ecbe5f2 100644 --- a/src/basic_memory/schemas/__init__.py +++ b/src/basic_memory/schemas/__init__.py @@ -40,11 +40,7 @@ from basic_memory.schemas.response import ( CreateEntityResponse, SearchNodesResponse, OpenNodesResponse, - AddObservationsResponse, - CreateRelationsResponse, DeleteEntitiesResponse, - DeleteRelationsResponse, - DeleteObservationsResponse, ) # For convenient imports, export all models @@ -70,11 +66,7 @@ __all__ = [ "CreateEntityResponse", "SearchNodesResponse", "OpenNodesResponse", - "AddObservationsResponse", - "CreateRelationsResponse", "DeleteEntitiesResponse", - "DeleteRelationsResponse", - "DeleteObservationsResponse", # Delete Operations "DeleteEntitiesRequest", "DeleteRelationsRequest", diff --git a/src/basic_memory/services/knowledge/relations.py b/src/basic_memory/services/knowledge/relations.py index 5ef897ce..f2ce613d 100644 --- a/src/basic_memory/services/knowledge/relations.py +++ b/src/basic_memory/services/knowledge/relations.py @@ -4,6 +4,7 @@ from typing import Sequence, List from loguru import logger +from basic_memory.models import Entity as EntityModel from basic_memory.schemas import Relation as RelationSchema from basic_memory.services.exceptions import EntityNotFoundError from basic_memory.services.relation_service import RelationService @@ -17,31 +18,41 @@ class RelationOperations(EntityOperations): super().__init__(*args, **kwargs) self.relation_service = relation_service - async def create_relations(self, relations: List[RelationSchema]) -> Sequence[RelationSchema]: - """Create relations and update affected entity files.""" + async def create_relations(self, relations: List[RelationSchema]) -> Sequence[EntityModel]: + """Create relations and return updated entities.""" logger.debug(f"Creating {len(relations)} relations") - created = [] + updated_entities = [] + update_entity_ids = set() for relation in relations: try: # Create relation in DB - db_relation = await self.relation_service.create_relation(relation) + await self.relation_service.create_relation(relation) - # Update files with their new relations - for entity_id in [relation.from_id, relation.to_id]: - # 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) - await self.entity_service.update_entity(entity_id, {"checksum": checksum}) - - created.append(db_relation) + # Keep track of entities we need to update + update_entity_ids.add(relation.from_id) + update_entity_ids.add(relation.to_id) except Exception as e: logger.error(f"Failed to create relation: {e}") continue - return created \ No newline at end of file + # Get fresh copies of all updated entities + for entity_id in update_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 diff --git a/tests/services/test_knowledge_service.py b/tests/services/test_knowledge_service.py index 181cb091..f2c4b80c 100644 --- a/tests/services/test_knowledge_service.py +++ b/tests/services/test_knowledge_service.py @@ -78,14 +78,8 @@ async def test_create_relations(knowledge_service: KnowledgeService, entity_serv ) ] - created = await knowledge_service.create_relations(relations) - assert len(created) == 1 - - # Verify relation was created - relation = created[0] - assert relation.from_id == entity1.id - assert relation.to_id == entity2.id - assert relation.relation_type == "test_relation" + updated_entities = await knowledge_service.create_relations(relations) + assert len(updated_entities) == 2 # Verify files were updated for entity_id in [entity1.id, entity2.id]: